针对 Robocasa365 的 FSDP 对模型参数分块时部分问题的解决方法

no_shardfull_shard:此前参数分块问题与代码级解决办法

下面是此前将 Pi0.5 Fixed ER 从 no_shard 改成真正 full_shard 时遇到的主要问题和实际解决方式。

问题 1:no_shard 并不解决大模型的训练状态复制

原始配置是:

1
2
3
actor.fsdp_config:
strategy: fsdp
sharding_strategy: no_shard

这看起来用了 FSDP,但 no_shard 的核心行为是:

  • 每个 rank 持有完整参数
  • 每个 rank 持有完整梯度
  • 每个 rank 持有完整 optimizer state

它更多接近传统 data parallel 的显存模型,而不是参数分片。

因此 4 卡可以提升吞吐,但无法把单卡容不下的全参 Pi0.5 训练状态拆到多张卡。若某个 rank 的单卡参数+optimizer+activation OOM,增加 GPU 数不一定解决。

解决方式是在 /export/pgs/heqijun/RLinf/examples/long_cl/config/robocasa365_long_cl.yaml 将配置改为:

1
2
actor.fsdp_config:
sharding_strategy: full_shard

这样参数、gradient、optimizer state 才会被分片。


问题 2:full-shard 后,rollout weight sync 拿到了 1D local shard

改成 full_shard 后,初始报错是:

1
2
3
4
ValueError: Shape mismatch for key
paligemma_with_expert.paligemma.model.language_model.embed_tokens.weight:
expected torch.Size([257152, 2048]),
got torch.Size([131661824])

这两个数的关系正好说明了问题:

$$
257152 \times 2048 / 4 = 131661824
$$

即 patch syncer 收到的是 4 卡中的一个一维 local shard,而 rollout 端期望的是完整二维 embedding matrix。

根因在:/export/pgs/heqijun/RLinf/rlinf/workers/actor/embodied_fsdp_actor_worker.py

原始实现:

1
2
3
4
5
def get_rollout_state_dict(self) -> dict:
return self.get_model_state_dict(
cpu_offload=False,
full_state_dict=False,
)

full_state_dict=False 对 full-shard 是正确的 checkpoint shard 导出方式,但不适用于当前 patch syncer。patch syncer 的协议是:

1
2
3
actor 导出完整命名参数 tensor
→ patch/delta encoder
→ rollout worker 将它 patch 到完整 rollout model

它无法理解 FSDP local shard。

解决方式是只在 actor → rollout 同步边界使用:

1
2
3
4
self.get_model_state_dict(
cpu_offload=False,
full_state_dict=True,
)

即:

1
2
训练期间仍然 full-shard;
只有同步给 rollout model 时临时 materialize 完整 state dict。

这样避免把 full tensor 常驻在每个训练 rank,但满足 rollout 端的完整 tensor shape contract。


问题 3:Pi0.5 tied parameters / alias keys 在 full state dict 下仍可能不一致

即便切到 full_state_dict=True,Pi0.5 / PaliGemma 中仍有共享 storage 的 tied parameter。

典型例子是 language-model embedding 等参数:不同 state-dict key 可能代表同一块底层 storage。

在 FSDP 展开 state dict 时,一个 alias key 可能已经是完整二维 tensor,另一个 alias key 却仍表现为局部一维 shard。这样 patch syncer 仍会对 alias key 报 shape mismatch。

因此在 /export/pgs/heqijun/RLinf/rlinf/workers/actor/embodied_fsdp_actor_worker.py 增加了:

1
_collect_state_dict_alias_groups(module)

它在 FSDP 包装前,依据共享 storage 的:

1
2
3
4
5
data pointer
storage offset
shape
stride
dtype

收集 tied-parameter alias groups。

之后在导出 rollout state dict 时:

  1. 请求完整 FSDP state dict;
  2. 对每个 alias group 找到已经 materialize 成完整 shape 的 canonical tensor;
  3. 将所有 alias key 指向这个完整 canonical tensor;
  4. 如果没有任何完整 tensor,明确抛错,而不是把错误 local shard 静默同步出去。

核心逻辑可概括为:

1
2
3
4
5
6
7
8
9
state_dict = self.get_model_state_dict(
cpu_offload=False,
full_state_dict=True,
)

for alias_group in alias_groups:
canonical_full_tensor = find_full_tensor(alias_group)
for name in alias_group:
state_dict[name] = canonical_full_tensor

这直接解决了:

1
2
expected [257152, 2048]
got [131661824]

这种 tied-key local-shard 错误。


问题 4:为 full-shard 导出行为增加回归测试

新增测试文件:

1
/export/pgs/heqijun/RLinf/tests/unit_tests/test_embodied_fsdp_actor_state_dict.py

测试覆盖两件事:

  1. 能根据 shared storage 正确识别不同名称的 alias parameter;
  2. get_rollout_state_dict() 必须请求:
1
full_state_dict=True

且 alias key 最终指向完整 shape 的 canonical tensor。

此外,实际的 4 GPU Fixed ER minimal smoke test 已验证:

1
2
3
4
5
6
initial weight sync 成功
rollout 成功
Fixed ER replay 插入与采样成功
critic update 成功
actor update 成功
alpha update 成功

所以当前 full-shard 路径不是理论修复,而是已经完成了真实多卡训练链路验证。


问题 5:Pi0.5 的 norm stats loader 曾被 LeRobot 版本耦合阻塞

这不是 FSDP 本身的问题,但它是 Pi0.5 跑通的前置障碍。

原来 OpenPI 加载 norm stats 时依赖:

1
from openpi.training import checkpoints as _checkpoints

这个 training checkpoint 模块会引入与当前环境中 LeRobot 版本绑定的训练数据加载依赖。即使 RLinf 自己的:

1
2
3
4
try:
from lerobot.datasets.lerobot_dataset import LeRobotDataset
except ModuleNotFoundError:
from lerobot.common.datasets.lerobot_dataset import LeRobotDataset

做了兼容,OpenPI 的内部 openpi.training.checkpoints 导入链仍会绕过这层兼容逻辑。

解决方法不是继续修补更深层的 LeRobot import,而是新增独立模块:

1
/export/pgs/heqijun/RLinf/rlinf/models/embodiment/openpi/norm_stats.py

只用 OpenPI 的稳定 normalization API:

1
from openpi.shared import normalize as _normalize

加载:

1
{assets_dir}/{asset_id}/norm_stats.json

然后统一替换 OpenPI、OpenPI-CFG、RLinf transforms pipeline 中对 training checkpoint loader 的调用。

这使得:

1
模型推理 / RL 权重加载

不再依赖:

1
2
OpenPI training data loader
→ 特定 LeRobot layout

同时,为 Pi0.5 RoboCasa config 显式设定:

1
use_quantile_norm=False

因为当前使用的 RoboCasa norm stats 是 mean/std 格式,而不是包含有效 q01/q99 的 quantile stats。

------------- 本文结束 感谢阅读 -------------