rtx4090_pi0_training_pitfalls

用 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

最终经验

  1. 多卡数量增加,不代表显存可以自动相加。
  2. 判断 FSDP 是否有效,必须看分片粒度,而不是只看有没有 --use_fsdp。
  3. root-only FSDP 可以降低常驻显存,但仍可能产生完整模型的瞬时恢复峰值。
  4. CPU offload 解决显存容量,代价是严重的 PCIe 传输开销。
  5. 8-bit optimizer 能降低优化器状态,但无法消除完整模型 all-gather 峰值。
  6. OpenPI JAX 的优势来自张量级 SPMD 分片,而不是简单的"JAX更省显存"。
  7. 四张 24 GB 卡训练 π0 时,EMA 的一份额外参数副本足以再次导致 OOM。
  8. 4090 上做全参数训练,显存不能只留几百 MiB 理论余量。
  9. LoRA 更适合快速验证;全参数微调更适合在数据与推理链路确认后进行。
  10. 训练完成只是开始,机器人项目最终仍要靠标定检查、动作审计和分级实机测试闭环。

一句话概括本次实验:

text 复制代码
LeRobot 全参:FSDP + CPU offload,能跑但慢。
LeRobot LoRA:适合 4090 快速验证。
OpenPI JAX 全参:4×4090、batch 4、关闭 EMA,是本机验证过的高效方案。
相关推荐
liuchangng1 小时前
类Jev项目Kev从入门到实战(7):如何选型:决策树与 checklist
算法·决策树·机器学习
IT研究室1 小时前
最新大数据毕业设计选题推荐-基于大数据的人工智能社交媒体情绪分析与可视化的设计与实现-大数据-Spark-Hadoop-Bigdata
大数据·人工智能·课程设计
Figo_Cheung1 小时前
Figo生成式人工智能潜空间计算最优参数规模研究
人工智能·空间计算
这张生成的图像能检测吗1 小时前
(论文速读)一种用于图像超分辨率的有效扩散变换结构
人工智能·计算机视觉·扩散模型·超分辨率
Eric.461 小时前
ComfyUI 本地部署 AI 漫剧与 AI 视频流水线:消费级显卡显存优化、角色一致性工程实战
人工智能·自动化·音视频·comfyui·ai漫剧
白猫不黑1 小时前
AI自动化漏洞挖掘从入门到进阶超详细学习路线
人工智能·学习·web安全·网络安全·信息安全·渗透测试·自动化
百度Geek说1 小时前
Agentic Harness Workflow 框架:把 AI Coding 变成工程化流程
人工智能
Data-Miner1 小时前
离线AI制表:模型适配与脚本固化实操
大数据·数据库·人工智能·excel
指针向南1 小时前
视频截图发黑先检查透明区域
图像处理·人工智能·计算机视觉·音视频