LeRobot 中的 FastWAM 世界动作模型:架构原理、训练评估与配置实战
【免费下载链接】lerobot🤗 LeRobot: Making AI for Robotics more accessible with end-to-end learning项目地址: https://gitcode.com/GitHub_Trending/le/lerobot
FastWAM(Fast World Action Model)是一类面向机器人控制的"世界动作模型"(World Action Model)策略:训练阶段保留完整的视频建模,推理阶段却不做测试时的未来想象,而是直接对动作去噪预测,从而大幅降低推理开销。本文以 LeRobot 仓库中的 fastwam README 为核心骨架,结合 FastWAM 配置类、策略实现、处理器管线 与 官方评测文档,讲解如何在 LeRobot 中用policy.type=fastwam完成训练、评测与真机部署,并深入解读其 MoT 混合注意力、缓存式动作推理与关键配置参数。
FastWAM 是什么:一个"训练想象、推理务实"的世界动作模型
FastWAM 的研究论文提出并回答了这样一个问题:世界动作模型在测试阶段真的需要"未来想象"吗?传统 World Action Model(WAM)在推理时会先生成未来的观测视频、再基于想象出的未来帧规划动作,代价是极高的测试时计算量。FastWAM 的核心主张是:把"未来想象"的能力留在训练期(作为视频自监督信号),推理期只用首帧观测直接去噪出动作。
LeRobot 将该模型以标准策略 API 集成:
- 通过
policy.type=fastwam完成配置(注册于 policies/init.py 与 factory.py); - 通过
lerobot-train训练、lerobot-eval评测、lerobot-rollout真机部署; - 通过
PreTrainedPolicy接口完成 checkpoint 的保存与加载。
从源码看,LeRobot 的 FastWAM 集成覆盖了四层能力(见 modeling_fastwam.py):
- 批数据适配:把 LeRobot 标准 batch(多路相机帧、单步
observation.state、action、语言task)转换为 FastWAM 原生 sample(video、action、context/context_mask、逐帧proprio),见_batch_to_training_sample; - 动作块推理:通过
predict_action_chunk输出形状为[batch, action_horizon, action_dim]的动作块,select_action内部维护一个n_action_steps长度的队列,消费完再预测新块; - checkpoint 保存/加载:只保存可训练的 MoT DiT 权重,冻结的 Wan VAE 与 UMT5 文本编码器从 diffusers/transformers 仓库按需加载,显著缩小 checkpoint 体积;
- 可配置的 LIBERO 夹爪后处理:通过
toggle_action_dimensions把夹爪动作映射到 FastWAM 评测管线约定的约定符号。
模型架构:视频专家 + 动作专家 + MoT 混合注意力
FastWAM 在 LeRobot 中的实现位于 wan/modular.py,整体由以下组件构成:
| 组件 | 来源 | 训练状态 |
|---|---|---|
视频专家WanVideoDiT(约 5B) | Wan-AI/Wan2.2-TI2V-5B | 可训练(默认) |
动作专家ActionDiT | FastWAM 自研(默认 30 层、hidden 1024) | 可训练 |
| Wan VAE | Wan-AI/Wan2.2-TI2V-5B-Diffusers | 冻结 |
| UMT5 文本编码器 | google/umt5-xxl | 冻结 |
本体感觉编码器proprio_encoder | 随机初始化 Linear | 可训练 |
MoT:跨专家混合注意力
MoT(Mixture-of-Transformers)把视频专家与动作专家的 transformer block按层重组:每一层由一个MoTLayer拥有该层所有专家(video/action)的 block,并将 Q/K/V 拼接后执行一次联合混合注意力(_forward_joint),随后再切回各专家做 MLP 与残差。这一设计有两个关键工程点:
- Attention 必须用 SDPA 而非 FlashAttention:MoT 路由需要任意的布尔
[query, key]注意力掩码(video→video 因果、action→action 全真、action→首帧视频 token 三种掩码组合,见_build_mot_attention_mask),而 FlashAttention 的 varlen API 无法表达这种任意掩码。因此安装flash-attn对 FastWAM 路径没有任何效果(scaled_dot_product_attention内部可能自行选择 PyTorch 的 flash/mem-efficient/math kernel,这与flash-attn包无关)。 - FSDP 包装单元:
MoTLayer被声明为_fsdp_wrap_modules,每个MoTLayer.forward是 FSDP 唯一可挂载的调用边界,分布式训练时按层 all-gather 参数。
缓存式动作推理:视频只 prefill 一次
FastWAM 推理的核心优化在infer_action(modular.py 第 1786 行起):
- 用首帧观测经 VAE 编码得到
first_frame_latents,timestep_video固定为 0; mot.prefill_video_cache只跑一遍视频专家,把每层的 K/V 缓存下来;_predict_action_noise_with_cache在随后的每个去噪步中,只运行动作专家,其注意力同时 attend 到缓存的视频 K/V 与当前动作 K/V;- 动作去噪调度器使用
WanContinuousFlowMatchScheduler,与 Wan2.2 的采样 sigma 兼容。
当compile_action_infer=true时,视频 prefill 与缓存动作去噪两条路径都会被torch.compile(mode="reduce-overhead", fullgraph=True)编译(并借助 CUDA Graph 标记),首次调用会编译预热、后续同形状调用复用计算图;test_fastwam_policy.py 中的test_cached_inference_paths_support_fullgraph_compile验证了这两条缓存路径在 fullgraph 编译下与 eager 输出一致。
安装:启用 fastwam 与 libero 两个 extra
FastWAM 依赖 transformers(UMT5 文本编码器/分词器)与 diffusers(Wan VAE),它们被收进fastwamextra(定义于 pyproject.toml 第 241 行)。从源码安装:
pip install -e ".[fastwam]"如需在 LIBERO 上评测,再叠加liberoextra(会额外安装数据集依赖、Linux 下的hf-libero与scipy):
pip install -e ".[fastwam,libero]"在基础安装上直接构造策略会立即失败并给出可操作的提示——FastWAMPolicy.__init__开头通过require_package("transformers", extra="fastwam")与require_package("diffusers", extra="fastwam")做 fail-fast 校验。
数据要求
FastWAM 期望的 LeRobot 数据集需要满足:
- 一个或多个视觉观测,且所有相机宽度之和等于
policy.image_size[1]。默认配置是单一图像特征observation.images.image,形状(3, 224, 448);若数据集有top、wrist两路相机,则需通过policy.input_features声明两路特征,高度同为224、宽度和为448。set_dataset_feature_metadata(configuration_fastwam.py)会在make_policy拿到数据集元数据后自动用真实相机键重建输入特征,每路相机宽度为image_size[1] // num_cameras,且模型内部的_stack_video_from_images会做对应 resize,因此支持异构源分辨率(如 480×640); observation.state:当policy.proprio_dim不为None时必需,其维度需等于proprio_dim;action:维度需等于policy.action_dim;- 语言指令:通过数据集
task字段提供,或直接提供预计算的context/context_mask张量。训练时_prompt_from_batch会把 task 套入提示模板"A video recorded from a robot's point of view executing the following instruction: {task}"再用 UMT5 编码。
值得注意的约束校验(FastWAMConfig.validate_features):num_video_frames - 1必须能被action_video_freq_ratio整除,且按比例下采样后的视频帧数必须满足model_video_frames % 4 == 1(VAE 时间维 4× 压缩的硬性要求,默认num_video_frames=33, ratio=4得到 9 帧),action_horizon还必须能被下采样后的视频 transitions 数整除。
使用:训练、评测与真机部署
训练一个新策略
lerobot-train \ --dataset.repo_id=your-org/your-dataset \ --policy.type=fastwam \ --policy.action_dim=7 \ --policy.proprio_dim=8 \ --policy.action_horizon=32 \ --policy.n_action_steps=10 \ --policy.image_size='[224,448]' \ --output_dir=./outputs/fastwam_training \ --job_name=fastwam_training \ --steps=300000 \ --batch_size=8 \ --policy.device=cuda训练时FastWAMPolicy.forward依次执行build_inputs(VAE 编码视频、UMT5 编码 prompt、追加 proprio token)、_sample_training_targets(视频与动作分别加噪并计算 flow-matching 目标)、_run_training_mot(联合前向)与双损失计算。总损失为:
loss = lambda_video * loss_video + lambda_action * loss_action其中loss_video是视频 latent 的 MSE(带image_is_pad掩码、按时间步加权),loss_action是动作的 MSE(带action_is_pad掩码、按时间步加权),权重由policy.loss='{"lambda_video": 1.0, "lambda_action": 1.0}'控制。训练日志会输出loss_video与loss_action两个指标。
用官方 checkpoint 复现 LIBERO 评测
LeRobot 发布了 LIBERO 上的无条件(unconditioned)2 相机 224 分辨率 checkpointZibinDong/fastwam_libero_uncond_2cam224。评测命令如下:
lerobot-eval \ --policy.path=ZibinDong/fastwam_libero_uncond_2cam224 \ --policy.device=cuda \ --policy.torch_dtype=float32 \ --policy.n_action_steps=10 \ --env.type=libero \ --env.task=libero_spatial \ --env.observation_height=224 \ --env.observation_width=224 \ --eval.batch_size=1 \ --eval.n_episodes=50 \ --seed=0 \ --env.episode_length=300官方文档给出的完整复现命令(README 内嵌行)在四个 LIBERO suite 上的结果如下,官方标注评测硬件为单卡 H20 140GB:
| Suite | Success rate | n_episodes |
|---|---|---|
| libero_spatial | 97.6% | 500 |
| libero_object | 99.0% | 500 |
| libero_goal | 95.0% | 500 |
| libero_10 | 94.0% | 500 |
| average | 96.4% | 2000 |
参数速查:
libero_goal、libero_spatial、libero_object:--env.episode_length=300;libero_10:需改用--env.task=libero_10 --env.episode_length=600(任务更长);- 官方 README 的复现命令中还出现了
--env.observation_height=256 --env.observation_width=256的高分辨率变体(首帧条件帧在模型内会被 resize 到模型自有分辨率,见_prepare_infer_image),两种观测分辨率均可使用; --policy.torch_dtype=float32:官方评测以 float32 运行(配置默认bfloat16,评测时建议按官方命令覆盖);--eval.n_episodes=50:单次跑 50 个 episode,四个 suite 各 10 次取总计 500/2000 的统计口径。
真机 rollout
lerobot-rollout \ --robot.type=so101_follower \ --robot.port=/dev/ttyACM0 \ --policy.path=your-org/fastwam-real-robot配置参数深度解读
图像特征与输入分辨率
policy.image_size是拼接后的 FastWAM 图像张量尺寸(height, width),每个图像特征必须是(3, height, camera_width),且所有camera_width之和等于配置宽度。processor_fastwam.py中不设置 resize 步骤——模型是输入分辨率的唯一权威:_stack_video_from_images/_prepare_infer_image会在训练与推理的所有路径上把每路相机 resize 到逐相机目标尺寸。若在 preprocessor 中再做一次 resize,会与微调数据集(继承自基础 checkpoint 的相机几何)冲突,导致拼接宽度变成相机数量的 N 倍。
归一化策略(normalization_mapping)值得特别说明:视觉输入使用 IDENTITY(不归一化),图像以[0, 1]传入,模型在 VAE 编码边界统一映射到[-1, 1]。原因(见 processor_fastwam.py 注释)是:微调时lerobot_train.py会用真实数据集的 per-channel 图像 std 覆盖归一化统计量,而真实帧间的亮度方差极小,会瞬间把图像推出[-1, 1]导致饱和。STATE 与 ACTION 仍用数据集统计做 MEAN_STD 归一化。
动作分块
policy.action_horizon(默认 32):训练监督与推理预测的未来动作数,也是每次predict_action_chunk输出的动作块长度;policy.n_action_steps(默认 32,评测官方用 10):策略消费完这么多步后才重新预测新块,必须满足n_action_steps <= action_horizon;action_video_freq_ratio(默认 4):动作按视频帧率的该倍数采样,模型实际看到的视频帧为(num_video_frames - 1) // ratio + 1帧(默认 33 → 9),每帧对应ratio个动作。
Wan 组件与 checkpoint 策略
model_id必须是Wan-AI/Wan2.2-TI2V-5B或显式本地路径(_validate_wan_model_id强制校验);- 文本编码器默认从
Wan-AI/Wan2.2-TI2V-5B-Diffusers加载,分词器来自google/umt5-xxl; - 冻结的 VAE / 文本编码器通过
object.__setattr__挂为未注册属性:它们不出现在state_dict()、parameters()与优化器中,checkpoint 只保存可训练 DiT(见 modular.py 的_apply覆写与测试test_save_pretrained_excludes_frozen_components); base_model_id默认lerobot/fastwam_base,当视频/动作 DiT 结构与基座兼容时自动路由到预训练权重(is_fastwam_base_compatible_config)。
冻结视频专家以节省显存与优化器开销
--policy.freeze_video_expert=truefreeze_video_expert会冻结约 5B 的视频专家(model.video_expert及其在 MoT layer 中重挂载的 block,见 modeling_fastwam.py),只训练动作专家与 proprio 编码器,显著削减 AdamW 优化器状态。注意此时应同时设置--policy.loss.lambda_video=0,跳过已无梯度的视频损失计算。
LIBERO 夹爪动作符号映射(Action Toggle)
官方 LIBERO checkpoint 约定policy.toggle_action_dimensions='[-1]'(README 注释指出默认即此行为):后处理器FastWAMActionToggleProcessorStep会把指定维度的动作a经sign(1 - 2a)映射到{-1, +1},对齐原版 FastWAM 评测管线使用的夹爪开合约定(见 processor_fastwam.py 与测试中的toggle_action_dimensions=[-1]用例)。
采样与调度器
num_inference_steps(默认 10,官方评测亦可加大):动作去噪步数;video_scheduler/action_scheduler:各自独立配置train_shift、infer_shift(默认 5.0)与num_train_timesteps(默认 1000);inference_seed(默认 42)与rand_device(默认 cpu):控制动作与视频 latent 的随机初始化,保证评测可复现;text_cfg_scale(默认 1.0)与negative_prompt:文本 CFG 相关,默认关闭;torch_dtype(默认 bfloat16):评测官方 checkpoint 时建议按官方命令显式设为 float32。
可复现性与评测注意事项
- 评测统计口径:README 表格中的 success rate 基于每个 suite 500 episodes(共 2000 episodes);单条复现命令只跑
n_episodes=50,用于快速验证,如需对齐官方数字应多次运行并累计; - 随机种子:
--seed=0固定动作采样噪声,配合inference_seed才能得到可复现结果; - 跨本体微调:
_load_as_safetensor支持 shape-aware 加载——当从 7-DoF/8 维 LIBERO checkpoint 微调到 6-DoF/6 维机械臂时,动作编码器/输出头与 proprio 编码器的形状不匹配张量会被丢弃并重新初始化,其余兼容权重全部加载(strict=False 语义,见 modeling_fastwam.py); - 旧 checkpoint 兼容:
MoT._load_from_state_dict会把 MoTLayer 重构前以mixtures.<name>.blocks.<i>.*命名的旧权重原位重映射到layers.<i>.blocks.<name>.*,因此官方发布的 checkpoint 可直接加载。
延伸阅读与引用
- 论文:《Fast-WAM: Do World Action Models Need Test-time Future Imagination?》(arXiv:2603.16666),作者 Tianyuan Yuan、Zibin Dong、Yicheng Liu、Hang Zhao;
- 基础视频模型:Wan-AI/Wan2.2-TI2V-5B;上游已发布 checkpoint:
yuanty/fastwam; - 仓库内的配套实现与测试:
wan子包说明见 wan/README.md(其中记录了 Wan2.2 源码的 vendored 子集与提交号、SDPA 替换 flash-attn 的裁剪说明);完整策略测试见 tests/policies/fastwam/test_fastwam_policy.py。
@article{yuan2026fastwam, title = {Fast-WAM: Do World Action Models Need Test-time Future Imagination?}, author = {Tianyuan Yuan and Zibin Dong and Yicheng Liu and Hang Zhao}, journal = {arXiv preprint arXiv:2603.16666}, year = {2026}, url = {https://arxiv.org/abs/2603.16666} }【免费下载链接】lerobot🤗 LeRobot: Making AI for Robotics more accessible with end-to-end learning项目地址: https://gitcode.com/GitHub_Trending/le/lerobot
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考