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

train_libero_wrapper.py — 本项目补丁 89 行

docker 挂载的猴子补丁:7 维 LIBERO 动作 pad 进 20 维布局 + PROPRI_DROPOUT 治捷径学习 · commit 97406f2 · 每一行都有右栏中文讲解,行号可作锚点直链(如 #L21)。

L1–89wall-x VLA 训练的 monkeypatch wrapper:拦截 LIBERO 7维动作/本体状态,NaN-pad 到 20 维前 7 槽;注入 PROPRI_DROPOUT 对抗捷径学习;修复上游 LeRobotDataset 的 episode 子集越界 bug;最后代理启动原始训练脚本。
1
#!/usr/bin/env python3
2
"""train_qact.py 的 wrapper:注入 libero 7 维 -> 20 维前 7 槽的 NaN pad。
3
 
4
上游 "Fix normalizer (#57)" 把 collate 的 pad-到-20 代码注释掉了,当前管线对
5
非 20 维数据 broken;与开环评估(openloop_libero.py)使用同一 monkeypatch,
6
该修复已在开环链路上验证。wandb 走 offline 模式,无需账号。
7
"""
模块级 docstring:说明这是对 train_qact.py 的 wrapper,核心功能是注入 7→20 维 NaN-pad。背景:上游 #57 注释掉了 collate 的 pad-to-20 代码,导致对非 20 维数据(LIBERO 只有 7 个 action dof)broken;与开环评估路径用同一 monkeypatch;wandb 走离线模式避免账号依赖。这是本 wrapper 的核心价值——修复该兼容性破洞。
8
import os
9
import sys
10
 
11
os.environ.setdefault("WANDB_MODE", "offline")
12
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
13
 
导入 os/sys 模块,并设置两个环保变量:WANDB_MODE=offline(禁用在线日志)+ TOKENIZERS_PARALLELISM=false(HuggingFace tokenizer 多进程警告)。这两行确保后续 import 的库不会有网络依赖或多进程冲突。空行 13 并入本组。
14
import torch
15
 
16
from wall_x.data.load_lerobot_dataset import PreprocessedDataset
17
 
18
MAX_DOF = 20
导入核心依赖 torch;导入 PreprocessedDataset(lerobot 数据 wrapper,负责 collate/normalization);定义全局常数 MAX_DOF=20,后续所有 pad 操作都基于这个维度。LIBERO 任务只用前 7 维,但模型架构期望 20 维(跨本体 dof_config),所以需要 pad。
19
 
20
 
21
def _pad_nan(t, dim):
22
    out = torch.full((*t.shape[:-1], MAX_DOF), float("nan"), dtype=t.dtype)
23
    out[..., :dim] = t[..., :dim]
24
    return out
_pad_nan() 函数:将任意形状的张量 t 的最后维从 dim 扩到 MAX_DOF(20),剩余位用 NaN 填充。形状变化:(…, dim) → (…, 20)。为什么用 NaN:后续 flow matching 里 dof_mask=0 的位会被 loss 忽略,但需要占位避免 OOM/形状不匹配;NaN 作为 sentinel 值清晰标示这些位无效。
25
 
26
 
27
_orig_getitem = PreprocessedDataset.__getitem__
保存 PreprocessedDataset 的原始 __getitem__() 方法到 _orig_getitem(后续 monkey patch 会调用它)。这是 Python hook 模式的标准做法——在覆盖前保存原始实现。
28
 
29
 
30
PROPRI_DROPOUT = float(os.environ.get("PROPRI_DROPOUT", "0"))
从环境变量读 PROPRI_DROPOUT,默认 0(无 dropout)。这个参数控制是否以概率 p 隐藏本体状态(proprioception),用于对抗捷径学习:官方 issue #58 实锤过微调后模型忽略图像、直接从 proprio 回归动作的失败模式。LIBERO 场景下本体信息充分,所以需要这个对抗机制。
31
 
32
 
33
def _padded_getitem(self, index):
34
    r = _orig_getitem(self, index)
35
    r["action"] = _pad_nan(r["action"], 7)
36
    # state 第 8 维是对称冗余的第二夹爪指,截去与 action 布局对齐
37
    # PROPRI_DROPOUT: 以概率 p 隐藏本体状态(全 NaN -> agent_pos_mask=0),
38
    # 对抗捷径学习——微调后模型忽略图像只回归 proprio(官方 issue #58 实锤的失败模式)
39
    if PROPRI_DROPOUT > 0 and torch.rand(1).item() < PROPRI_DROPOUT:
40
        r["agent_pos"] = torch.full((*r["agent_pos"].shape[:-1], MAX_DOF), float("nan"))
41
    else:
42
        r["agent_pos"] = _pad_nan(r["agent_pos"], 7)
43
    return r
_padded_getitem() 被 monkey patch 成 PreprocessedDataset.__getitem__,覆盖原始的 batch 返回逻辑。核心三步:(1) 调用 _orig_getitem() 获得原始 batch;(2) 把 r['action'] 从 (B,A,7) pad 到 (B,A,20)(A=action_horizon);(3) 把 r['agent_pos'] 从 (B,S,7) pad 到 (B,S,20)(S=seq_len),同时注入 PROPRI_DROPOUT 隐藏本体:若随机数 < p,则整个 agent_pos 置 NaN(对应 agent_pos_mask=0)。实战视角:PROPRI_DROPOUT 是对抗微调陷阱的关键调参,官方实验用 0.5 能显著抑制捷径学习。
44
 
45
 
46
PreprocessedDataset.__getitem__ = _padded_getitem
第一个 monkey patch:用 _padded_getitem 替换 PreprocessedDataset.__getitem__。从此所有 PreprocessedDataset 实例的 batch 都会自动做 7→20 pad + PROPRI_DROPOUT 隐藏。这个 patch 时机很关键——必须在 load_dataset() 之前,否则会 getitem() 后再 patch,来不及。
47
 
48
# 上游 load_dataset 忽略 lerobot_config.episodes,写死用全量 95% 切分(25 万帧/epoch,
49
# 3090 上 12.5h/epoch)。拦截 LeRobotDataset 构造:episodes 超过 1000(即全量 train
50
# split)时替换为 libero_spatial 前 120 集;test split(~85 集)与开环(1 集)不受影响。
51
SPATIAL_EPISODES = list(range(1261, 1693))
定义 SPATIAL_EPISODES = list(range(1261, 1693)):LIBERO 的 libero_spatial 任务集对应的 episode ID 范围(120 个 episode)。注释解释:上游 load_dataset 忽略 lerobot_config.episodes 参数,写死用全量 95% split(约 25 万帧/epoch,3090 上 12.5h/epoch);这里拦截 LeRobotDataset 构造,当 episodes 超过 1000(即全量 train split)时替换为 libero_spatial 前 120 集。test split 和开环不受影响。
52
 
53
import wall_x.data.load_lerobot_dataset as _dsmod
54
 
导入 wall_x.data.load_lerobot_dataset 模块,保存其中的原始 LeRobotDataset class(后续 patch)。两行一起形成第二个 monkeypatch 的前置。
55
_OrigLeRobotDataset = _dsmod.LeRobotDataset
56
 
57
 
58
def _patched_lerobot_dataset(*args, **kwargs):
59
    eps = kwargs.get("episodes")
60
    if eps is not None and len(eps) > 1000:
61
        print(f"[wrapper] override train episodes: {len(eps)} -> {len(SPATIAL_EPISODES)} (libero_spatial)")
62
        kwargs["episodes"] = SPATIAL_EPISODES
63
    return _OrigLeRobotDataset(*args, **kwargs)
_patched_lerobot_dataset() wrapper 函数:拦截 LeRobotDataset 构造。逻辑是若 kwargs['episodes'] 存在且长度 > 1000(说明传了全量),则打印日志并替换为 SPATIAL_EPISODES(120 个)。否则直接转发原始 LeRobotDataset 构造。实战视角:LIBERO spatial 任务集是官方推荐用的数据子集(只含空间推理类任务),比全量 100 任务更稳定。
64
 
65
 
66
_dsmod.LeRobotDataset = _patched_lerobot_dataset
第二个 monkey patch:用 _patched_lerobot_dataset wrapper 替换 _dsmod.LeRobotDataset(module 级别的类引用)。这样 load_lerobot_dataset 里所有 LeRobotDataset() 调用都会走 wrapper 逻辑。
67
 
68
# lerobot 0.3.4 bug:episodes 子集模式下,帧数据的 episode_index 列是全局编号,
69
# 但 _get_query_indices 直接拿它当局部 episode_data_index 下标 → 越界
70
# (官方 load_test_dataset 用 episodes=[0] 只是碰巧 0 号不越界)。补全局→局部映射。
71
from lerobot.datasets.lerobot_dataset import LeRobotDataset as _LRDCls
72
 
73
_orig_gqi = _LRDCls._get_query_indices
修复 lerobot 0.3.4 的 episode 子集越界 bug。背景:当用 episodes 参数截取数据子集时,frame 数据的 episode_index 列保持全局编号(如 1261),但 _get_query_indices() 直接拿它当本地 episode_data_index 下标,导致越界(官方 load_test_dataset 用 episodes=[0] 只是碰巧 0 号不越界)。保存原始 _get_query_indices() 到 _orig_gqi,后续覆盖修复。
74
 
75
 
76
def _fixed_gqi(self, idx, ep_idx):
77
    if self.episodes is not None:
78
        ep_idx = self.episodes.index(int(ep_idx))
79
    return _orig_gqi(self, idx, ep_idx)
_fixed_gqi() 修复函数:重写 LeRobotDataset._get_query_indices(idx, ep_idx)。核心修复是:若 self.episodes 不为 None(即用了子集),则把全局 ep_idx 用 .index() 方法映射到本地位置(0-indexed),再调用原始 _orig_gqi()。这样就解决了全局→本地 episode index 的不匹配。
80
 
81
 
82
_LRDCls._get_query_indices = _fixed_gqi
第三个 monkey patch:用 _fixed_gqi 替换 LeRobotDataset._get_query_indices 方法。从此所有使用 episodes 子集的数据加载都会走修复过的逻辑。
83
 
84
sys.path.insert(0, "/opt/wall-x")
85
sys.argv = ["train_qact.py", "--config", "/opt/config_libero_train.yml", "--seed", "42"]
86
 
87
import runpy
88
 
89
runpy.run_path("/opt/wall-x/train_qact.py", run_name="__main__")
三个 monkey patch 都到位后,配置启动参数:sys.path.insert(0, '/opt/wall-x') 把 docker 内的源码目录加入头部(优先级最高);sys.argv 伪造成 ['train_qact.py', '--config', '/opt/config_libero_train.yml', '--seed', '42'](跳过 python train_libero_wrapper.py 这个 wrapper 脚本,直接让下游当做是 train_qact.py 被调用)。最后 runpy.run_path('/opt/wall-x/train_qact.py', run_name='__main__') 在当前进程上下文执行原始训练脚本,所有 monkeypatch 都生效。这是代理执行模式,核心思想是在原脚本启动前注入 hook,避免修改原脚本。

源码零改写,与仓库 /tmp/claude-1000/-home-xgwang-learn-02-wall-x/eb76002a-0a31-4eb7-b08f-c783c42cd312/scratchpad/train_libero_wrapper.py 逐字节一致(commit 97406f2)。 生成于 wall-x LIBERO 微调项目 · ← 返回导读