
最近一直在折腾大模型预训练的分布式切分方案从单卡跑不动、到 FSDP 状态分片、再到最后把整张训练集群像网格一样铺开踩了不少坑也积累了一套从原理到底层逻辑再到具体配置的完整思路。这篇文章想把FSDP 怎么切这件事讲透为什么非切不可、状态分片到底分的是什么、FSDP 和那些并行策略张量并行、流水线并行、上下文并行之间是什么关系以及最终组成一张训练网格时该怎么排布资源。内容主要面向正在准备或已经启动大模型预训练、希望系统理解 FSDP 及其与并行策略配合关系的工程师和算法同学有一定 PyTorch 和分布式训练基础的话读起来会更顺畅纯新手也能从显存的账本部分获得直观认知。1. 显存的账本一张卡为什么装不下大模型1.1 参数之外真正吃掉显存的隐形三件套很多人第一次接触大模型预训练时脑子里的计算很简单70B 参数每个参数 2 字节BF16那就是 140GB单卡 80GB 装不下所以需要多卡。这个理解没错但严重低估了实际显存开销。真正让显存爆掉的不是参数本身而是参数的影子——梯度、优化器状态以及训练过程中产生的激活值。拿 Adam 优化器举例。Adam 需要保存每个参数的一阶动量momentum和二阶动量variance再加上参数本身和梯度一份模型参数在实际训练中占用显存的份数大概是项目数据类型每参数占用70B 模型总占用仅状态不含激活值参数BF162 字节140 GB梯度BF162 字节140 GBAdam 一阶动量FP324 字节280 GBAdam 二阶动量FP324 字节280 GB合计-12 字节840 GB也就是说如果用 Adam 训练一个 70B 模型光静态状态就需要 840GB 显存而这还只是模型状态的部分不含激活值。对比一下8 张 H10080GB加起来才 640GB连静态状态都塞不下。这个账算明白之后你就知道分布式训练里切分是必须的不是可选项。1.2 激活值最容易被低估的动态显存模型状态下再大好歹是静态的——占多少空间是确定的。激活值就麻烦了它随 batch size、序列长度、模型层数动态变化。预训练场景里通常输入序列很长几千甚至上万 token激活值会非常惊人。举个例子一个 7B 参数的 Transformer层数 32隐藏维度 4096序列长度 4096batch size 8。单层激活值前向过程中需要保留用于反向传播的中间张量大概在几百 MB 到 1GB 级别乘上 32 层总激活值奔着 20GB 以上去了。这还是不开启激活检查点activation checkpointing的情况。所以真正训练时显存压力来自模型状态 激活值两条线。FSDP 解决的是模型状态那条线激活值那条线主要通过激活检查点每层不保存完整激活反向时重算来压缩。这也是为什么预训练框架里 FSDP 基本都会和 activation checkpointing 一起开两个工具解决的是不同的显存问题。2. FSDP 状态分片的底层逻辑切什么、怎么切、什么时候拼回来2.1 从 DDP 到 FSDP从各持全量到各持碎片理解 FSDP 之前得先知道 DDPDistributed Data Parallel怎么工作。DDP 下每张卡都有完整的一份模型参数副本前向反向各自算各自的反向结束后对梯度做 all-reduce把多卡梯度求平均然后用平均梯度去更新各自手里的完整参数。这种模式下显存占用是单卡完整模型状态 × 卡数没有任何节省。8 张卡训练 70B 模型需要 8 × 840GB这在显存层面直接不可行。DDP 的定位从来不是解决显存问题而是通过数据并行把计算量摊到多卡上。FSDPFully Sharded Data Parallel全分片数据并行的思路完全不同它把模型状态下放到所有 GPU 上做分片shard——每张卡只持有参数、梯度、优化器状态的 1/N 份N 是并行卡数。这正好对应 DeepSpeed 的 ZeRO-3 阶段。ZeRO-1 只分片优化器状态ZeRO-2 分片优化器状态 梯度ZeRO-3 则把参数也分片。PyTorch 原生 FSDP 默认就走 ZeRO-3 路线这也是它和 DDP 最本质的区别。2.2 FSDP 一次前向反向的内部动作拆解FSDP 的切分不是切完就不管了训练过程中需要不断把分片拼回来参与计算。一次完整的 forward backwardFSDP 内部做了这样几件事all-gather 参数为了计算某一层或某一个 FSDP 单元的前向结果需要该层的完整参数。FSDP 会从所有卡上收集该层参数的碎片拼成完整参数放进临时缓冲区。前向计算用完整参数计算该层输入输出之后释放不用的参数碎片。反向传播梯度算出来后在每张卡本地就是对这个层参数的局部梯度这其实就是完整梯度的一部分因为我的输入数据是我这份计算图也是由我这份数据产生的。reduce-scatter 梯度反向结束后对梯度执行 reduce-scatter把梯度平均后按分片规则分发回每张卡上。整个过程就像一群人一起做一套卷子每人的草稿纸只有一部分但通过互相传真保证了每一步计算都用上了全部的知识点。关键点在于通信次数大量增加——每过一个 FSDP 单元就要 all-gather 一次参数、reduce-scatter 一次梯度而且参数是在前向和反向各 gather 一次。2.3 FSDP 和 ZeRO 的关系辨析PyTorch FSDP 在设计上和 DeepSpeed ZeRO 高度对应但有几个实现层面的差异值得注意分片粒度ZeRO-3 是层级粒度每层做完 forward 就可以丢弃参数。FSDP 默认也按层划分但你可以通过auto_wrap_policy自定义一个 FSDP 单元包含多少层。通信原语FSDP 用 all-gather reduce-scatterZeRO-3 也是这套但 ZeRO 有stage3_gather_16bit_weights_on_model_save之类的后处理选项FSDP 则提供了summon_full_params这样的上下文方法。与 offload 的配合FSDP 支持CPUOffload可以把参数和梯度进一步卸载到 CPU 内存ZeRO-3 也有cpu_offload选项。两者思路一致但 FSDP 是 PyTorch 原生和torch.compile、activation checkpointing等生态协作更自然。从实用角度我不建议你在一个训练任务里同时混用 FSDP 和 DeepSpeed ZeRO两者会对模型状态做重复分片反而浪费通信。选一个为主即可。PyTorch 生态里现在 FSDP 用得更多因为它原生支持torch.compile且对 Transformer 结构有更好的自动包装策略。3. FSDP 代码层配置从裸写到一份可复用的训练模板3.1 核心参数sharding_strategy、auto_wrap_policy 和 forward_prefetchFSDP 用起来并不复杂但配置项多且互相影响。先把三个最关键的参数讲清楚sharding_strategy控制分片程度对应三种典型模式FULL_SHARD参数、梯度、优化器状态三者全部分片ZeRO-3显存最省通信开销最大。SHARD_GRAD_OP梯度、优化器状态分片参数不分片ZeRO-2前向不需要 all-gather通信少一些但单卡参数全量驻留。NO_SHARD就是 DDP只对梯度做 all-reduce不分片任何状态。预训练大模型基本首选FULL_SHARD因为显存是第一瓶颈。如果你的模型刚好能塞进单卡、只是想加速SHARD_GRAD_OP反而更稳。auto_wrap_policy决定 FSDP 切分的基本单位——哪些层包在一起形成一个 FSDP 单元。Transformer 里常用transformer_auto_wrap_policy按 block 为单位包装即每个 Transformer block 是一个 FSDP 单元。这样前向走到某个 block 时就只对这个 block 做 all-gather计算完立刻释放后续 block 的显存占用不会和前面的叠在一起。如果不用 wrap policy整个模型默认是一个巨大的 FSDP 单元那前向开始就得 gather 全模型参数显存优势大打折扣。forward_prefetch是个收敛速度和通信效率的开关设为True时当前层还在计算下一层的参数已经提前 all-gather 了通信和计算重叠通常能显著缩短单步训练时间。3.2 一份实测可跑的 FSDP 训练骨架下面给出一个精简但可以直接跑的训练骨架覆盖了 FSDP 初始化的关键环节。这里以 HuggingFace Transformers 的模型为例import torch import torch.distributed as dist from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp import ( FullStateDictConfig, StateDictType, MixedPrecision, CPUOffload, ShardingStrategy, ) from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy, enable_wrap from transformers import AutoModelForCausalLM, AutoTokenizer from transformers.models.llama.modeling_llama import LlamaDecoderLayer def setup(rank, world_size): dist.init_process_group(nccl, rankrank, world_sizeworld_size) def fsdp_model(model, rank, world_size): # 1) 定义混合精度策略参数和梯度用 BF16优化器状态保留 FP32 mp_policy MixedPrecision( param_dtypetorch.bfloat16, reduce_dtypetorch.bfloat16, buffer_dtypetorch.bfloat16, ) # 2) Transformer 自动包装策略按 LlamaDecoderLayer 为 FSDP 单元 wrap_policy transformer_auto_wrap_policy( transformer_layer_cls{LlamaDecoderLayer}, ) # 3) 初始化 FSDP model FSDP( model, sharding_strategyShardingStrategy.FULL_SHARD, mixed_precisionmp_policy, auto_wrap_policywrap_policy, cpu_offloadCPUOffload(offload_paramsTrue), device_idrank, ) return model几个配置要点MixedPrecision里参数用 BF16把内存占用减半但优化器状态保持 FP32 精度保证训练数值稳定。CPUOffload(offload_paramsTrue)是显存极限时的选择但这会把大量参数搬运到 CPU训练速度明显下降。我一般只在显存实在撑不住时才开而且配合forward_prefetch能缓解 CPU 和 GPU 之间的搬运延迟。device_idrank必须设置否则模型可能落在默认设备上FSDP 初始化直接报错。3.3 常见报错与参数联动wrap policy 和 state_dict 的坑FSDP 训练中我最常碰到的报错里有两个跟配置联动强相关一是wrap policy 与模型结构不匹配。比如用transformer_auto_wrap_policy时传错了层类LlamaDecoderLayer写成了LlamaModelFSDP 不会报错但包装粒度会退化成整个模型一个单元显存立刻暴涨。排查方法是打印模型的 FSDP 包装结构def print_fsdp_units(model, rank0): if rank 0: for name, module in model.named_modules(): if isinstance(module, FSDP): print(fFSDP unit: {name}, params: {sum(p.numel() for p in module.parameters())})二是保存和加载全量权重。FSDP 里直接torch.save(model.state_dict())保存的是分片规则下的状态不同卡数量下无法直接互通。正确做法是先转换为全量状态再保存def save_full_state_dict(model, rank, path): dist.barrier() if rank 0: full_policy FullStateDictConfig(offload_to_cpuTrue, rank0_onlyTrue) with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, full_policy): state_dict model.state_dict() torch.save(state_dict, path) dist.barrier()这段代码里rank0_onlyTrue和offload_to_cpuTrue表示只在 0 号卡上拼出完整权重并放到 CPU避免多卡同时保存互相覆盖。加载时反过来用SHARDED_STATE_DICT方式让 FSDP 自己把全量权重切回当前集群的分片布局。4. 从切模型状态到一张训练网格FSDP 与张量并行、流水线并行的组合方式4.1 为什么只有 FSDP 还不够FSDP 解决了显存问题也带来了通信开销。当模型规模继续增长——比如从 70B 冲到 405B、甚至千亿参数级别时FSDP 单独扛不住两个问题第一单层参数太大导致 all-gather 成本过高。假设一个超大 Transformer 的 hidden size 是 16K单层参数可能超过 5GBFSDP 在每一步前向 / 反向都要对这么一大坨参数做全集群 all-gather通信量随卡数增加而上升网络很容易成为瓶颈。第二跨节点通信的延迟和带宽浪费。FSDP 是一个全局的 all-gather / reduce-scatter它不感知物理拓扑。如果一张卡上的参数碎片分散在不同机房通信延迟会显著拉高。这个时候就需要引入张量并行Tensor Parallel, TP和流水线并行Pipeline Parallel, PP。张量并行把单层内部的矩阵乘法按行 / 列切开分布在多张卡上协同计算——像把一个矩阵乘法拆给 4 个人分别算一部分再做拼接流水线并行则把模型按层切成多个阶段每个阶段放在不同设备上数据按序流过这些设备——类似工厂流水线每个工位只处理自己的环节。这两者的共同特点是它们把通信限制在一个较小的设备组内而不是让每个算子在整张集群上做全局集合通信。4.2 一张训练网格的排布方式所谓一张网格就是把所有并行策略有机地排布到 GPU 资源上。以 64 张 GPU 的集群为例一种典型布局是张量并行组 44 张卡吃下一个超大 Transformer 层的计算流水线并行组 4模型分为 4 个阶段每阶段 16 层左右数据并行组 44 个数据并行副本彼此独立处理不同 batchFSDP 的 ZeRO-3 分片维度 数据并行组内部的 4 张卡在这个布局下全局 GPU 网格的尺寸是 TP4 × PP4 × DP4总共 64 张卡。每个 DP 副本内部再使用 FSDP 对参数、梯度、优化器状态做分片。这就是题中从 FSDP 状态分片到一张网格的完整含义FSDP 负责横向的显存切分TP/PP 负责纵向的计算并行两者组合后形成一张二维或三维的并行网格。从实际配置角度不要把 FSDP 的world_size即分片卡数和全局 GPU 数混淆。在组合训练中FSDP 只应该在数据并行组内部生效其分片维度 数据并行组的大小这里的 4而不是 64。否则同一个参数会被跨 TP 组做分片集合通信引发拓扑错乱和链路瓶颈。4.3 混合并行下的 FSDP 参数设置在网格化的训练框架比如 Megatron-LM 或 NeMo里FSDP 一般不会直接暴露为最外层 API而是通过use_fsdp或virtual_pipeline_model_parallel_size之类开关启用。但如果你在 PyTorch 原生环境里手搓这套网格FSDP 的初始化逻辑要跟着 DP 组走from torch.distributed.fsdp import ShardingStrategy from torch.distributed.fsdp.fully_sharded_data_parallel import ShardingStrategy # 假设已经通过 dist.new_group 创建了 tp_group, pp_group, dp_group # FSDP 只在 dp_group 内构造sharding_strategyFULL_SHARD 也只针对 dp 维度 sharding_strategy ShardingStrategy.FULL_SHARD model FSDP( model, process_groupdp_group, # 关键FSDP 的通信范围限定在 DP 组 sharding_strategysharding_strategy, auto_wrap_policywrap_policy, device_idlocal_rank, )注意process_group这个参数不传时 FSDP 默认使用全世界的通讯组global process group。一旦你的环境里同时存在 TP 和 PP这就会导致 FSDP 的集合通信把 TP 组内的卡也牵扯进来通信量直接翻倍不说还会干扰张量并行内部的矩阵计算。把 FSDP 限制在 DP 组内是最稳妥的做法。4.4 网格排布时的显存与调度估算排布网格时有一个粗略但经验很好用的显存估算公式单卡显存 ≈ 单卡模型状态分片 本节点内张量并行分摊的层权重 激活值 / 数据并行副本数举一个实际场景假设 70B 模型32 层hidden8192FFN28672单层参数大约 1.5 ~ 2GBBF16。64 卡训练配置为 TP4、PP4、FSDP-DP4。FSDP 分片后单卡模型状态840GB / 4 ≈ 210GB在 DP 4 的组内分片但这个估算已经包含了参数、梯度、优化器状态。如果四个 DP 副本彼此独立那全局静态状态仍然是每副本 210GB × 4 ≈ 840GB跟全局总量一致。张量并行让每一层权重分摊到 4 张卡上所以单卡持有的层权重约为整层的 1/4。激活值方面TP 能把单层激活分摊到 4 卡PP 能让激活值只出现在当前流水线阶段因此整体激活占用被显著压缩。这个估算方式不需要精确但能在训练前帮你判断显存是否够用以及瓶颈在哪个并行维度。如果算出来单卡显存逼近上限且通信主要发生在 TP 组那就该调大 TP如果显存余量充足但每步速度上不去瓶颈在 all-reduce 或 all-gather那就该调大 DP / 精细化 FSDP 的 wrap 粒度。5. 状态分片后的通信账为什么慢、怎么把它藏起来5.1 算一笔通信量FSDP 每步到底在传多少数据FSDP 的通信开销和分片卡数 N直接相关。以 70B 模型、BF16 精度为例FULL_SHARD 下每步训练前向每层 / 每个 FSDP 单元 all-gather 全量参数总量是 140GB所有参数全部 gather 一遍。反向每层 / 每个 FSDP 单元 reduce-scatter 梯度总量也是 140GB。如果不开启 backward 时的参数重 gather 优化反向还需要再 all-gather 一次参数又是 140GB。也就是说一步训练的最优通信量大约在 280GB ~ 420GB 之间具体取决于 FSDP 是否会在反向时复用前向已 gather 的参数。对比 DDP梯度 all-reduce 同样要传 140GB但 DDP 没有参数双 gather 的开销。跨节点场景比如 8 节点每节点 8 卡中FSDP 的通信是全局的意味着 42 张卡以及更多卡之间频繁做全集群集合通信。现代集群普遍用 InfiniBand400Gbps或 RoCE理论上 280GB 在理想带宽下大概需要 5.6 秒 / 步。这个数字非常吓人所以如果数量级不加以优化FSDP 在跨节点场景很容易把训练速度拖入泥潭。5.2 隐藏通信时间的有效手段理论上通信量降不下来那就得把通信时间藏进计算时间里。我实测有效的做法按优先级排第一开启 forward_prefetch。上文提过它让下一层参数在前一层还在计算时就启动 all-gather通信和计算重叠。效果在深层 Transformer 上很明显单步时间普遍能缩短 10% ~ 20%。第二合理设置梯度分桶bucket。PyTorch 的reduce_scatter是按 bucket 分批进行的对 FSDP 而言bucket 大小直接影响通信粒度。太小则通信次数过多、延迟放大太大可能导致显存峰值上升。我用 25MB ~ 50MB 的 bucket 效果较好。FSDP 的forward_prefetch和backward_prefetch可以配合 bucket 做预取from torch.distributed.fsdp import BackwardPrefetch model FSDP( model, backward_prefetchBackwardPrefetch.BACKWARD_PRE, forward_prefetchTrue, )第三梯度累积gradient accumulation。FSDP 每步都是全量 reduce-scatter梯度累积能把通信次数降为原来的 1/kk 为累积步数代价是每步更新频率降低、batch 变为原来 k 倍。在大模型预训练里这几乎是标配因为真实 batch size 本来就需要很大。注意梯度累积在 FSDP 中的实现正确性与use_orig_params、sync_module_states等参数有关跑通之前最好先做一次数值对齐测试。5.3 拓扑感知为什么 FSDP 的通信要看物理节点FSDP 全局集合通信如果穿越太多物理节点延迟会线性叠加。一个务实的优化是先让 FSDP 在单节点内完成分片再通过数据并行组将多个节点串联。也就是说FSDP 的分片维度尽量不超过单节点的 GPU 数量或者至少让一个 TP/PP 组内的卡落在同一节点里。举一个配置案例4 节点 × 8 卡 32 卡。优先把 TP8 放在单节点内8 卡过 NVLink带宽远超网卡PP1FSDP 分片维度 8理想节点内通信如果模型大到单节点的 TP8 仍然放不下再跨节点补一个 PP2 或 TP16 但使用拓扑感知的通信分组。这套尽量让内部高速通信留在节点内的路子基本是所有大规模预训练框架的共同实践。Megatron 的tensor_model_parallel_size × pipeline_model_parallel_size之所以常被安排成节点数的整数倍就是这个原因。6. 实测中的几个反直觉现象与调参取舍6.1 FSDP 不一定比 DDP 慢在某些场景下它反而更快很多人对 FSDP 的直觉是多了一堆集合通信肯定更慢。我实测过 7B 模型在 4 卡 A100400W上的对比在 batch size 相对较小、单卡能勉强容纳模型的场景下FSDP 每步确实比 DDP 慢 5% ~ 10%但当 batch size 增加到接近显存上限时DDP 因为显存不足直接 OOMFSDP 反而能稳定训练。换句话说FSDP 的核心价值不是更快而是在同样的显存下能训练更大的模型 / 更大的 batch。6.2sync_module_states和use_orig_params的取舍PyTorch 2.x 里 FSDP 增加了一些新参数其中最容易踩坑的是use_orig_params。它设True时保留了原始模型参数的视图理论上让torch.compile和部分模块如 LoRA更好兼容但代价是会额外占用一小部分参数元数据且在full state dict保存时行为有所不同。预训练这种纯全量微调场景我一般设False默认因为全量训练不需要参数裁剪或 LoRA 那类操作省事且显存最紧。sync_module_states默认False表示每个 rank 只加载自己分配到的分片权重设True则从 0 号卡广播完整参数。预训练从 checkpoint 加载时建议设True因为 checkpoint 里通常保存的是全量权重让 0 号卡持有全量再广播给其它卡比每张卡都从磁盘读一次更快也更省磁盘随机 IO。6.3 CPU offload 到底什么时候开前面说过CPUOffload(offload_paramsTrue)能显著降低单卡显存但训练速度会大幅下滑尤其是跨节点场景——因为参数要从 CPU 内存搬运到 GPU再参与 all-gather。我实测 70B 模型在 32 卡集群上开 CPU offload 后单步训练时间大概翻了一倍不止。除非模型实在塞不进 GPU显存否则不要开。更优雅的方案是先把 FSDP 的分片维度加大也就是增加卡数再考虑把激活检查点开满最后才轮到 CPU offload。当前几板斧都用尽仍 OOM再评估 CPU offload 的成本是否可接受。6.4 一个值得注意的现象FSDP 的显存峰值发生在 forward 前一刻FSDP 在 all-gather 完成的瞬间显存占用达到局部峰值——此时临时缓冲区里躺着完整的层参数加上模型其他部分的参数碎片和激活。如果你显存刚好在临界点上OOM 往往发生在这一步。解决办法有几个顺手的小手段减小bucket_cap_mb让通信更细碎但减少峰值缓冲占用。开启forward_prefetch后注意它虽然提高了速度但可能让下一层的 gather 和当前层的临时缓冲区叠加略微推高峰值。如果 OOM 与 prefetch 相关可以尝试关闭forward_prefetch再观察。调整auto_wrap_policy的包装粒度让单个 FSDP 单元不要过大。Transformer 场景我习惯一个 decoder layer 或相邻两层作为一个单元过细则通信次数增多过粗则显存峰值严重。7. 从理论到落地一张 64 卡网格的完整设计举例这一节从一个可落地的真实规模出发演示怎么把前文所有概念组织成一份可执行的训练方案。假设我们有一台 8 节点 × 8 卡共 64 卡的集群每卡 80GB如 H100要预训练一个 70B 级别的 causal LM。第一步计算纯 FSDP 需要多少卡。按前文全状态 840GB、每卡 80GB 算纯 FSDP 至少需要 11 张卡的显存840 / 80 ≈ 10.5。考虑到激活值、中间缓冲等额外开销纯 FSDP 至少需要 16 张卡以上才安全。但 16 卡通信量巨大单步时间会非常难看。第二步引入 TP 和 PP。把模型切分成 4 个流水线阶段PP4每阶段约 8 层再用 TP4 分摊单层计算和权重。这时每个流水线阶段的单卡模型状态大约是全量的 1/4TP 1/1PP 环节内不做 FSDP 分片若继续叠加 DP 副本再除以 DP 数。假设 DP4则单卡模型状态约为 840GB / 4 ≈ 210GB扣掉 80GB 的硬件上限明显不够。所以真正可行的是 TP4、PP4、FSDP-DP2单卡模型状态约 840 / 2 420GB仍然超了。因此在 80GB 卡上预训练 70B 模型纯 FSDP 方案或简单二维并行方案都不现实必须同时开启激活检查点甚至考虑 CPU offload。第三步把激活检查点算进去。开启激活检查点后激活值占用通常降到原来的 1/3 ~ 1/5取决于 checkpoint 粒度模型状态仍然是主要开销。70B 模型的实际经验是TP4、PP8或 4、DP2、FSDP FULL_SHARD 组合下单卡显存可以压在 75GB 以内。具体拆分TP4 让单层权重变为原来的 1/4。PP8 把模型切成 8 段每段只保留 1/8 的层。DP2 让模型状态在数据并行维度再减半。FSDP FULL_SHARD 在 DP2 内把剩余状态继续对半分。三层切割叠加后单卡静态显存大约能降到接近 15GB 量级840GB / 4 / 8 / 2 × 若干系数剩下的显存大头就是激活值和通信缓冲。这个方案的通信特征也健康TP 通信在节点内 NVLink、PP 通信是点对点、FSDP 通信范围只在 2 卡之间不会出现全集群风暴。实际跑的时候还需要注意global batch size 与 DP 副本的匹配。DP2 意味着一次同步更新只累积 2 份数据加上 micro-batch 和梯度累积比如 micro-batch 8、梯度累积 16则 global batch 2 × 8 × 16 256 条序列。这个值在预训练里算中等偏小必要时可以把梯度累积加到 32 ~ 64但注意收敛曲线会随 global batch 变化学习率也要按比例调整。8. 预训练切分经验小结我最常回头看的几条心得所有参数和策略最终都要落到一条原则切分的目标不是把参数塞进显存而是在通信、计算、显存三者的夹角中和出一个能稳定跑完的训练系统。我最常回头看的三条心得是显存计算优先于通信优化。先把账算清楚——模型状态多少、激活值多少、checkpoint 开不开、TP/PP/FSDP 各自切走多少——显存够才谈得上速度和扩展性。我见过不止一次团队在通信优化上费尽心思结果发现显存根本放不下方案推倒重来。FSDP 的分片维度一定不能无脑等于集群总卡数。跟 TP/PP 并存时要把 FSDP 的通信组限定在数据并行组内部否则网格拓扑错乱通信风暴只是时间问题。线上环境里的稳定性远比 benchmark 数字宝贵。多卡预训练动辄跑几周FSDP 的summon_full_params、保存 checkpoint、加载不均匀权重等机制的稳定性比单步速度重要得多。每次改配置后先拿小模型跑一个快速收敛验证再放大到全量是一套我教训很深之后形成的固定流程。最后再说一个细节很多人在切分状态时忽略了一个看似无关紧要的参数——limit_all_gathers。开启后同一批 FSDP 单元的参数 all-gather 会被合并执行而不是逐层各自发起。我实测在多层结构上开这个参数常常能把通信次数减少 20% ~ 30%代价是峰值显存略增。在显存尚有余量时值得试试。