Cosmos 3 · 训练框架 ▸ 分布式并行
整个框架"围绕大规模分布式训练"而建。真正落地在 ParallelDims 里的是四条轴,分两类:
分区轴(dp_shard · dp_replicate)把所有 rank 瓜分掉;
叠加轴(cp · cfgp)不额外占 rank,
在同一批 rank 上再切一层分组。理解这套的关键,就是分清这两类。
fully_shard 把参数 / 梯度 / 优化器状态切成 N 片分到各 rank。-1 时自动取满 world_size。dp_replicate × dp_shard = world_size。[1, 32]。两种 attention I/O 布局:sequence_sharded / replicated。1 或 2,仅 VFM 推理。dp_mesh (2×4) 消耗 rankdp_replicate=2 × dp_shard=4 = 8。列 = FSDP 分片(各持 ¼ 模型),行 = HSDP 副本。dp_replicate 个 rank 共读同一份 shard 文件,把读取从 O(world) 降到 O(dp_shard)。cp_mesh / cfgp_mesh 不占 rankcp / cfgp 切分组——一个 rank 同时属于某 dp 组和某 cp 组。dispatch_attention_fn 被后续包装器捕获selective(按 op 正则 MUST_SAVE/RECOMPUTE)或 full(整块)fully_shard 沿 dp_mesh 分片,最后装register_fsdp_forward_method 注册,确保 unshard/reshard 钩子照常触发。
| 场景 | dp_shard | dp_replicate | cp | cfgp |
|---|---|---|---|---|
| VLM 训练 | ✓ | 可选 | 1 | 1 |
| VFM 训练 | ✓ | 可选 | 可选 | 1 |
| VFM 推理 | ✓ | 强制 1 | 叠加 | 叠加(≤2) |
docs/code_structure.md 把并行笼统写作 “FSDP / TP / CP / PP”,但当前 ParallelDims 描述符实际只构建
dp_mesh / cp_mesh / cfgp_mesh 三张 mesh——即 FSDP/HSDP + CP + (推理)CFG 并行。
张量并行(TP)与流水并行(PP)尚未进入这个统一描述符。