复现说明:本文记录在 PushT 环境下复现 LeWorldModel (LeWM) 的完整过程。实验在个人笔记本(RTX 5060 Laptop, 8GB VRAM)上完成,受限于显存与时间仅训练了 4 个 epoch,远未达到论文报告的收敛水平。如需直接体验完整效果,建议直接使用官方预训练模型(见第五节)。

一、实验概述

LeWorldModel(LeWM) 是一个基于 JEPA(Joint Embedding Predictive Architecture)的世界模型,通过联合嵌入预测训练,在潜空间中学习环境的动力学,从而支持基于模型的规划与控制。本次实验基于官方代码,在 PushT 机器人操作任务上完成从环境搭建、数据准备、模型训练到 CEM-MPC 规划评估的完整流程。

训练总时长:约 13 小时(03:16 ~ 16:29),共 4 个 epoch。

Epoch 完成时间 单 epoch 耗时
1 06:40 ~3h24m
2 09:56 ~3h16m
3 13:12 ~3h16m
4 16:29 ~3h17m

训练配置

参数
模型 ViT-tiny + ARPredictor (~15M 参数量)
Batch Size 32
图像尺寸 224×224
优化器 AdamW, lr=5e-5, weight_decay=1e-3
精度 bfloat16 AMP
Embed Dim 192
History Size 3
SIGReg λ 0.09
预测步数 1
数据集 pusht_expert_train.h5 (46GB)

图1:WandB 记录的训练过程。上方为 epoch 推进曲线(4 个 epoch),下方为训练损失下降趋势,可以看到损失在 4 个 epoch 内仍在持续下降,远未到达平台期。

二、遇到的问题及解决

复现过程并非一帆风顺。从环境搭建到最终评估,先后遇到了近 20 个问题,可归纳为以下五类。这里将它们完整记录下来,希望对后续复现 LeWM 的朋友有所帮助。

2.1 环境依赖问题

克隆仓库后按照 README 安装依赖,运行时仍然遇到一系列缺失模块错误——这是因为项目中有些子模块(如 stable_worldmodel)引入了额外的可选依赖,但并未在主 requirements.txt 中声明。

# 问题 原因 解决方式
1 stable_pretraining 未安装 README 未明确列出所有依赖 pip install stable-pretraining
2 lightning 未安装 同上 pip install lightning
3 imageio 未安装 stable_worldmodel 可选依赖缺失 pip install imageio
4 imageio[ffmpeg] 未安装 视频编码器缺失 pip install imageio[ffmpeg]
5 pygame 未安装 PushT 环境可视化依赖 pip install pygame
6 pymunk 未安装 PushT 物理引擎依赖 pip install pymunk
7 shapely 未安装 PushT 几何计算依赖 pip install shapely
8 cv2 (opencv) 未安装 图像处理依赖 pip install opencv-python

建议后来的复现者直接在环境中预装上述包,可以省去很多 ModuleNotFoundError 的排查时间。

2.2 CUDA / GPU 兼容问题

RTX 5060 Laptop 是 Blackwell 架构(计算能力 sm_120),对 PyTorch 版本有硬性要求。我踩了两个坑:

# 问题 原因 解决方式
1 CUDA out of memory batch_size=128 超出 8GB 显存 降至 batch_size=32
2 RTX 5060 不被 PyTorch 2.5.1 支持 Blackwell (sm_120) 需 PyTorch ≥ 2.7 升级至 torch 2.11.0+cu128

关键经验:如果你使用的是 RTX 50 系列显卡,务必使用 PyTorch 2.7 以上版本。旧版 PyTorch 的预编译二进制不包含 sm_120 的 kernel,运行时会直接报 CUDA 错误。

2.3 数据集路径问题

数据路径是新手最容易踩坑的地方。LeWM 使用了 stable_worldmodel 库的统一数据加载机制,默认期望特定的目录结构和数据格式。

# 问题 原因 解决方式
1 FileNotFoundError: Cannot resolve dataset STABLEWM_HOME 环境变量未设置 设置 STABLEWM_HOME=stable_wm_data/
2 同上 数据文件需放在 $STABLEWM_HOME/datasets/ 子目录 创建 datasets/ 目录并移入 .h5 文件
3 No format detected hdf5plugin 包未安装导致 HDF5 格式无法注册 pip install hdf5plugin
4 配置引用 .lance 但数据是 .h5 默认训练配置与实际数据格式不匹配 修改 config/train/data/pusht.yaml

其中第 4 条比较隐蔽——官方默认配置针对的是 .lance 格式数据集,但 PushT 的专家数据是 .h5 格式。需要手动修改配置文件中的数据集路径,否则加载阶段就会失败。

2.4 Windows 兼容性问题

在 Linux 服务器上这不会是一个问题,但在 Windows 上运行就遇到了:

# 问题 原因 解决方式
1 AttributeError: module 'signal' has no attribute 'SIGUSR1' stable_pretraining 直接引用了 Unix 信号量 print_signal_info 添加 getattr 平台兼容判断

具体修改方式是找到 stable_pretraining 中调用 signal.SIGUSR1 的位置,改为:

1
signal.SIGUSR1 if hasattr(signal, 'SIGUSR1') else None

这类 Unix 信号量在 Windows 下不存在,直接用 getattr 做防御性访问即可。

2.5 评估阶段问题

训练完成后进入评估,又遇到两个问题:

# 问题 原因 解决方式
1 save_panel_videos 空帧崩溃 video_stride=10 导致未录制视频的 env 帧为空 修改 save_panel_videos 跳过空帧 env
2 WandB entity 错误 配置中 entity 与登录账号不匹配 修改为正确的 WandB entity 名称

评估视频使用三栏面板布局(agent 预测 rollout / dataset ground truth / goal 状态),但由于 video_stride=10 的设置,并非每个 env 都有录制帧,需要在保存逻辑中跳过空帧。

三、评估结果

采用 CEM-MPC 规划进行 50 个 episode 的评估,参数如下:

  • CEM 采样:300 samples × 30 iterations
  • 规划 horizon:5
  • 模型:训练 4 epoch 的 checkpoint
1
2
Success Rate: 2.0%
Episode Success: 1 / 50

图2:WandB 记录的验证过程指标。

结果分析

仅训练 4 个 epoch 导致模型远未收敛。做个简单的对比:

维度 论文配置 本次配置
训练轮数 100 epochs 4 epochs
Batch Size 128 32
单 epoch 步数 全量 / 128 全量 / 32
总梯度更新量 ~论文基准 < 论文的 1%

4 epoch 的模型本质上还是随机初始化附近的状态,原因很明确:

  • ViT-tiny 编码器的视觉表征尚未充分学习,提取的潜空间特征质量不足
  • ARPredictor 的时序预测精度有限,无法准确建模多步动力学
  • SIGReg 正则化在早期 epoch 中尚未平衡好预测损失与分布约束的权重

成功率为 2%(1/50),与随机策略接近,说明 4 epoch 的模型尚未学到有效的世界模型表征。这并非模型本身的问题,纯粹是训练量不足。

四、资源受限说明

本次实验受限于以下约束条件:

  1. GPU 显存:RTX 5060 Laptop 约 8GB 可用 VRAM,无法支持论文默认的 batch_size=128。降至 32 后,单 epoch 耗时约 3h15m。按比例估算,完整 100 epoch 训练将需要约 325 小时(近 14 天),对个人笔记本而言难以承受。

  2. 训练时间:受限于单卡和个人时间安排,仅训练了 4 个 epoch,相当于论文推荐训练量的 4%。

  3. 对比参考:论文使用 A5000/A100 等 24GB+ 数据中心 GPU,batch_size=128 下 100 epoch 的训练量换算成有效梯度步数,远高于本实验。

图3:从 03:16 到 16:29,共约 13 小时连续训练的 WandB 时间线截图。

好在官方在 HuggingFace 上提供了预训练权重(quentinll/lewm-pusht),可以直接下载用于评估,跳过耗时的训练阶段。

五、如何获得更好的结果

根据本次复现的经验教训,推荐以下三种方案:

方案一:使用官方预训练模型(推荐)

最简单高效的方式,适合只想评估模型效果的场景:

1
2
hf download quentinll/lewm-pusht --local-dir stable_wm_data/hf_pusht
python eval.py --config-name=pusht.yaml policy=pusht/lewm

预训练模型经过了完整的 100 epoch 训练,评估成功率远超从零训练的模型。

方案二:从头训练(需要充足资源)

如果目标是完整复现论文结果,建议满足以下条件:

  • GPU:24GB+ 显存(如 A5000、A100、4090)
  • 配置:恢复默认 batch_size=128, max_epochs=100
  • 预计训练时间:A5000 上约 2~3 天

方案三:轻量级训练优化(有限资源)

如果只有消费级 GPU 但仍想从头训练,可以尝试以下优化:

  • img_size 减小至 112(降低 ViT 计算量约 4 倍)
  • embed_dim 减小至 128(减小潜空间维度)
  • 启用 precision=bf16 和 gradient checkpointing(以时间换显存)
  • 适当增加 batch_size(在显存允许范围内,越大越有利于 SIGReg 的分布约束)

这些措施可以在一定程度上弥补显存不足,但收敛效果和最终成功率仍会低于论文配置。

六、关键文件清单

以下是本项目中涉及的核心文件,供需要的读者参考:

文件 说明
train.py 训练入口脚本
eval.py 评估入口(已添加 video_stride 支持)
jepa.py JEPA 模型架构定义
module.py SIGReg 正则化器、Transformer、Predictor
config/train/lewm.yaml 训练主配置
config/train/data/pusht.yaml PushT 数据配置(已修正为 .h5
config/eval/pusht.yaml PushT 评估配置(已添加 video_stride: 10
stable_wm_data/checkpoints/lewm/ 训练 checkpoint 保存目录
stable_wm_data/lewm/ 评估视频输出目录
stable_wm_data/datasets/ 数据集存放目录

总结

这次 LeWM 复现是一次有价值的实践——虽然受限于消费级硬件最终只得到了 2% 的成功率,但整个过程中对 JEPA 架构的训练流程、环境搭建、数据管道和 CEM-MPC 规划有了更深入的理解。

几点核心体会:

  • 世界模型的训练门槛不低:LeWM 论文推荐的 100 epoch × batch_size=128 的全量训练,对消费级 GPU 来说计算成本太高。轻量级训练可以跑通流程,但离可用效果还有很大差距。
  • 预训练模型是务实之选:HuggingFace 上的 quentinll/lewm-pusht 权重可以直接跳过训练阶段,适合快速上手评估。
  • Windows 环境不是一等公民:深度学习代码默认 Linux 平台,Windows 上会遇到 signal.SIGUSR1 等平台 API 差异。后续可以考虑用 WSL2 替代原生 Windows 开发环境。
  • Blackwell 架构需要新版 PyTorch:RTX 50 系列用户务必使用 PyTorch ≥ 2.7,否则 CUDA kernel 无法运行。

参考链接