Cosmos 3 · 训练框架 ▸ 分布式并行

并行 / 分片策略 —— 分区轴 × 叠加轴

整个框架"围绕大规模分布式训练"而建。真正落地在 ParallelDims 里的是四条轴,分两类: 分区轴dp_shard · dp_replicate)把所有 rank 瓜分掉; 叠加轴cp · cfgp)不额外占 rank, 在同一批 rank 上再切一层分组。理解这套的关键,就是分清这两类。

① 四条并行轴 ParallelDims
dp_shard分区 · FSDP2
数据并行 · 分片
fully_shard 把参数 / 梯度 / 优化器状态切成 N 片分到各 rank。-1 时自动取满 world_size。
dp_replicate分区 · HSDP
数据并行 · 复制
把整组分片模型再复制若干份,跨副本 all-reduce 梯度。铁律 dp_replicate × dp_shard = world_size
cp叠加 · overlay
上下文并行(切序列)
把长视频 / 长上下文序列切到多 rank 协同算注意力。范围 [1, 32]。两种 attention I/O 布局:sequence_sharded / replicated
cfgp叠加 · 仅推理
CFG 并行(引导)
把 classifier-free guidance 的条件 / 无条件两路前向拆到不同 rank 并行。取值 12,仅 VFM 推理。
② 两类轴怎么落到 GPU 上(示意 · world_size = 8)
🧊 分区轴 · dp_mesh (2×4) 消耗 rank
dp_replicate=2 × dp_shard=4 = 8。列 = FSDP 分片(各持 ¼ 模型),行 = HSDP 副本。
replica 0
g0¼ 参数
g1¼ 参数
g2¼ 参数
g3¼ 参数
replica 1
g4¼ 参数
g5¼ 参数
g6¼ 参数
g7¼ 参数
shard 0
shard 1
shard 2
shard 3
↕ 副本间 all-reduce 梯度 ↔ 分片间 all-gather 参数(前向按需重组)。dp_replicate 个 rank 共读同一份 shard 文件,把读取从 O(world) 降到 O(dp_shard)。
🕸️ 叠加轴 · cp_mesh / cfgp_mesh 不占 rank
同一批 8 个 rank,再按 cp / cfgp 切分组——一个 rank 同时属于某 dp 组和某 cp 组。
g0
g1
g2
g3
g4
g5
g6
g7
cp 组 · 序列½·½
cp 组
cp 组
cp 组
分区:dp_replicate × dp_shard = world_size (瓜分全部 rank)
叠加:cfgp × cp | world_size (整除即可,不新占 rank)
③ 装配顺序(parallelize_unified_mot)· 顺序有讲究 parallelize_unified_mot.py
STEP 1
apply_cp
先装上下文并行,让 CP-aware 的 dispatch_attention_fn 被后续包装器捕获
STEP 2
apply_ac
激活重计算:selective(按 op 正则 MUST_SAVE/RECOMPUTE)或 full(整块)
STEP 3
apply_compile
torch.compile 编译加速(可选)
STEP 4
apply_fsdp
fully_shard 沿 dp_mesh 分片,最后装
为何 CP 最先、FSDP 最后:CP 改写注意力 dispatch,必须先就位才能被 AC / compile / FSDP 的 wrapper 一层层包住; reasoner 的自回归前向另经 register_fsdp_forward_method 注册,确保 unshard/reshard 钩子照常触发。
④ 三种场景各用哪些轴
场景dp_sharddp_replicatecpcfgp
VLM 训练可选11
VFM 训练可选可选1
VFM 推理强制 1叠加叠加(≤2)
一句话:先分清分区(dp:瓜分 rank、切模型/复制模型)与叠加(cp/cfgp:借同一批 rank、切序列/切引导路)。 dp 解决"模型/数据放不下",cp 解决"单条序列太长",cfgp 解决"推理要跑两遍引导"。四轴可自由组合,靠上面两条不变式约束。
诚实备注:docs/code_structure.md 把并行笼统写作 “FSDP / TP / CP / PP”,但当前 ParallelDims 描述符实际只构建 dp_mesh / cp_mesh / cfgp_mesh 三张 mesh——即 FSDP/HSDP + CP + (推理)CFG 并行。 张量并行(TP)与流水并行(PP)尚未进入这个统一描述符。