wall-x 逐行精读 导读modeling_qwen2_5_vl_act.pyvla_mixin.pyaction_head.pytrain_libero_wrapper.py Notebook结果页

action_head.py — 动作头 808 行

Normalizer 归一化查表、时间步正弦编码、ActionProcessor(flow matching 的加噪/投影/损失) · commit 97406f2 · 每一行都有右栏中文讲解,行号可作锚点直链(如 #L12)。

L1–122Normalizer 类(按 dataset_name 查表归一化动作)、辅助函数(分布式打印、正弦位置编码、卷积采样),以及 normalize_data 方法核心实现。支持 per-robot/per-dataset 统计,LIBERO 中按前7维 dof_mask 截断。
1
import torch
2
 
3
import torch.nn as nn
4
 
5
from typing import Union
6
import math
7
 
8
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) 用)。
10
 
11
 
空行。
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。
19
 
20
 
21
class Normalizer(nn.Module):
空行+Normalizer 类定义开始,继承 nn.Module。该类管理动作空间的归一化统计量(per-robot/per-dataset),支持 normalize/unnormalize 两个方向,是 action_head 的配套模块。
22
    @classmethod
23
    def from_ckpt(cls, ckpt_path):
24
        instance = cls.__new__(cls)
@classmethod 装饰器+from_ckpt 类方法签名。从保存的 checkpoint 文件恢复 Normalizer 实例(离线构建模式),对标标准 pytorch_lightning 的 from_checkpoint 模式。
25
        nn.Module.__init__(instance)
26
 
27
        instance.min = nn.ParameterDict()
28
        instance.delta = nn.ParameterDict()
29
        instance.min_key = "min"
30
        instance.delta_key = "delta"
31
 
创建空实例并手动调用 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")
33
 
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=False
41
                    )
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 并跳过(健壮性)。
46
 
47
        return instance
48
 
返回初始化完成的 Normalizer 实例;空行分隔。
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。
53
 
54
        self.min_key = min_key
55
        self.delta_key = delta_key
存储 min_key 和 delta_key 为实例属性,后续 normalize_data 通过这两个 key 从 ParameterDict 查索,支持自定义键名(如备选命名约定)。
56
 
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
93
 
将累积的列表转换为 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]。
106
 
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 跨多行打印确保可读性。实战:快速检查是否加载了正确的统计值。
111
 
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
122
 
空行+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。
127
        for x, dataset_name, mask in zip(xs, dataset_names, dof_mask):
128
            x = (x + 1) / 2
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
返回反归一化后的动作张量。
140
 
141
 
空行。
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 = dim
146
        if dim % 2 != 0:
147
            raise ValueError(f"embedding_dim ({dim}) must be divisible by 2")
148
        self.min_period = min_period
149
        self.max_period = max_period
__init__ 初始化:dim 必须偶数(sin/cos 各占一半),min_period/max_period 控制频率范围(默认 4e-3~4.0)。校验 dim%2==0,否则抛 ValueError。
150
 
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 emb
forward 方法:输入 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) 的正弦编码。
161
 
空行。
162
 
163
class Downsample1d(nn.Module):
164
    def __init__(self, dim):
165
        super().__init__()
166
        self.conv = nn.Conv1d(dim, dim, 3, 2, 1)
167
 
168
    def forward(self, x):
169
        return self.conv(x)
Downsample1d:1D 卷积下采样(kernel=3, stride=2, padding=1),保持通道数不变。用于 UNet encoder 路径压缩序列长度。
170
 
空行。
171
 
172
class Upsample1d(nn.Module):
173
    def __init__(self, dim):
174
        super().__init__()
175
        self.conv = nn.ConvTranspose1d(dim, dim, 4, 2, 1)
176
 
177
    def forward(self, x):
178
        return self.conv(x)
Upsample1d:转置卷积上采样(kernel=4, stride=2, padding=1),保持通道数不变。用于 UNet decoder 路径恢复序列长度。
179
 
空行。
180
 
181
class Conv1dBlock(nn.Module):
182
    """
183
    Conv1d --> GroupNorm --> Mish
184
    """
Conv1dBlock 类及 docstring:简单的 Conv1d→GroupNorm→Mish 堆栈,无跨层连接。是残差块的基础单元。
185
 
186
    def __init__(self, inp_channels, out_channels, kernel_size, n_groups=8):
187
        super().__init__()
188
 
189
        self.block = nn.Sequential(
190
            nn.Conv1d(
191
                inp_channels, out_channels, kernel_size, padding=kernel_size // 2
192
            ),
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)。
196
 
197
    def forward(self, x):
198
        return self.block(x)
forward 方法:直接调用 self.block (Sequential)。
199
 
200
 
空行。
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__()
204
 
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))。
211
 
212
        # FiLM modulation https://arxiv.org/abs/1709.07871
213
        # predicts per-channel scale and bias
214
        cond_channels = out_channels * 2
215
        self.out_channels = out_channels
216
        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。
222
 
223
        # make sure dimensions compatible
224
        self.residual_conv = (
225
            nn.Conv1d(in_channels, out_channels, 1)
226
            if in_channels != out_channels
227
            else nn.Identity()
228
        )
残差连接处理:若 in_channels != out_channels 则添加 Conv1d(in_channels→out_channels, kernel=1) 做维度匹配;否则用 Identity()。
229
 
230
    def forward(self, x, cond):
231
        """
232
        x : [ batch_size x in_channels x horizon ]
233
        cond : [ batch_size x cond_dim]
234
 
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)
240
 
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(逐通道仿射变换)。
245
 
246
        out = self.blocks[1](out)
247
        out = out + self.residual_conv(x)
248
        return out
第二个 Conv1dBlock 处理调制后的 out;加上残差连接(residual_conv(x))返回最终输出。residual_conv 处理通道数不匹配(在 __init__ 已定义好)。
249
 
250
 
空行。
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 FiLM
265
          in addition to diffusion step embedding. This is usually obs_horizon * obs_dim
266
        diffusion_step_embed_dim: Size of positional encoding for diffusion iteration k
267
        down_dims: Channel size for each UNet level.
268
          The length of this array determines numebr of levels.
269
        kernel_size: Conv kernel size
270
        n_groups: Number of groups for GroupNorm
271
        """
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 控制卷积和归一化的细节配置。
272
 
273
        super().__init__()
空行后 super().__init__():初始化 nn.Module 基类,使当前类继承 PyTorch 模块管理机制(包含分隔空行与初始化语句)。
274
        all_dims = [input_dim] + list(down_dims)
275
        start_dim = down_dims[0]
构建维度映射: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 掩蔽不用的维度。
276
 
空行。
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。
285
 
空行。
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),接收拼接后的时间+全局条件。此处不做上下采样,只做特征变换。
306
 
空行。
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。
331
 
空行。
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。
356
 
空行。
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 维动作空间。
361
 
空行。
362
        self.diffusion_step_encoder = diffusion_step_encoder
363
        self.up_modules = up_modules
364
        self.down_modules = down_modules
365
        self.final_conv = final_conv
将所有子模块挂载为类属性:self.diffusion_step_encoder(时间编码)、self.up_modules(上采样块)、self.down_modules(下采样块)、self.final_conv(最终卷积),使其纳入模型参数跟踪和设备迁移管理,forward() 方法会依次调用这些模块。
366
 
367
        # print("number of parameters: {:e}".format(
368
        #     sum(p.numel() for p in self.parameters()))
369
        # )
370
 
被注释掉的参数计数打印语句。若需统计模型规模可解注释调用,用 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 step
380
        global_cond: (B,global_cond_dim)
381
        output: (B,T,input_dim)
382
        """
forward 的 docstring:输入输出形状说明,强调 timestep 可以是 tensor/float/int,global_cond 是观测特征(通常为图文编码)。
383
        # (B,T,C)
384
        sample = sample.moveaxis(-1, -2)
385
        # (B,C,T)
386
 
输入坐标变换及空行。sample 从 (B,T,input_dim) 改到 (B,C,T) = (B,input_dim,T),moveaxis(-1,-2) 对调最后两维为后续 Conv1d 做准备。
387
        # 1. time
388
        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 can
391
            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 ML
397
        timesteps = timesteps.expand(sample.shape[0])
398
 
时间步 t 的规范化处理。若 timestep 是 Python int/float,转为 long tensor;若是 0 维标量 tensor,升维至 1D;最后 expand 广播到 batch_size (B,),避免 CPU-GPU 同步延迟。注释强调应该尽量传 tensor 以加速。
399
        global_feature = self.diffusion_step_encoder(timesteps)
400
 
时间步编码: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, 5
U-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。
410
 
411
        for mid_module in self.mid_modules:
412
            x = mid_module(x, global_feature)  # bs, 2048, 5
U-Net 中间瓶颈层(bottleneck)。两个 ConditionalResidualBlock1D 作用于最深特征 (B,1024,T/8),无采样,纯 FiLM 调制和非线性变换。
413
 
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, 20
U-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)。
419
 
420
        x = self.final_conv(x)  # bs, 14, 20
421
 
最终卷积投影及空行。self.final_conv 含两层 Conv1d,(B,256,T)→(B,256,T)→(B,input_dim,T)。此处 input_dim=14(LIBERO 7 dof action)。
422
        # (B,C,T)
423
        x = x.moveaxis(-1, -2)
424
        # (B,T,C)
425
        return x
输出坐标变换及返回。x 从 (B,C,T) 转回 (B,T,C) = (B,T,14),返回预测的去噪 action。U-Net 内部全用 (B,C,L) 格式,前后端与 transformer 的 (B,L,C) 格式适配。
426
 
427
 
ConditionalUnet1D 类结束,两行空行分隔。
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_dim
438
        self.transformer_dim = transformer_dim
439
        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_dim
448
            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
            )
455
 
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](离线预训练使用的大模型)。
462
 
463
            # load pretrained model
464
            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。
473
 
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 倍数据)。
487
 
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)。
497
 
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_pred
DDPM 前向过程及噪声预测。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 非零)。
503
 
forward 方法结束,空行。
L504–632DP_Action_head.predict()推理方法:使用DDPM调度器逐步去噪完成动作生成;ActionProcessor.__init__()与set_normalizer()、sample_time()方法:初始化Flow Matching相关组件(Beta分布、时间嵌入、w1/w2/w3/action_proj_back投影层),设置normalizer对象,实现Beta分布时间步采样
504
    @torch.no_grad()
505
    def predict(self, condition, naction=None):
@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)
510
 
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 scheduler
518
        self.noise_scheduler.set_timesteps(
519
            self.noise_scheduler.config.num_train_timesteps
520
        )
torch.randn(noise_shape)采样高斯噪声初始化;naction_pred=noise作为完全去噪的起点;scheduler.set_timesteps()初始化DDPM的时间步序列,用于后续迭代去噪
521
 
522
        for k in self.noise_scheduler.timesteps:
523
            # predict noise
524
            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)
527
 
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逆向公式的结果,迭代下去逐步去噪
532
 
533
        return naction, naction_pred
534
 
535
 
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 = config
540
        self.dof_config = config.dof_config
class ActionProcessor(nn.Module):融合VLM隐状态和本体状态、生成flow matching损失的核心动作头模块;__init__接收config对象,初始化dof_config、agent_pos_config等关键参数
541
        self.agent_pos_config = config.agent_pos_config
542
        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
544
 
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
        )
548
 
549
        self.action_hidden_size = config.action_hidden_size
550
        self.state_hidden_size = config.state_hidden_size
551
        self.hidden_size = config.hidden_size
print_rank_last()在分布式环境仅在最后rank打印一次,用于诊断dof配置;action_hidden_size、state_hidden_size、hidden_size从config读取,分别控制动作编码维数(2048)、本体状态编码维数、VLM隐层维数
552
 
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拼接后投影
562
 
563
        # noise scheduler configing
564
        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)对时间步做正弦位置编码
573
 
574
            # project to hidden space
575
            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平滑
599
 
600
            # project back to action space
601
            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")
605
 
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_action
608
        self.normalizer_propri = normalizer_propri
set_normalizer(normalizer_action, normalizer_propri)是setter方法,绑定两个Normalizer实例;这两个Normalizer对象按dataset_name查表存储min/delta参数,用于训练时的min-max归一化
609
 
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)
615
 
注释掉的调试代码(5行):原本用于打印normalizer_propri/normalizer_action的min/delta统计;若需验证LIBERO等dataset的动作/本体状态值域范围,可解注此部分来审视dataset特定的归一化参数
616
    def sample_time(self, batch_size, device, dtype):
617
        """
618
        Sampling Time Step
619
        Generates random numbers in the range [0, 1] using a Beta distribution, and then scales them.
620
 
621
        Parameters:
622
            batch_size (int): Batch size
623
            device: Device type
624
            dtype: Data type
625
 
626
        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 time
sample=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合成
632
 
空行(文件或下一个函数的起点)
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.dtype
643
        )
644
        if dof_mask is not None:
645
            if self.config.proj_with_mask:
646
                proprioception = torch.cat(
647
                    [proprioception, dof_mask], dim=-1
648
                )  # .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.dtype
651
        )
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。
655
 
656
        if self.state_hidden_size < self.hidden_size:
657
            # padding to hidden size
658
            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)
665
 
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)。
667
 
空行分隔符。
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]
674
 
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 供后续张量操作。
683
 
684
            # 1. add noise to action_chunk
685
            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_chunk
689
            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。
690
 
691
            # 2. sinusoidal positional encoding for timesteps
692
            time_embed = self.time_embed(time).to(torch.float32)
693
 
694
            self.noise = noise
695
            self.noisy_action = noisy_action  # for new x-pred
696
 
时间位置编码+中间变量保存。空行+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)
700
 
701
            noisy_action = noisy_action.to(dtype=self.w1.weight.dtype)
702
            action_embed = self.w1(noisy_action)
703
 
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 过滤无用自由度。
705
 
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 = None
use_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 条件化。两种分支目的都是让网络感知时间步。
723
 
724
            if self.action_hidden_size < self.hidden_size:
725
                # padding to hidden size
726
                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 维度对齐。
737
 
738
        return action_time_embed, flow, adarms_cond
739
 
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_dim
742
        # timestep: bs
step 函数定义+简注。推理时单步去噪处理,从当前 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 部署中确保生成动作符合任务约束。
748
 
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 采样。
752
 
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 对称,目标仍是让输出与时间步一致的去噪结果。
770
 
771
            if self.action_hidden_size < self.hidden_size:
772
                # padding to hidden size
773
                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 维度。
780
 
781
        return embed, adarms_cond
782
 
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 精度保证数值稳定性(特别是损失计算)。
792
            action_pred = self.action_proj_back(
793
                action_hidden_states[:, : self.action_hidden_size]
794
            )
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
800
 
如果 dof_mask 存在(形状 B×1×20 或 B×chunk_len×20),reshape 为 (B*S, 20) 平面形式,然后逐元素乘到损失,使对应维度为 0 的自由度损失清零。LIBERO实战:只有前 7 个维度有效,后 13 个维度 dof_mask=0,训练时这些维度损失自动忽略。
801
            if flow_loss_mask is not None:
802
                flow_loss_mask = (
803
                    flow_loss_mask.unsqueeze(-1)
如果 flow_loss_mask 存在(序列级掩码,形状 B*S),表示某些时间步的损失应被忽略。unsqueeze(-1) 从 (B*S,) 变成 (B*S, 1),为后续的链式操作做准备。
804
                    .reshape(-1, 1)
805
                    .expand(-1, loss.shape[-1])
806
                )
807
                loss = loss * flow_loss_mask
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 微调项目 · ← 返回导读