LeWM 世界模型复现
复现说明:本文记录在 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 | Success Rate: 2.0% |
图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 的模型尚未学到有效的世界模型表征。这并非模型本身的问题,纯粹是训练量不足。
四、资源受限说明
本次实验受限于以下约束条件:
GPU 显存:RTX 5060 Laptop 约 8GB 可用 VRAM,无法支持论文默认的
batch_size=128。降至 32 后,单 epoch 耗时约 3h15m。按比例估算,完整 100 epoch 训练将需要约 325 小时(近 14 天),对个人笔记本而言难以承受。训练时间:受限于单卡和个人时间安排,仅训练了 4 个 epoch,相当于论文推荐训练量的 4%。
对比参考:论文使用 A5000/A100 等 24GB+ 数据中心 GPU,
batch_size=128下 100 epoch 的训练量换算成有效梯度步数,远高于本实验。
图3:从 03:16 到 16:29,共约 13 小时连续训练的 WandB 时间线截图。
好在官方在 HuggingFace 上提供了预训练权重(quentinll/lewm-pusht),可以直接下载用于评估,跳过耗时的训练阶段。
五、如何获得更好的结果
根据本次复现的经验教训,推荐以下三种方案:
方案一:使用官方预训练模型(推荐)
最简单高效的方式,适合只想评估模型效果的场景:
1 | hf download quentinll/lewm-pusht --local-dir stable_wm_data/hf_pusht |
预训练模型经过了完整的 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 无法运行。
参考链接:

