L1–122Normalizer 类(按 dataset_name 查表归一化动作)、辅助函数(分布式打印、正弦位置编码、卷积采样),以及 normalize_data 方法核心实现。支持 per-robot/per-dataset 统计,LIBERO 中按前7维 dof_mask 截断。
1
import torch3
import torch.nn as nn
5
from typing import Union
6
import math8
from diffusers.schedulers.scheduling_ddpm import DDPMScheduler
9
from torch.distributions import Beta
导入核心库:torch 主框架、nn 模块、类型提示(Union)、math 模块、DDPM 调度器(扩散模型用)、Beta 分布(flow matching 训练时采样时间步 t~Beta(1.5,1) 用)。
12
def print_rank_last(message):
13
"""If distributed is initialized, print only on last rank."""14
if torch.distributed.is_initialized():15
if torch.distributed.get_rank() == (torch.distributed.get_world_size() - 1):
16
print(message, flush=True)
17
else:18
print(message, flush=True)
print_rank_last 辅助函数:分布式训练时只在最后一个 rank 打印,避免多进程重复输出。检查 torch.distributed 初始化状态;若已初始化则获取 rank 并比较 get_world_size()-1,只有最后进程执行 print。
空行+Normalizer 类定义开始,继承 nn.Module。该类管理动作空间的归一化统计量(per-robot/per-dataset),支持 normalize/unnormalize 两个方向,是 action_head 的配套模块。
@classmethod 装饰器+from_ckpt 类方法签名。从保存的 checkpoint 文件恢复 Normalizer 实例(离线构建模式),对标标准 pytorch_lightning 的 from_checkpoint 模式。
25
nn.Module.__init__(instance)27
instance.min = nn.ParameterDict()
28
instance.delta = nn.ParameterDict()
29
instance.min_key = "min"
30
instance.delta_key = "delta"
创建空实例并手动调用 nn.Module.__init__ 初始化(绕过 __init__);初始化两个 ParameterDict 容器 min/delta 和对应的 key 名('min','delta');稍后逐个 robot 填充。实战:LIBERO 中 checkpoint 需恢复多套 min/max,各 dataset 不同。
32
ckpt = torch.load(ckpt_path, map_location="cpu")
34
for key, value in ckpt.items():
35
# Parse key: "min.robot_name" -> prefix="min", name="robot_name"36
try:37
prefix, name = key.split(".", 1)
加载 checkpoint 文件为 dict (map_location='cpu' 避免 GPU 内存压力);遍历 checkpoint 中每个 key-value 对,格式为 'prefix.name'(如 'min.libero_10'),用 split('.', 1) 拆分前缀和名称,为后续路由到对应 ParameterDict。
38
if hasattr(instance, prefix):
39
getattr(instance, prefix)[name] = nn.Parameter(40
value, requires_grad=False41
)
42
print("prefix", prefix)
43
print("name", name)
44
except ValueError:45
continue条件判断:若前缀名(min/delta)与实例已有属性匹配,则 getattr 取出该 ParameterDict,用 nn.Parameter(value, requires_grad=False) 包装张量后存入;调试打印 prefix 和 name。异常处理 ValueError 捕获格式不匹配的 key 并跳过(健壮性)。
49
def __init__(
50
self, action_statistic_dof, dof_config, min_key="min", delta_key="delta"
51
):
52
super(Normalizer, self).__init__()
__init__ 方法签名:接收 action_statistic_dof(dict of dict,格式 {robot_name: {dof_key: {min_key: [...], delta_key: [...]}}}),dof_config(dict 映射 dof 名到维数),以及 min_key/delta_key 字符串(默认 'min'/'delta');调用 super().__init__() 初始化 nn.Module。
存储 min_key 和 delta_key 为实例属性,后续 normalize_data 通过这两个 key 从 ParameterDict 查索,支持自定义键名(如备选命名约定)。
57
action_statistic = {}58
for robot_name in action_statistic_dof.keys():
59
action_statistic[robot_name] = {}60
all_dof_min = []
61
all_dof_delta = []
初始化空 action_statistic dict 和两个空列表 all_dof_min/all_dof_delta;遍历 action_statistic_dof.keys() 获取所有 robot_name,为每个 robot 创建空字典,初始化列表用于累积该 robot 所有 dof 的 min/delta 值。
62
for k in dof_config:
63
if k in action_statistic_dof[robot_name]:
64
if (65
min_key in action_statistic_dof[robot_name][k]66
and delta_key in action_statistic_dof[robot_name][k]
67
):
68
all_dof_min.extend(action_statistic_dof[robot_name][k][min_key])
69
all_dof_delta.extend(
70
action_statistic_dof[robot_name][k][delta_key]
71
)
内层 for 循环遍历 dof_config 的所有关键字(arm/gripper/height/base 等);若该 dof k 在该 robot 的 action_statistic_dof 中存在,且同时包含 min_key 和 delta_key,则 extend 两个列表。张量形状变化:单个 min/delta 可能是 [7] 或 [1],extend 后逐个 robot 累积成 [D] (D=20)。实战关键:LIBERO 只用前7维,dof_config 中 arm=7,gripper/height/base 等贡献剩余 13 维。
72
else:73
if robot_name == "x2_normal" or "libero" in robot_name:
74
print_rank_last(
75
f"Normalizer (Warning): min_key {min_key} or delta_key {delta_key} "
76
)
77
print_rank_last(
78
f"not in action_statistic_dof[{robot_name}][{k}], use default min 0.0 and delta 1.0"
79
)
80
all_dof_min.extend([0.0] * dof_config[k])81
all_dof_delta.extend([1.0] * dof_config[k])else 分支:若 min_key 或 delta_key 缺失,使用默认值 0.0 和 1.0(相当于不进行归一化,原样通过);仅在 robot_name='x2_normal' 或 'libero' 子串时打印警告,避免其他 robot 的无关输出。实战:LIBERO 数据可能某些 dof 没有统计,降级为恒等映射。
82
else:83
if robot_name == "x2_normal" or "libero" in robot_name:
84
print_rank_last(
85
f"Normalizer (Warning): Action {k} not in action_statistic_dof for {robot_name}, use default min 0.0 and delta 1.0"
86
)
87
all_dof_min.extend([0.0] * dof_config[k])88
all_dof_delta.extend([1.0] * dof_config[k])else 分支:若 action key k 不在 action_statistic_dof[robot_name] 中,说明该 dof 类型对该 robot 完全未知,extend [0.0]*dof_config[k] 和 [1.0]*dof_config[k];同样仅 x2/libero robot 打印警告。确保最终列表长度与 dof_config 总和一致(形状对齐)。
89
all_dof_min = torch.tensor(all_dof_min)
90
all_dof_delta = torch.tensor(all_dof_delta)
91
action_statistic[robot_name][min_key] = all_dof_min
92
action_statistic[robot_name][delta_key] = all_dof_delta
将累积的列表转换为 torch tensor,然后存回 action_statistic 字典中(替换列表)。张量形状:all_dof_min/delta 都是一维 tensor [D],D=sum(dof_config.values()) 通常为 20。空行分隔。
94
self.min = nn.ParameterDict(95
{96
k: nn.Parameter(action_statistic[k][min_key], requires_grad=False)97
for k in action_statistic.keys()
98
}
99
)
创建 self.min 为 nn.ParameterDict,字典推导式遍历 action_statistic 中每个 robot,将对应的 min tensor 用 nn.Parameter(..., requires_grad=False) 包装,确保不更新(推理/评估时固定)。访问模式:self.min[dataset_name] 返回 Parameter,形状 [D]。
100
self.delta = nn.ParameterDict(101
{102
k: nn.Parameter(action_statistic[k][delta_key], requires_grad=False)103
for k in action_statistic.keys()
104
}
105
)
创建 self.delta 为 nn.ParameterDict,同 self.min 的逻辑,存储缩放因子。访问模式:self.delta[dataset_name] 返回 Parameter [D]。together min/delta 定义了该 robot 的完整归一化范围 [min, min+delta]。
107
for k, v in action_statistic.items():
108
print_rank_last(
109
f"Normalizer: {k} min {action_statistic[k][min_key]} delta {action_statistic[k][delta_key]}"
110
)
空行+统计输出循环:遍历所有 robot,打印每个 robot 的 min 和 delta 统计值(调试/验证)。message 跨多行打印确保可读性。实战:快速检查是否加载了正确的统计值。
112
def normalize_data(self, xs, dataset_names):
113
new_xs = []
114
dataset_names = [name for name in dataset_names if name != "x2_multimodal"]
115
for x, dataset_name in zip(xs, dataset_names):
116
x = (x - self.min[dataset_name]) / (self.delta[dataset_name])
117
x = x * 2 - 1
118
x = torch.clamp(x, -1, 1)
119
new_xs.append(x)
120
new_xs = torch.stack(new_xs)
121
return new_xs空行+normalize_data 方法定义。接收动作张量列表 xs 和对应的 dataset_names;过滤掉 'x2_multimodal' dataset(多模态混合 dataset 无单独统计);对每个 (x, dataset_name) 对,执行 (x-min)/delta * 2 - 1 将原始范围 [min, min+delta] 映射到 [-1, 1],最后 clamp 防越界。输出:归一化后的张量列表,torch.stack 合并为 batch 张量 (B,D)。实战关键:LIBERO 多个 dataset(libero_10/libero_goal 等)各有统计,此处按 dataset_name 查表实现 per-dataset 归一化,后续训练的 target action 都要先 normalize;推理时也要反向 unnormalize 才能执行。
L123–251Normalizer.unnormalize_data 反归一化方法(123-139)+ 位置编码/下采样/上采样/残差块等 UNet 基础组件类定义(142-251 起)。
123
def unnormalize_data(self, xs, dataset_names, dof_mask=None):
Normalizer 的反归一化方法签名。将 flow matching 生成的 [-1,1] 范围速度场转换回原始动作空间。被 generate_flow_action 的后处理阶段调用。
124
new_xs = []
125
dataset_names = [name for name in dataset_names if name != "x2_multimodal"]
126
dof_mask = dof_mask if dof_mask is not None else [None] * len(xs)
初始化输出列表 new_xs;过滤 dataset_names 排除 x2_multimodal(仅保留 libero/x2_normal 等实际数据集);dof_mask 默认为 [None, None, ...] 表示不做维度掩码,使用完整 action space。
zip 遍历动作张量 xs、数据集名 dataset_names、dof_mask。第 128 行从 [-1,1] 缩放到 [0,1]:(x+1)/2 是 flow matching 反归一化的第一步。
129
if mask is not None:
130
mask = mask[0].bool()131
action_space_delta = self.delta[dataset_name][mask]132
action_space_min = self.min[dataset_name][mask]133
else:134
action_space_delta = self.delta[dataset_name]135
action_space_min = self.min[dataset_name]根据 dof_mask 判断是否做部分维度反归一化。若 mask 非空(形如 (1,D) bool 张量,D=20 总 dof),取 mask=True 的维度对应的 delta/min;否则用完整的 delta/min。LIBERO 实战:mask 的前 7 位 True(对应左右臂 7 维),后 13 位 False;此时只对前 7 维反归一化。
136
x = x * action_space_delta + action_space_min
137
new_xs.append(x)
138
new_xs = torch.stack(new_xs)
第二步反归一化:x = x * delta + min,将 [0,1] 映射回原始值域(例如 LIBERO 的关节速度范围);stack 合并所有样本的动作,返回 (B, seq_len, action_dim) 或 (B, seq_len, 7)。
139
return new_xs返回反归一化后的动作张量。
142
class SinusoidalPosEmb(nn.Module):
SinusoidalPosEmb 类:正弦位置编码器,用于将 diffusion 时间步 t(或 flow matching 的 s)编码为连续向量。被 ConditionalUnet1D 的 diffusion_step_encoder 调用(第 279 行)。
143
def __init__(self, dim: int, min_period: float = 4e-3, max_period: float = 4.0):
144
super().__init__()
145
self.dim = dim146
if dim % 2 != 0:
147
raise ValueError(f"embedding_dim ({dim}) must be divisible by 2")
148
self.min_period = min_period149
self.max_period = max_period__init__ 初始化:dim 必须偶数(sin/cos 各占一半),min_period/max_period 控制频率范围(默认 4e-3~4.0)。校验 dim%2==0,否则抛 ValueError。
151
def forward(self, x):
152
device = x.device
153
half_dim = self.dim // 2
154
emb = math.log(10000) / (half_dim - 1)
155
emb = torch.exp(
156
torch.arange(half_dim, device=device, dtype=torch.float32) * -emb
157
)
158
emb = x[:, None] * emb[None, :]
159
emb = torch.cat((emb.sin(), emb.cos()), dim=-1)160
return embforward 方法:输入 x 形如 (B,) 或 (B,1) 的时间步;计算频率向量 emb=exp(arange(half_dim)*(-log(10000)/(half_dim-1)))(来自 Transformer 位置编码);x[:, None]*emb[None, :] 得角频率 (B, half_dim);最后 cat(sin,cos) 返回 (B, dim) 的正弦编码。
空行。
163
class Downsample1d(nn.Module):
164
def __init__(self, dim):
165
super().__init__()
166
self.conv = nn.Conv1d(dim, dim, 3, 2, 1)
168
def forward(self, x):
169
return self.conv(x)
Downsample1d:1D 卷积下采样(kernel=3, stride=2, padding=1),保持通道数不变。用于 UNet encoder 路径压缩序列长度。
空行。
172
class Upsample1d(nn.Module):
173
def __init__(self, dim):
174
super().__init__()
175
self.conv = nn.ConvTranspose1d(dim, dim, 4, 2, 1)
177
def forward(self, x):
178
return self.conv(x)
Upsample1d:转置卷积上采样(kernel=4, stride=2, padding=1),保持通道数不变。用于 UNet decoder 路径恢复序列长度。
空行。
Conv1dBlock 类及 docstring:简单的 Conv1d→GroupNorm→Mish 堆栈,无跨层连接。是残差块的基础单元。
186
def __init__(self, inp_channels, out_channels, kernel_size, n_groups=8):
187
super().__init__()
189
self.block = nn.Sequential(190
nn.Conv1d(
191
inp_channels, out_channels, kernel_size, padding=kernel_size // 2192
),
193
nn.GroupNorm(n_groups, out_channels),
194
nn.Mish(),
195
)
__init__:Conv1d(inp_channels→out_channels, kernel=kernel_size, padding=kernel_size//2)→GroupNorm(n_groups=8, out_channels)→Mish()。padding 保持序列长度不变(same padding)。
201
class ConditionalResidualBlock1D(nn.Module):
202
def __init__(self, in_channels, out_channels, cond_dim, kernel_size=3, n_groups=8):
ConditionalResidualBlock1D:条件残差块,支持 FiLM 调制(Feature-wise Linear Modulation)。在两个 Conv1dBlock 之间用外部条件动态调整缩放和偏置,增强条件化能力。
203
super().__init__()
205
self.blocks = nn.ModuleList(206
[
207
Conv1dBlock(in_channels, out_channels, kernel_size, n_groups=n_groups),
208
Conv1dBlock(out_channels, out_channels, kernel_size, n_groups=n_groups),
209
]
210
)
__init__:定义两个 Conv1dBlock 堆栈(in_channels→out_channels→out_channels)。cond_dim:条件向量维度(通常为 diffusion_step_embed_dim(256) + global_cond_dim(obs_horizon*obs_dim))。
212
# FiLM modulation https://arxiv.org/abs/1709.07871213
# predicts per-channel scale and bias214
cond_channels = out_channels * 2215
self.out_channels = out_channels216
self.cond_encoder = nn.Sequential(217
nn.Mish(),
218
nn.Linear(cond_dim, cond_channels),
219
nn.Dropout(0.1),220
nn.Unflatten(-1, (-1, 1)),
221
)
FiLM 调制器(cond_encoder):Mish()→Linear(cond_dim→out_channels*2)→Dropout(0.1)→Unflatten(-1,(-1,1)),最终输出形状 (B, out_channels*2, 1),一半用作 scale,一半用作 bias。
223
# make sure dimensions compatible224
self.residual_conv = (225
nn.Conv1d(in_channels, out_channels, 1)226
if in_channels != out_channels227
else nn.Identity()228
)
残差连接处理:若 in_channels != out_channels 则添加 Conv1d(in_channels→out_channels, kernel=1) 做维度匹配;否则用 Identity()。
230
def forward(self, x, cond):
231
"""232
x : [ batch_size x in_channels x horizon ]233
cond : [ batch_size x cond_dim]235
returns:236
out : [ batch_size x out_channels x horizon ]237
"""forward 方法 docstring:x 形如 (B, in_channels, horizon),cond 形如 (B, cond_dim),返回 (B, out_channels, horizon)。horizon 是动作序列长度(通常对应 action_horizon)。
238
out = self.blocks[0](x)
239
embed = self.cond_encoder(cond)241
embed = embed.reshape(embed.shape[0], 2, self.out_channels, 1)
242
scale = embed[:, 0, ...]243
bias = embed[:, 1, ...]244
out = scale * out + bias
第一个 Conv1dBlock 得到 out(B, out_channels, horizon);cond_encoder 生成 embed(B, out_channels*2, 1);reshape 成 (B, 2, out_channels, 1) 分出 scale 和 bias;FiLM 调制 out = scale*out + bias(逐通道仿射变换)。
第二个 Conv1dBlock 处理调制后的 out;加上残差连接(residual_conv(x))返回最终输出。residual_conv 处理通道数不匹配(在 __init__ 已定义好)。
251
class ConditionalUnet1D(nn.Module):
ConditionalUnet1D 类定义:U-Net 架构用于 flow matching 的去噪。支持 diffusion_step 编码和 global observation 的条件化,是整个 action_head 的核心降噪网络。
L252–370ConditionalUnet1D.__init__ 方法:构建用于扩散模型的条件 U-Net 架构,支持按 diffusion timestep 和全局观察条件调制,包含时间步编码器、中间层、下采样块、上采样块、最终卷积输出层。
252
def __init__(
253
self,254
input_dim,
255
global_cond_dim,
256
diffusion_step_embed_dim=256,257
down_dims=[256, 512, 1024],
258
# down_dims=[512, 1024, 2048],259
kernel_size=5,260
n_groups=8,261
):
ConditionalUnet1D.__init__ 函数签名:input_dim 是动作维度,global_cond_dim 是观察条件拼接后的维度,diffusion_step_embed_dim 是扩散步 embedding 维度(默认256),down_dims 是各层通道数列表(默认[256,512,1024]),kernel_size 是卷积核大小(默认5),n_groups 是 GroupNorm 分组数(默认8)。
262
"""263
input_dim: Dim of actions.264
global_cond_dim: Dim of global conditioning applied with FiLM265
in addition to diffusion step embedding. This is usually obs_horizon * obs_dim266
diffusion_step_embed_dim: Size of positional encoding for diffusion iteration k267
down_dims: Channel size for each UNet level.268
The length of this array determines numebr of levels.269
kernel_size: Conv kernel size270
n_groups: Number of groups for GroupNorm271
"""Docstring:详细说明各参数用途。input_dim 对应动作维度(LIBERO 场景为 7 维);global_cond_dim 通常为观察窗口长度×观察维度(如 obs_horizon * obs_dim);diffusion_step_embed_dim 用于对扩散迭代步数 k 进行正弦位置编码;down_dims 列表长度决定了 U-Net 的深度层数;kernel_size 和 n_groups 控制卷积和归一化的细节配置。
空行后 super().__init__():初始化 nn.Module 基类,使当前类继承 PyTorch 模块管理机制(包含分隔空行与初始化语句)。
构建维度映射:all_dims = [input_dim] + down_dims,如 input_dim=20, down_dims=[256,512,1024] 则 all_dims=[20,256,512,1024];start_dim = down_dims[0] = 256,用作最终卷积前的通道数。实战中 input_dim=20 是完整动作空间(包含 7 DOF 动作、头部、高度、车),LIBERO 后续会用 dof_mask 掩蔽不用的维度。
空行。
277
dsed = diffusion_step_embed_dim
278
diffusion_step_encoder = nn.Sequential(
279
SinusoidalPosEmb(dsed),
280
nn.Linear(dsed, dsed * 4),281
nn.Mish(),
282
nn.Linear(dsed * 4, dsed),283
)
创建扩散步编码器 diffusion_step_encoder:SinusoidalPosEmb 对 diffusion step t 做正弦位置编码 (B,) → (B,256),再经两层 Linear(256→1024→256) 和 Mish 激活,输出 (B,256) 维的时间条件向量,后续与全局观察条件拼接。
284
cond_dim = dsed + global_cond_dim
计算条件维度 cond_dim = diffusion_step_embed_dim + global_cond_dim = 256 + global_cond_dim,用于 FiLM 调制中的条件投影。每个 ConditionalResidualBlock1D 会用 cond_dim 通过线性层投出 per-channel scale 和 bias。
空行。
286
in_out = list(zip(all_dims[:-1], all_dims[1:]))
创建维度配对表 in_out = [(20,256), (256,512), (512,1024)],用于循环构建下采样和上采样块,每对表示该级的输入和输出通道数。
287
mid_dim = all_dims[-1]mid_dim = all_dims[-1] = 1024,为 U-Net 最底层(瓶颈层)的通道数。
288
self.mid_modules = nn.ModuleList(289
[
290
ConditionalResidualBlock1D(
291
mid_dim,
292
mid_dim,
293
cond_dim=cond_dim,
294
kernel_size=kernel_size,
295
n_groups=n_groups,
296
),
297
ConditionalResidualBlock1D(
298
mid_dim,
299
mid_dim,
300
cond_dim=cond_dim,
301
kernel_size=kernel_size,
302
n_groups=n_groups,
303
),
304
]
305
)
创建 self.mid_modules(瓶颈层的两个残差块):两个 ConditionalResidualBlock1D,均为 (1024,1024,cond_dim=...) 配置,处理最底层特征 (B,1024,T_min),接收拼接后的时间+全局条件。此处不做上下采样,只做特征变换。
空行。
307
down_modules = nn.ModuleList([])
308
for ind, (dim_in, dim_out) in enumerate(in_out):
309
is_last = ind >= (len(in_out) - 1)
310
down_modules.append(
311
nn.ModuleList(
312
[
313
ConditionalResidualBlock1D(
314
dim_in,
315
dim_out,
316
cond_dim=cond_dim,
317
kernel_size=kernel_size,
318
n_groups=n_groups,
319
),
320
ConditionalResidualBlock1D(
321
dim_out,
322
dim_out,
323
cond_dim=cond_dim,
324
kernel_size=kernel_size,
325
n_groups=n_groups,
326
),
327
Downsample1d(dim_out) if not is_last else nn.Identity(),
328
]
329
)
330
)
创建下采样模块列表 down_modules:对每对 (dim_in, dim_out) 循环(如第 0 轮 (20,256)、第 1 轮 (256,512) 等),每轮创建一个 ModuleList 包含两个 ConditionalResidualBlock1D(做通道变换 dim_in→dim_out 再 dim_out→dim_out)和一个 Downsample1d(步长 2 的卷积)或 Identity(最后一级不下采样)。下采样路径 (B,C,T) → (B,C,T/2),特征保存至 h 列表以供上采样跳连。实战注意:T 初始为 action_horizon 长度(如 32),经过 3 级下采样变为 T/8。
空行。
332
up_modules = nn.ModuleList([])
333
for ind, (dim_in, dim_out) in enumerate(reversed(in_out[1:])):
334
is_last = ind >= (len(in_out) - 1)
335
up_modules.append(
336
nn.ModuleList(
337
[
338
ConditionalResidualBlock1D(
339
dim_out * 2,340
dim_in,
341
cond_dim=cond_dim,
342
kernel_size=kernel_size,
343
n_groups=n_groups,
344
),
345
ConditionalResidualBlock1D(
346
dim_in,
347
dim_in,
348
cond_dim=cond_dim,
349
kernel_size=kernel_size,
350
n_groups=n_groups,
351
),
352
Upsample1d(dim_in) if not is_last else nn.Identity(),
353
]
354
)
355
)
创建上采样模块列表 up_modules:对反向的维度对循环(in_out[1:] 反向,如先 (512,256)、再 (256,20)),每轮创建 ModuleList 含两个残差块和一个上采样。第一个残差块处理拼接后的跳连特征,输入通道数为 dim_out*2(跳连的 dim_out + 来自下一层的 dim_out,故相加后翻倍),输出到 dim_in;第二个残差块输出维度仍为 dim_in;Upsample1d 或 Identity 用于恢复空间分辨率 (B,C,T) → (B,C,T*2)。实战:最终上采样后恢复到原始 action_horizon。
空行。
357
final_conv = nn.Sequential(
358
Conv1dBlock(start_dim, start_dim, kernel_size=kernel_size),
359
nn.Conv1d(start_dim, input_dim, 1),360
)
创建最终输出卷积 final_conv = Sequential(Conv1dBlock(256→256), Conv1d(256→input_dim)):先用 Conv1dBlock 做滤波(GroupNorm+Mish),再用 1x1 卷积投射回原始动作维度(20)。输出形状 (B,20,T),对应完整 20 维动作空间。
空行。
362
self.diffusion_step_encoder = diffusion_step_encoder363
self.up_modules = up_modules364
self.down_modules = down_modules365
self.final_conv = final_conv将所有子模块挂载为类属性:self.diffusion_step_encoder(时间编码)、self.up_modules(上采样块)、self.down_modules(下采样块)、self.final_conv(最终卷积),使其纳入模型参数跟踪和设备迁移管理,forward() 方法会依次调用这些模块。
367
# print("number of parameters: {:e}".format(368
# sum(p.numel() for p in self.parameters()))369
# )被注释掉的参数计数打印语句。若需统计模型规模可解注释调用,用 sum(p.numel() for p in self.parameters()) 计算总参数数,一般 U-Net 参数量在百万到千万级别。开发调试时常用。
L371–503ConditionalUnet1D 的扩散去噪网络前向通路(含时间步编码、FiLM 条件调制、U-Net 下上采样)及 DP_Action_head 训练/推理接口(DDPM 噪声预测和 action 采样)
371
def forward(
372
self,373
sample: torch.Tensor,
374
timestep: Union[torch.Tensor, float, int],
375
global_cond=None,376
):
ConditionalUnet1D.forward() 函数签名:sample (B,T,input_dim) 为噪声 action,timestep (B,) 为扩散步数,global_cond (B,global_cond_dim) 为条件向量。被 DP_Action_head.noise_pred_net 调用,或在推理时被 generate_flow_action 逐步调用。
377
"""378
x: (B,T,input_dim)379
timestep: (B,) or int, diffusion step380
global_cond: (B,global_cond_dim)381
output: (B,T,input_dim)382
"""forward 的 docstring:输入输出形状说明,强调 timestep 可以是 tensor/float/int,global_cond 是观测特征(通常为图文编码)。
输入坐标变换及空行。sample 从 (B,T,input_dim) 改到 (B,C,T) = (B,input_dim,T),moveaxis(-1,-2) 对调最后两维为后续 Conv1d 做准备。
387
# 1. time388
timesteps = timestep
389
if not torch.is_tensor(timesteps):
390
# TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can391
timesteps = torch.tensor(
392
[timesteps], dtype=torch.long, device=sample.device
393
)
394
elif torch.is_tensor(timesteps) and len(timesteps.shape) == 0:
395
timesteps = timesteps[None].to(sample.device)396
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML397
timesteps = timesteps.expand(sample.shape[0])时间步 t 的规范化处理。若 timestep 是 Python int/float,转为 long tensor;若是 0 维标量 tensor,升维至 1D;最后 expand 广播到 batch_size (B,),避免 CPU-GPU 同步延迟。注释强调应该尽量传 tensor 以加速。
时间步编码:SinusoidalPosEmb(timesteps) 将 t∈[0,num_train_timesteps) 映射为正弦位置向量 (B,256),后续作为 FiLM 调制的条件基础。
401
if global_cond is not None:
402
global_feature = torch.cat([global_feature, global_cond], axis=-1)可选的全局条件融合:若传入 global_cond(VLM 特征),沿最后维 cat 拼接,得到完整条件向量 (B,cond_dim=256+global_cond_dim),用于所有 ConditionalResidualBlock1D 的 FiLM 调制。
403
x = sample
404
h = []
405
for idx, (resnet, resnet2, downsample) in enumerate(self.down_modules):
406
x = resnet(x, global_feature)
407
x = resnet2(x, global_feature)
408
h.append(x)
409
x = downsample(x) # bs, 2048, 5U-Net encoder(下采样路径)初始化及循环。x 赋为 sample,h 为跳跃连接列表。遍历 down_modules,每层两个 ConditionalResidualBlock1D 后 Downsample1d,中间激活保存到 h。流程:(B,14,T)→(B,256,T/2)→(B,512,T/4)→(B,1024,T/8),其中 T 是 action_horizon。
U-Net 中间瓶颈层(bottleneck)。两个 ConditionalResidualBlock1D 作用于最深特征 (B,1024,T/8),无采样,纯 FiLM 调制和非线性变换。
414
for idx, (resnet, resnet2, upsample) in enumerate(self.up_modules):
415
x = torch.cat((x, h.pop()), dim=1)416
x = resnet(x, global_feature)
417
x = resnet2(x, global_feature)
418
x = upsample(x) # bs, 512, 20U-Net decoder(上采样路径)。遍历 up_modules,与 h.pop() 的跳跃连接沿通道维 cat 拼接,两个 ConditionalResidualBlock1D + Upsample1d。最后一层用 Identity 代替 upsample 避免额外上采样。流程:(B,2048,T/8)→cat→(B,1024,T/4)→(B,512,T/2)→(B,256,T)。
最终卷积投影及空行。self.final_conv 含两层 Conv1d,(B,256,T)→(B,256,T)→(B,input_dim,T)。此处 input_dim=14(LIBERO 7 dof action)。
输出坐标变换及返回。x 从 (B,C,T) 转回 (B,T,C) = (B,T,14),返回预测的去噪 action。U-Net 内部全用 (B,C,L) 格式,前后端与 transformer 的 (B,L,C) 格式适配。
428
class DP_Action_head(nn.Module):
429
def __init__(
430
self,431
action_dim=14,432
transformer_dim=896,433
global_cond_dim=1806,434
load_pretrained=True,435
):
DP_Action_head 类定义及 __init__ 签名。Diffusion Policy 风格的动作生成头,参数包括 action_dim(目标维度)、transformer_dim(VLM 条件维度 896)、global_cond_dim(预训练条件维度 1806)、load_pretrained(是否加载预训练权重)。
436
super().__init__()
437
self.action_dim = action_dim438
self.transformer_dim = transformer_dim439
self.load_pretrained = load_pretrained初始化基类及存储属性:super().__init__(),保存 action_dim、transformer_dim、load_pretrained 标志,这些在 forward/predict 中用于控制执行流程。
440
self.noise_scheduler = DDPMScheduler(441
num_train_timesteps=132,442
beta_schedule="squaredcos_cap_v2",
443
clip_sample=True,444
prediction_type="epsilon",
445
)
初始化 DDPM 噪声调度器。num_train_timesteps=132(比离线预训练的少,加速微调)、beta_schedule='squaredcos_cap_v2'(方差调度策略)、prediction_type='epsilon'(预测噪声而非 x0)。
446
if self.load_pretrained:
447
self.global_cond_dim = global_cond_dim448
self.condition_proj = nn.Sequential(449
nn.Linear(self.transformer_dim, 2 * self.transformer_dim),
450
nn.ReLU(),
451
nn.Linear(2 * self.transformer_dim, 2 * self.global_cond_dim),
452
nn.ReLU(),
453
nn.Linear(2 * self.global_cond_dim, self.global_cond_dim),
454
)
456
self.noise_pred_net = ConditionalUnet1D(457
input_dim=self.action_dim,458
# down_dims=[256,512,1024],459
down_dims=[512, 1024, 2048],
460
global_cond_dim=self.global_cond_dim,461
)
预训练模式(load_pretrained=True)初始化。condition_proj 是 3 层 MLP 把 VLM 特征 (B,896)→(B,1806);ConditionalUnet1D 用更大的 down_dims=[512,1024,2048](离线预训练使用的大模型)。
463
# load pretrained model464
action_pretrained_path = "/x2robot/liangyuxin/workspace/DiffusionPolicy/big_mix_0718_mn/30_noise_pred_net.pth"
465
print("load noise_pred_net from:", action_pretrained_path, flush=True)
466
self.noise_pred_net.load_state_dict(torch.load(action_pretrained_path))加载预训练权重:从离线 Diffusion Policy 模型 'big_mix_0718_mn' 的 checkpoint 30_noise_pred_net.pth 读取,仅在 LIBERO 微调时使用。该模型对多任务/多机器人集合预训练过。
467
else:468
self.noise_pred_net = ConditionalUnet1D(469
input_dim=self.action_dim,470
down_dims=[256, 512, 1024],
471
global_cond_dim=self.transformer_dim,472
)
非预训练模式(load_pretrained=False)初始化。ConditionalUnet1D 用更小的 down_dims=[256,512,1024](从零训练),condition_dim 直接为 transformer_dim(896) 无需 condition_proj。
474
def forward(self, naction, condition, sample_times):
475
bs = naction.shape[0]476
noise_shape = (
477
naction.shape[0] * sample_times,478
naction.shape[1],479
naction.shape[2],480
)
forward(naction, condition, sample_times) 方法:训练时的前向接口。naction (B,T,14) 为已归一化的 action(经过 action_norm.normalize),condition (B,896) 为图文特征,sample_times 为 augmentation 倍数。生成 (B*sample_times,T,14) 形状的噪声。
481
noise = torch.randn(noise_shape, device=naction.device)
482
naction = (
483
naction.unsqueeze(1)484
.repeat(1, sample_times, 1, 1)
485
.reshape(bs * sample_times, naction.shape[1], naction.shape[2])
486
)
naction 数据增强。unsqueeze(1) 加 batch 维→repeat(1,sample_times,1,1) 复制→reshape 到 (B*sample_times,T,14),与采样的噪声对齐,增加训练数据多样性(实战:sample_times=2,生成 2 倍数据)。
488
timesteps = torch.randint(
489
0,490
self.noise_scheduler.config.num_train_timesteps,491
(bs * sample_times,),
492
device=naction.device,
493
).long()
随机采样扩散时间步。randint(0,num_train_timesteps) 产生 (B*sample_times,) 的整数时间步,范围 [0,132),用于后续 DDPM 加噪。
494
condition = condition.to(self.condition_proj[0].weight.data.dtype)
495
if self.load_pretrained:
496
condition = self.condition_proj(condition)condition 类型和维度处理。转为浮点(匹配 Linear 层权重 dtype);若预训练则过 condition_proj 投到 global_cond_dim(1806),否则保持 transformer_dim(896)。
498
noisy_actions = self.noise_scheduler.add_noise(naction, noise, timesteps)499
noise_pred = self.noise_pred_net(500
noisy_actions, timesteps, global_cond=condition
501
)
502
return noise, noise_predDDPM 前向过程及噪声预测。add_noise() 执行 noisy_action=(1-sqrt(alpha_bar_t))*noise + sqrt(alpha_bar_t)*action(前向过程);noise_pred_net() 在加噪状态下预测噪声;返回真实噪声和预测噪声用于 MSE 损失。实战:此 forward 被 train_step 每个 batch 调用,MSE 系数乘以 dof_mask(LIBERO 只有前 7 dof 非零)。
forward 方法结束,空行。
L504–632DP_Action_head.predict()推理方法:使用DDPM调度器逐步去噪完成动作生成;ActionProcessor.__init__()与set_normalizer()、sample_time()方法:初始化Flow Matching相关组件(Beta分布、时间嵌入、w1/w2/w3/action_proj_back投影层),设置normalizer对象,实现Beta分布时间步采样
@torch.no_grad()禁用梯度计算以加速推理;def predict(self, condition, naction=None)是DP_Action_head的推理方法入口,condition是图文编码向量(B,transformer_dim),naction是可选初始动作或None
506
bs = condition.shape[0]507
condition = condition.to(self.condition_proj[0].weight.data.dtype)
508
if self.load_pretrained:
509
condition = self.condition_proj(condition)提取batch_size:bs=condition.shape[0];condition转到condition_proj的weight.dtype确保类型兼容;若load_pretrained=True则通过MLP投影condition:(B,896)→condition_proj→(B,1806)
511
if naction is not None:
512
noise_shape = (naction.shape[0], naction.shape[1], naction.shape[2])
513
else:514
noise_shape = (bs, 16, self.action_dim) # tobe parameterized
根据输入确定noise_shape:若naction非None则noise_shape=(B,action_horizon,action_dim),否则默认(bs,16,action_dim);这里16是写死的action_horizon参数,实战中LIBERO用16
515
noise = torch.randn(noise_shape, device=condition.device)
516
naction_pred = noise
517
# init scheduler518
self.noise_scheduler.set_timesteps(519
self.noise_scheduler.config.num_train_timesteps520
)
torch.randn(noise_shape)采样高斯噪声初始化;naction_pred=noise作为完全去噪的起点;scheduler.set_timesteps()初始化DDPM的时间步序列,用于后续迭代去噪
522
for k in self.noise_scheduler.timesteps:
523
# predict noise524
noise_pred = self.noise_pred_net(525
sample=naction_pred, timestep=k, global_cond=condition
526
)
for k in self.noise_scheduler.timesteps遍历噪声时间步(从高→低);每步调noise_pred_net(sample=naction_pred, timestep=k, global_cond=condition)预测当前状态下的噪声,返回shape(B,16,action_dim)
528
# inverse diffusion step (remove noise)529
naction_pred = self.noise_scheduler.step(530
model_output=noise_pred, timestep=k, sample=naction_pred
531
).prev_sample
scheduler.step(model_output=noise_pred, timestep=k, sample=naction_pred)执行一步逆扩散,更新naction_pred到更清晰的状态;.prev_sample是DDPM逆向公式的结果,迭代下去逐步去噪
return naction, naction_pred返回元组:naction来自输入(可能为None),naction_pred是经完整去噪循环后的最终动作预测(B,16,action_dim);接下来是ActionProcessor类定义(空行534-535)
536
class ActionProcessor(nn.Module):
537
def __init__(self, config):
538
super().__init__()
539
self.config = config540
self.dof_config = config.dof_configclass ActionProcessor(nn.Module):融合VLM隐状态和本体状态、生成flow matching损失的核心动作头模块;__init__接收config对象,初始化dof_config、agent_pos_config等关键参数
541
self.agent_pos_config = config.agent_pos_config542
self.action_dim = sum([v for k, v in self.dof_config.items()])
543
self.propri_dim = sum([v for k, v in self.agent_pos_config.items()])
self.action_dim=sum(dof_config.values())得总dof维数(20=双臂14+头2+height1+车3);self.propri_dim=sum(agent_pos_config.values())得本体状态维数;LIBERO实战只用action_dim的前7维,后13维dof_mask=0
545
print_rank_last(
546
f"self.dof_config: {self.dof_config}; action_dim: {self.action_dim}; self.agent_pos_config: {self.agent_pos_config}; propri_dim: {self.propri_dim}"
547
)
549
self.action_hidden_size = config.action_hidden_size550
self.state_hidden_size = config.state_hidden_size551
self.hidden_size = config.hidden_sizeprint_rank_last()在分布式环境仅在最后rank打印一次,用于诊断dof配置;action_hidden_size、state_hidden_size、hidden_size从config读取,分别控制动作编码维数(2048)、本体状态编码维数、VLM隐层维数
553
if not self.config.use_state_string_representation:
554
if self.config.proj_with_mask:
555
self.propri_proj = nn.Linear(556
self.propri_dim * 2, self.state_hidden_size, bias=False
557
)
558
else:559
self.propri_proj = nn.Linear(560
self.propri_dim, self.state_hidden_size, bias=False
561
)
若use_state_string_representation=False(默认),创建propri_proj层:Linear(propri_dim*2→state_hidden_size)若proj_with_mask=True,或Linear(propri_dim→state_hidden_size)若False;实战中LIBERO使用proj_with_mask将本体状态与dof_mask拼接后投影
563
# noise scheduler configing564
if getattr(self.config, "use_flow_action_expert", True):
565
noise_scheduler_config = config.noise_scheduler
566
self.beta_alpha = noise_scheduler_config.get("beta_alpha", 1.5)
567
self.beta_beta = noise_scheduler_config.get("beta_beta", 1.0)
568
self.s = noise_scheduler_config.get("s", 0.999)
569
alpha_tensor = torch.tensor(self.beta_alpha, dtype=torch.float32).to("cuda")
570
beta_tensor = torch.tensor(self.beta_beta, dtype=torch.float32).to("cuda")
571
self.beta_dist = Beta(alpha_tensor, beta_tensor)572
self.time_embed = SinusoidalPosEmb(self.action_hidden_size)
若use_flow_action_expert=True(通常为真),初始化Flow Matching所需的参数:Beta(alpha_tensor=1.5, beta_tensor=1.0)分布用于采样非均匀时间步;s=0.999用于将采样值映射到高噪端;SinusoidalPosEmb(action_hidden_size)对时间步做正弦位置编码
574
# project to hidden space575
if self.config.proj_with_mask:
576
self.w1 = nn.Linear(577
self.action_dim * 2, self.action_hidden_size, bias=False
578
)
579
else:580
self.w1 = nn.Linear(581
self.action_dim, self.action_hidden_size, bias=False
582
)
创建w1层(动作投影到隐空间):Linear(action_dim*2→action_hidden_size)若proj_with_mask=True,或Linear(action_dim→action_hidden_size)若False;实战中proj_with_mask=True时输入为[noisy_action(20)|dof_mask(20)]共40维→2048维
583
if not self.config.use_adarms:
584
self.w2 = nn.Linear(585
self.action_hidden_size * 2, self.action_hidden_size, bias=False
586
)
587
self.w3 = nn.Linear(588
self.action_hidden_size, self.action_hidden_size, bias=False
589
)
590
self.act_fn = nn.SiLU()591
else:592
self.time_mlp_in = nn.Linear(593
self.action_hidden_size, self.action_hidden_size
594
)
595
self.time_mlp_out = nn.Linear(596
self.action_hidden_size, self.action_hidden_size
597
)
598
self.act_fn = nn.SiLU()根据use_adarms选择融合时间特征的方式:若False则创建w2/w3和SiLU激活(w2:2048*2→2048, w3:2048→2048)拼接时间嵌入和动作特征;若True则创建time_mlp_in/out做AdaRM(Adaptive Residual Modulation)条件融合;SiLU=x*sigmoid(x)比ReLU平滑
600
# project back to action space601
self.action_proj_back = nn.Linear(602
self.action_hidden_size, self.action_dim, bias=False
603
)
604
self.mse_loss = nn.MSELoss(reduction="none")
action_proj_back=Linear(action_hidden_size→action_dim)把2048维隐层投回20维速度场;MSELoss(reduction='none')计算逐元素损失,后续与dof_mask相乘确保只对有效dof(LIBERO前7维)计算监督;空行605处分隔__init__和set_normalizer方法
606
def set_normalizer(self, normalizer_action, normalizer_propri):
607
self.normalizer_action = normalizer_action608
self.normalizer_propri = normalizer_propriset_normalizer(normalizer_action, normalizer_propri)是setter方法,绑定两个Normalizer实例;这两个Normalizer对象按dataset_name查表存储min/delta参数,用于训练时的min-max归一化
610
# dataset_name = self.config["data"]["lerobot_config"]["repo_id"]611
# print("normalizer_propri min", self.normalizer_propri.min.__getattr__(dataset_name), flush=True)612
# print("normalizer_propri delta", self.normalizer_propri.delta.__getattr__(dataset_name), flush=True)613
# print("normalizer_action min", self.normalizer_action.min.__getattr__(dataset_name), flush=True)614
# print("normalizer_action delta", self.normalizer_action.delta.__getattr__(dataset_name), flush=True)注释掉的调试代码(5行):原本用于打印normalizer_propri/normalizer_action的min/delta统计;若需验证LIBERO等dataset的动作/本体状态值域范围,可解注此部分来审视dataset特定的归一化参数
616
def sample_time(self, batch_size, device, dtype):
617
"""618
Sampling Time Step619
Generates random numbers in the range [0, 1] using a Beta distribution, and then scales them.621
Parameters:622
batch_size (int): Batch size623
device: Device type624
dtype: Data type626
Returns:627
torch.Tensor: Sampled time steps, with shape [batch_size]628
"""sample_time(batch_size, device, dtype)方法实现Flow Matching所需的非均匀时间步采样;docstring说明:从Beta(1.5,1.0)分布采样[0,1]范围的值,经(1-sample)*0.999压缩到[0,0.999]以偏向高噪端;返回shape [batch_size]的浮点时间向量
629
sample = self.beta_dist.sample([batch_size]).to(device=device, dtype=dtype)630
time = (1 - sample) * self.s
631
return timesample=self.beta_dist.sample([batch_size]).to(device, dtype)产生Beta采样;time=(1-sample)*self.s压缩到[0,0.999];return time返回[batch_size]的时间向量,用于FM训练中的noisy_action合成
空行(文件或下一个函数的起点)
L633–782proprioception_proj、forward、step 三个核心函数:proprioception_proj 处理本体感受投影并 padding 到隐层维度;forward 是训练时加噪-融合流程(加高斯噪声、时间编码、w1 投影后用 w2/w3 或 time_mlp 融合);step 是推理去噪的单步处理(与 forward 逻辑相似,但面向预生成的 noisy_action)。
633
def proprioception_proj(
634
self, proprioception, dataset_names=None, dof_mask=None, use_history=False
635
):
636
"""637
proprioception: [batch_size, 1, action_dim]638
dataset_names: [batch_size]639
dof_mask: [batch_size, action_dim]640
"""proprioception_proj 函数定义+docstring。输入本体感受 [batch_size, 1, action_dim]、数据集名称和 dof_mask,输出投影后的嵌入 [batch_size, 1, hidden_size]。被 VLA 模型推理时调用来编码机械臂位置/速度信息。
641
proprioception = proprioception.to(device=self.propri_proj.weight.device).to(642
dtype=self.propri_proj.weight.dtype643
)
644
if dof_mask is not None:
645
if self.config.proj_with_mask:
646
proprioception = torch.cat(
647
[proprioception, dof_mask], dim=-1648
) # .unsqueeze(1)本体感受设备转换+dof_mask 条件化拼接。641-643 转换设备和 dtype;644-648 若配置 proj_with_mask=true 且有 dof_mask,拼接在最后维。(B,1,D)→(B,1,2*D)。LIBERO 中 dof_mask 标记哪些 dof 有效,拼接让网络感知任务约束。
649
proprioception = proprioception.to(device=self.propri_proj.weight.device).to(650
dtype=self.propri_proj.weight.dtype651
)
652
proprio_embed = self.propri_proj(653
proprioception
654
) # [batch_size, 1, state_hidden_size]再次设备转换+propri_proj 投影。649-651 处理 cat 后可能的设备不对齐(冗余设计);652-654 通过 propri_proj 线性层投影到 state_hidden_size (B,1,2*D)或(B,1,D)→(B,1,state_hidden_size),通常 128-256。
656
if self.state_hidden_size < self.hidden_size:
657
# padding to hidden size658
padding_size = self.hidden_size - self.state_hidden_size
659
padding = torch.zeros(
660
(proprio_embed.shape[0], 1, padding_size),
661
device=proprio_embed.device,
662
dtype=proprio_embed.dtype,
663
)
664
proprio_embed = torch.cat([proprio_embed, padding], dim=-1)666
return proprio_embed # [batch_size, 1, hidden_size]
padding to hidden_size。若 state_hidden_size < hidden_size(如隐层 2048),用零向量补齐到 hidden_size,使本体嵌入与图文 token 维度对齐后可拼接融合。(B,1,state_hidden_size)→(B,1,hidden_size)。
空行分隔符。
668
def forward(self, action_chunk, dataset_names, dof_mask=None):
669
"""670
Parameters:671
action_chunk (torch.Tensor): Action sequence, shape [batch_size, action_chunk_len, action_dim]672
dataset_names: [batch_size]673
dof_mask: [batch_size, action_dim]675
Returns:676
torch.Tensor: Processed action representation, shape [batch_size, seq_len, hidden_size]677
"""forward 函数定义+docstring。训练时动作头前向:输入 action_chunk [B, action_chunk_len, D]、数据集名称和 dof_mask;输出动作嵌入、flow(速度目标)和可选的 adarms 条件。调用方:VLA forward():1427。
678
with torch.autocast("cuda", dtype=torch.float32):
679
action_chunk = action_chunk.to(dtype=torch.float32)
680
batch_size = action_chunk.shape[0]681
device = action_chunk.device
682
dtype = action_chunk.dtype
autocast 上下文+基本参数获取。678-679 强制 float32 精度避免 bfloat16 噪声采样精度丢失;680-682 获取 batch_size、device、dtype 供后续张量操作。
684
# 1. add noise to action_chunk685
noise = torch.randn_like(action_chunk)
686
time = self.sample_time(batch_size, device, dtype)687
time_expanded = time.unsqueeze(-1).unsqueeze(-1)
688
noisy_action = (1 - time_expanded) * noise + time_expanded * action_chunk689
flow = action_chunk - noise
Flow matching 加噪核心。空行+采样高斯噪声、时间 t~Beta(1.5,1)、计算 noisy_action=(1-t)*noise+t*action、flow=action-noise。t 近 1 时接近原动作,t 近 0 时接近噪声。(B,A,D) 运算,A=action_horizon。
691
# 2. sinusoidal positional encoding for timesteps692
time_embed = self.time_embed(time).to(torch.float32)694
self.noise = noise695
self.noisy_action = noisy_action # for new x-pred
时间位置编码+中间变量保存。空行+sin/cos 位置编码把标量 t 映射到隐维 (B,hidden_size);保存 noise/noisy_action 供 loss 计算,time_expanded (B,1,1) 供梯度流。
697
# 3.action_chunk_nosiy + t_pos_emb -> MLP_act_chunk -> action_chunk_nosiy_emb_with_t (dim=trans * chunk)698
if dof_mask is not None:
699
noisy_action = torch.cat([noisy_action, dof_mask], dim=-1)701
noisy_action = noisy_action.to(dtype=self.w1.weight.dtype)702
action_embed = self.w1(noisy_action)704
self.time_expanded = time_expanded # for new x-pred
动作噪声条件化投影+时间保存。698-699 若有 dof_mask 拼接到 noisy_action;701-702 通过 w1 线性层投影 (B,A,2*D)或(B,A,D)→(B,A,action_hidden_size);704 保存 time_expanded 避免重复 unsqueeze。LIBERO 中 dof_mask 过滤无用自由度。
706
if not self.config.use_adarms:
707
time_embed = (
708
time_embed.unsqueeze(1)709
.repeat(1, action_embed.shape[1], 1)
710
.to(dtype=self.w2.weight.dtype)711
)
712
concat_embed = torch.cat([action_embed, time_embed], dim=-1)713
concat_embed = self.w2(concat_embed)714
action_time_embed = self.w3(self.act_fn(concat_embed))
715
adarms_cond = Noneuse_adarms=false 时的时间-动作融合分支。time_embed repeat 到每个 horizon (B,A,hidden_size)、与 action_embed cat→w2(线性)→act_fn(SiLU);输出 action_time_embed (B,A,action_hidden_size);adarms_cond=None。
716
else:717
time_embed = self.time_mlp_in(time_embed)718
time_embed = self.act_fn(time_embed)719
time_embed = self.time_mlp_out(time_embed)720
time_embed = self.act_fn(time_embed)721
action_time_embed = action_embed
722
adarms_cond = time_embed
use_adarms=true 时的融合分支。time_embed 过 MLP (time_mlp_in→SiLU→time_mlp_out→SiLU) 作为 AdaLN 条件参数;action_embed 保持原样;adarms_cond (B,hidden_size) 供后续 AdaLN 条件化。两种分支目的都是让网络感知时间步。
724
if self.action_hidden_size < self.hidden_size:
725
# padding to hidden size726
padding_size = self.hidden_size - self.action_hidden_size
727
padding = torch.zeros(
728
(
729
action_time_embed.shape[0],730
action_time_embed.shape[1],padding to hidden_size(前半)。空行+if 判断 action_hidden_size < hidden_size;计算 padding_size;初始化零向量 (B,A,padding_size)。
731
padding_size,
732
),
733
device=action_time_embed.device,
734
dtype=action_time_embed.dtype,
735
)
736
action_time_embed = torch.cat([action_time_embed, padding], dim=-1)padding to hidden_size(后半)。零向量设备/dtype 转换;与 action_time_embed cat 到最后维。(B,A,action_hidden_size)→(B,A,hidden_size),使输出与 VLA transformer token 维度对齐。
forward 返回。空行+返回三个值:action_time_embed (B,A,hidden_size) 供 VLA 编码;flow (B,A,D) 是 MSE 监督目标;adarms_cond 供 AdaLN。空行。
740
def step(self, timestep, noisy_action, dof_mask=None):
741
# noisy_action: bs, pred_horizon, action_dim742
# timestep: bsstep 函数定义+简注。推理时单步去噪处理,从当前 noisy_action 和时间步推出去噪后的嵌入。被 generate_flow_action() 的 Euler 积分循环 5 步调用(部署时)。
743
with torch.autocast("cuda", dtype=torch.float32):
744
if dof_mask is not None and self.config.proj_with_mask:
745
if dof_mask.shape[1] == 1:
746
dof_mask = dof_mask.unsqueeze(1).repeat(1, noisy_action.shape[1], 1)
747
noisy_action = torch.cat([noisy_action, dof_mask], dim=-1)推理时 autocast+dof_mask 处理。743 float32 autocast;744-747 若有 dof_mask 且 proj_with_mask,检查 shape[1] 是否为 1 并 repeat 到 horizon 维度。(B,1,D)→(B,A,D),然后拼接到 noisy_action。LIBERO 部署中确保生成动作符合任务约束。
749
noisy_action = noisy_action.to(dtype=self.w1.weight.dtype)750
time_embed = self.time_embed(timestep).to(torch.float32) # bs,hidden_size
751
action_embed = self.w1(noisy_action)w1 投影+时间编码。空行+749 dtype 转换;750 时间编码 (B,hidden_size);751 w1 投影 (B,A,action_hidden_size)。与 forward 不同的是 timestep 来自 Euler 迭代计数而非 Beta 采样。
753
if not self.config.use_adarms:
754
time_embed = time_embed.unsqueeze(1).repeat(1, action_embed.shape[1], 1)
755
time_embed = time_embed.to(device=noisy_action.device).to(
756
dtype=noisy_action.dtype
757
)
758
concat_embed = torch.cat([action_embed, time_embed], dim=-1)759
concat_embed = self.w2(concat_embed)760
embed = self.w3(self.act_fn(concat_embed)) # is this right?
761
adarms_cond = None推理时 use_adarms=false 分支。空行+time_embed repeat (B,A,hidden_size);设备/dtype 转换;cat→w2→act_fn→w3 融合。输出 embed (B,A,action_hidden_size);adarms_cond=None。
762
else:763
time_embed = time_embed.to(dtype=self.time_mlp_in.weight.dtype)764
time_embed = self.time_mlp_in(time_embed)765
time_embed = self.act_fn(time_embed)766
time_embed = self.time_mlp_out(time_embed)767
time_embed = self.act_fn(time_embed)768
embed = action_embed
769
adarms_cond = time_embed
推理时 use_adarms=true 分支。time_embed 过 MLP (time_mlp_in→act_fn→time_mlp_out→act_fn) 作为 adarms_cond;action_embed 直接作为 embed。与 forward 对称,目标仍是让输出与时间步一致的去噪结果。
771
if self.action_hidden_size < self.hidden_size:
772
# padding to hidden size773
padding_size = self.hidden_size - self.action_hidden_size
774
padding = torch.zeros(
775
(embed.shape[0], embed.shape[1], padding_size),
776
device=embed.device,
777
dtype=embed.dtype,
778
)
779
embed = torch.cat([embed, padding], dim=-1)推理时 padding to hidden_size。空行+if 判断+计算 padding_size;初始化零向量 (B,A,padding_size);设备/dtype 转换;与 embed cat。(B,A,action_hidden_size)→(B,A,hidden_size),对齐 VLA transformer 维度。
step 返回。空行+返回 (embed, adarms_cond):embed (B,A,hidden_size) 是去噪后的动作表示供 VLA 解码;adarms_cond (B,hidden_size) 或 None 供 AdaLN 条件化。
L783–808ActionProcessor.flow_loss() 方法:计算流匹配损失。从 VLM 隐状态投影回动作空间得到速度预测,与流目标(action-noise)计算 MSE 损失,然后用 dof_mask(维度有效性掩码)和 flow_loss_mask(时间步掩码)对损失进行选择性应用。
783
def flow_loss(
784
self,785
action_hidden_states,
786
flow,
787
action_chunk,
788
dof_mask=None,789
flow_loss_mask=None,790
):
flow_loss() 方法定义,被 vla_mixin.py:forward() 调用以计算流匹配目标的损失。接收 action_hidden_states(VLM提取的动作token隐状态,形状B*S×H)、flow(action-noise流目标,B*S×D)、action_chunk(原始动作块)、dof_mask(自由度有效掩码)、flow_loss_mask(序列级时间步掩码)。
791
with torch.autocast("cuda", dtype=torch.float32):
进入 torch.autocast("cuda", dtype=torch.float32) 上下文,强制后续计算为 float32 精度保证数值稳定性(特别是损失计算)。
action_proj_back()(Linear 权重维度为 action_hidden_size→action_dim)将隐状态投影回动作空间。输入形状 (B*S, H) 其中 action_hidden_states[:, :self.action_hidden_size] 提取前 2048 维,输出 action_pred 形状 (B*S, 20)。这是速度场预测(flow prediction),对标扩散中的噪声预测。
795
v_pred = action_pred
v_pred = action_pred,别名赋值便于理解这是速度预测(velocity prediction)而非噪声预测。
796
loss = self.mse_loss(v_pred, flow)self.mse_loss(reduction='none') 计算速度预测 v_pred 与流目标 flow 的逐元素平方误差,输出形状 (B*S, 20),每个样本每个维度一个独立的损失值待后续掩码。
797
if dof_mask is not None:
798
dof_mask = dof_mask.reshape(-1, dof_mask.shape[-1])
799
loss = loss * dof_mask
如果 dof_mask 存在(形状 B×1×20 或 B×chunk_len×20),reshape 为 (B*S, 20) 平面形式,然后逐元素乘到损失,使对应维度为 0 的自由度损失清零。LIBERO实战:只有前 7 个维度有效,后 13 个维度 dof_mask=0,训练时这些维度损失自动忽略。
如果 flow_loss_mask 存在(序列级掩码,形状 B*S),表示某些时间步的损失应被忽略。unsqueeze(-1) 从 (B*S,) 变成 (B*S, 1),为后续的链式操作做准备。
reshape(-1, 1) 保持形状,.expand(-1, loss.shape[-1]) 广播到 (B*S, 20) 与损失同形。最后 loss *= flow_loss_mask 逐元素乘法,使掩码为 0 的时间步的所有维度损失都变为 0。
808
return loss返回处理后的损失张量,形状 (B*S, 20)。调用方(vla_mixin.py) 会对其 .mean() 取平均值,然后乘以 flow_loss_weight 加入总损失。
源码零改写,与仓库 wall-x/wall_x/model/action_head.py 逐字节一致(commit 97406f2)。
生成于 wall-x LIBERO 微调项目 · ← 返回导读