用 4 张 RTX 4090 训练 π0:LeRobot FSDP、CPU Offload 与 OpenPI JAX 踩坑实录
本文记录一次真实的 π0 全参数微调实验。硬件为 4 张 RTX 4090(单卡约 24 GB),先后测试了 LeRobot/PyTorch 与 OpenPI/JAX 两条训练路线。重点不是给出泛化的显存估算,而是解释实际发生的 OOM、为什么 FSDP 没能直接解决问题,以及最终可稳定运行的配置。
结论先行
在这台 4×RTX 4090 工作站上:
- LeRobot π0 全参数训练即使开启 FSDP,仍会在反向传播阶段 OOM。
- 原因不是 FSDP 没生效,而是 π0 当时只能使用 root-only FSDP,反向阶段每张卡仍需临时恢复约 6.52 GiB 的完整 BF16 参数。
- LeRobot 开启 CPU offload 后能够完成全参数训练,但速度下降到约 20~34 秒/step。
- LeRobot LoRA 更适合在 24 GB 显卡上快速验证训练与推理流程。
- OpenPI JAX 使用张量级 SPMD 分片,4 卡可以直接进行全参数训练。
- JAX 配置必须关闭 EMA;否则每卡新增约 3.26 GiB,24 GB 显存基本没有足够的临时计算余量。
- 最终验证可用的 JAX 配置是:全局 batch size 4、4 卡 FSDP、关闭 EMA、关闭 XLA 显存预分配。
- 实测 17,020 步约耗时 5 小时 19 分钟,稳定速度约 1.1~1.2 秒/step。
实验环境
硬件:
text
GPU:4 × NVIDIA RTX 4090
单卡可用容量:约 23.52 GiB(nvidia-smi 显示约 24564 MiB)
总物理显存:约 96 GB
一个常见误解是:4 张 24 GB 显卡等于一张 96 GB 显卡。实际上,设备之间的显存并不是统一地址空间。模型能否利用总显存,取决于框架如何切分参数、梯度、优化器状态,以及计算过程中是否会临时恢复完整张量。
π0 的参数到底占多少显存?
LeRobot 训练日志记录的可训练参数量为:
text
3,501,372,176
只计算模型参数本身:
text
BF16:3,501,372,176 × 2 byte ÷ 1024³ ≈ 6.52 GiB
FP32:3,501,372,176 × 4 byte ÷ 1024³ ≈ 13.04 GiB
但全参数训练还需要保存或产生:
- 梯度;
- Adam 一阶动量;
- Adam 二阶动量;
- 前向激活与反向计算图;
- collective 通信缓冲区;
- 临时矩阵和框架运行时缓存。
因此,"模型权重只有 6.52 GiB"并不代表 24 GB 显卡可以训练它。训练的决定因素通常是峰值显存,而不是静态模型文件大小。
为什么普通多卡 DDP 不够?
数据并行 DDP 会在每张 GPU 上保存完整模型:
text
GPU 0:完整 π0
GPU 1:完整 π0
GPU 2:完整 π0
GPU 3:完整 π0
DDP 可以提高吞吐量,但不会把单卡上的模型参数、梯度和优化器状态简单除以 GPU 数量。实际实验中,普通四卡 DDP 在 Adam 状态初始化阶段就超过了单卡 24 GB 的容量。
这也是我们转向 FSDP 的原因。
LeRobot FSDP 为什么还是 OOM?
当时的关键配置为:
bash
--use_fsdp
--fsdp_version 1
--fsdp_sharding_strategy FULL_SHARD
--fsdp_auto_wrap_policy no_wrap
--fsdp_offload_params false
FULL_SHARD 会切分参数、梯度和优化器状态;问题出在 no_wrap。
root-only FSDP
理想的细粒度 FSDP 应该按 Transformer block 包装:
text
π0
├── Vision block 1 ← 一个 FSDP 单元
├── Vision block 2 ← 一个 FSDP 单元
├── Language block 1 ← 一个 FSDP 单元
├── Language block 2 ← 一个 FSDP 单元
└── Action expert ← 一个 FSDP 单元
计算一个 block 时,只需临时恢复该 block 的参数,完成后即可重新分片。
但 no_wrap 形成的是:
text
FSDP(
整个 35 亿参数的 π0
)
也就是说,整个模型只有一个根 FSDP 单元。
参数在静止状态下确实被分片。例如三卡训练时,6.52 GiB 的 BF16 参数可以近似分为:
text
GPU 0:约 2.17 GiB
GPU 1:约 2.17 GiB
GPU 2:约 2.17 GiB
但模型执行需要完整参数。FSDP 会进行 all-gather:
text
各卡参数分片
↓ all-gather
每张卡临时得到完整 FSDP 单元
↓ 计算
重新分片并释放完整副本
由于这个 FSDP 单元就是整个 π0,每张 GPU 在反向传播时都要临时恢复完整的 35 亿参数。
6.52 GiB 不是巧合
实际 OOM 日志为:
text
torch.OutOfMemoryError: CUDA out of memory.
Tried to allocate 6.52 GiB.
而完整 BF16 π0 参数恰好是:
text
3,501,372,176 × 2 byte ÷ 1024³ ≈ 6.52 GiB
两者完全吻合。这说明失败点不是普通图像 batch 或零散缓存,而是反向阶段恢复完整根参数所需的大块缓冲区。
三卡无 offload 测试的实际状态为:
text
单卡总容量:约 23.52 GiB
PyTorch 已分配:约 19.59 GiB
剩余空间:约 3.07~3.09 GiB
下一次申请:6.52 GiB
因此 OOM 是确定结果:
text
3.08 GiB < 6.52 GiB
为什么不用按层 FSDP?
我们并非没有尝试细粒度包装。实际测试过 nested FSDP 和 FSDP v2,但 LeRobot 当时的 π0 实现会手工执行内部 Transformer layer,和嵌套 FSDP 对模块及参数的管理方式不兼容。
遇到的错误包括:
- nested FSDP 与 π0 手工 Transformer layer 执行路径不兼容;
- FSDP v2 中普通 Tensor 与 DTensor 混用;
- 视觉塔卷积进入分布式算子时出现 Tensor/DTensor 类型冲突。
所以 no_wrap 并不是随意选择,而是在保持模型语义可运行的前提下退回到 root-only FSDP。
代价就是:常驻状态可以分片,但完整根模块的计算峰值依然存在。
四卡加 8-bit optimizer 为什么还不够?
增加到四张卡并使用 8-bit optimizer 后,常驻显存确实下降了。某次实测中:
text
PyTorch 已分配:16.35 GiB
剩余空间:6.33 GiB
反向阶段需要:6.52 GiB
只差:
text
6.52 - 6.33 ≈ 0.19 GiB
但显存判断没有"差一点就算成功"。只要连续大块内存无法分配,训练就会终止。
这里也能看出 FSDP 和 8-bit optimizer 并非没有生效:它们已经显著降低了常驻占用。真正无法消除的是 root-only FSDP 每卡都需要的 6.52 GiB 完整参数峰值。
CPU offload 为什么有效?
最终 LeRobot 全参数训练采用:
bash
--use_fsdp
--fsdp_version 1
--fsdp_sharding_strategy FULL_SHARD
--fsdp_auto_wrap_policy no_wrap
--fsdp_offload_params true
--fsdp_use_orig_params true
CPU offload 会把暂时不用的参数分片放到系统内存,从而降低 GPU 的常驻显存基线。这样反向阶段再次申请 6.52 GiB 完整根参数时,GPU 才能腾出足够空间。
实际结果:
text
不开 CPU offload:第一步之后 backward OOM
打开 CPU offload:连续优化步骤成功,并最终完成训练
代价是 PCIe 数据搬运。烟雾测试的更新时间为:
text
第 1 步:约 33.99 秒
第 2 步:约 21.70 秒
CPU offload 解决的是容量问题,不是计算效率问题。
OpenPI JAX 为什么可以直接训练?
OpenPI JAX 虽然也把配置项称为 FSDP,但其实现方式与 PyTorch root-only FSDP 不同。
它会遍历训练状态中的每个数组。对于大于 4 MiB 的矩阵或高维张量,选择一个可被设备数整除的维度,然后直接沿该维度切分。
例如一个大型 MLP 权重:
text
shape = (18, 2, 2048, 16384)
总大小 = 4608 MiB
四卡时沿最后一维切分:
text
GPU 0:(18, 2, 2048, 4096)
GPU 1:(18, 2, 2048, 4096)
GPU 2:(18, 2, 2048, 4096)
GPU 3:(18, 2, 2048, 4096)
每卡只保存:
text
4608 MiB ÷ 4 = 1152 MiB
另一个 embedding 权重总大小为 2009 MiB,四卡后每卡约 502.25 MiB。
关键差别是:XLA 会把矩阵运算本身编译为分布式 SPMD 计算。
如果权重沿输出维度切分:
text
W = [W0 | W1 | W2 | W3]
GPU 0:Y0 = X × W0
GPU 1:Y1 = X × W1
GPU 2:Y2 = X × W2
GPU 3:Y3 = X × W3
输出也可以保持分片,不需要先在每张 GPU 上拼出完整 W。如果沿输入维度切分,XLA 则计算局部结果并通过 collective reduce 合并。
因此,通信发生在分片计算结果之间,而不是在每次反向前给每张卡恢复一套完整 35 亿参数模型。
JAX 连 Adam 状态也一起分片
OpenPI 对整个训练状态 pytree 应用相同的 sharding 规则,包括:
text
params
optimizer.mu
optimizer.nu
大型权重、Adam 一阶状态和 Adam 二阶状态都按四卡切分。这使显存压力比较均匀,而不是只切参数、却把优化器状态或其他副本完整留在某张卡上。
为什么必须关闭 EMA?
本次所用 OpenPI 代码中的默认 EMA 配置为:
python
ema_decay=0.99
EMA 并不是几个统计值,而是额外保存一整套模型参数:
text
训练参数 params
EMA 参数 ema_params
JAX 训练状态中的参数按 FP32 计算时,总参数大小约为:
text
3,501,372,176 × 4 byte ÷ 1024³ ≈ 13.04 GiB
四卡分片后,EMA 每卡仍需新增约:
text
13.04 GiB ÷ 4 ≈ 3.26 GiB
关闭 EMA 的成功训练中,GPU 峰值记录约为:
text
GPU 0:17073 MiB
GPU 1:17073 MiB
GPU 2:17073 MiB
GPU 3:20282 MiB
GPU 3 只剩:
text
24564 - 20282 = 4282 MiB
如果打开 EMA:
text
20282 + 约 3339 ≈ 23621 MiB
只剩不到 1 GiB,无法可靠容纳 XLA 临时输出、collective 缓冲区和运行时波动。
早期 JAX 测试也确实出现过 OOM:训练状态初始化时再申请 502~576 MiB 即失败。关闭 EMA 是最终配置能够稳定训练的重要条件,而不是可有可无的优化。
最终验证成功的 JAX 配置
训练配置核心参数:
python
batch_size = 4
fsdp_devices = 4
ema_decay = None
num_workers = 1
环境变量:
bash
export HF_HUB_OFFLINE=1
export TRANSFORMERS_OFFLINE=1
export XLA_PYTHON_CLIENT_PREALLOCATE=false
export XLA_PYTHON_CLIENT_MEM_FRACTION=0.80
其中:
- 全局 batch size 为 4,与四张 GPU 对齐;
- 模型参数和 Adam 状态跨四卡分片;
- 关闭 EMA,避免第二套模型参数;
- 禁止 XLA 启动时预占大部分显存;
- 训练前确保 GPU 上没有其他大显存任务。
最终实测:
text
训练步数:17,020
稳定速度:约 1.1~1.2 秒/step
总耗时:约 5 小时 19 分钟
最终 checkpoint:正常保存
训练跑通不等于机器人能完成任务
显存问题解决后,还必须把模型训练问题和机器人系统问题分开。
实际项目中还遇到过:
- 相机名称与物理相机语义对应错误;
- 数据采集与推理使用的关节标定不一致;
- 物体摆放范围太散;
- 真正决定抓取成败的对准、下降和闭合只占长轨迹的一小部分;
- 动作过快,关键抓取阶段缺少足够帧;
- action chunk 太长,执行时使用过期图像;
- 安全限幅连续截断后,原始轨迹被扭曲;
- EEF 模型经过 IK 后缺少关节连续性与残差约束。
所以至少要分别验证:
text
训练是否数值正常
数据语义是否正确
标定是否一致
离线动作是否合理
实机单步是否安全
闭环任务是否真正成功
Loss 下降只能证明优化器在工作,不能证明机器人会抓取。
action chunk 与 execute horizon
这两个概念很容易混淆:
text
chunk size:模型一次预测多少个动作
execute horizon:真正执行其中多少步后重新观察和推理
即使模型一次预测 20 或 50 步,也不应在没有验证的情况下全部执行。更稳妥的实机策略是:
text
预测 20~50 步
只执行 1~5 步
重新采集图像和关节状态
再次推理
长 chunk 全量执行容易造成误差累积,也可能跨过对准、下降或闭合等关键阶段。
动作限幅不是轨迹修复器
类似下面的设置:
text
max_relative_target = 2°
只能防止单次指令突然跳得太远。如果模型连续输出错误方向:
text
每次错误 2° × 多次执行 = 仍然会逐渐偏离
完整的实机安全层还应包含:
- 关节绝对限位;
- EEF 单步位移与旋转限位;
- IK 残差检查;
- IK 关节连续性约束;
- 连续触发相对限幅时自动停止;
- 工作空间边界;
- 较短的 execute horizon;
- dry-run、单步、短闭环、完整闭环的分级验证流程。
路线选择建议
| 目标 | 建议方案 |
|---|---|
| 单卡快速验证数据和推理流程 | LeRobot LoRA |
| 少量 GPU 快速实验 | LeRobot LoRA |
| LeRobot 全参数训练 | FSDP v1 + root-only + CPU offload |
| 4×4090 高效全参数训练 | OpenPI JAX 四卡分片 |
| 需要 EMA | 更大显存、更多设备、CPU EMA 或离线权重平均 |
| 实机早期验证 | dry-run + 单步 + 短 execute horizon |
最终经验
- 多卡数量增加,不代表显存可以自动相加。
- 判断 FSDP 是否有效,必须看分片粒度,而不是只看有没有
--use_fsdp。 - root-only FSDP 可以降低常驻显存,但仍可能产生完整模型的瞬时恢复峰值。
- CPU offload 解决显存容量,代价是严重的 PCIe 传输开销。
- 8-bit optimizer 能降低优化器状态,但无法消除完整模型 all-gather 峰值。
- OpenPI JAX 的优势来自张量级 SPMD 分片,而不是简单的"JAX更省显存"。
- 四张 24 GB 卡训练 π0 时,EMA 的一份额外参数副本足以再次导致 OOM。
- 4090 上做全参数训练,显存不能只留几百 MiB 理论余量。
- LoRA 更适合快速验证;全参数微调更适合在数据与推理链路确认后进行。
- 训练完成只是开始,机器人项目最终仍要靠标定检查、动作审计和分级实机测试闭环。
一句话概括本次实验:
text
LeRobot 全参:FSDP + CPU offload,能跑但慢。
LeRobot LoRA:适合 4090 快速验证。
OpenPI JAX 全参:4×4090、batch 4、关闭 EMA,是本机验证过的高效方案。