
好不容易多卡并行结果显存还是爆了先说一个我自己的真实经历。去年接了一个把中等规模模型往多卡上迁移的活儿当时想得很简单反正手头有4张A100DDP一上数据并行嘛一张卡不够就四张卡一起扛。结果代码改完跑起来optimizer.step()直接OOM。我对着nvidia-smi看了半天才反应过来——DDP只是把数据切开了模型权重、梯度和优化器状态在每张卡上都是完整的一份副本显存压力压根没解决。后来换成PyTorch FSDPFully Sharded Data Parallel全分片数据并行同样四张卡同样的模型和batch size显存占用直接就下来了。这篇文章就是围绕FSDP这套分片并行方案把我从原理到落地、再到调优踩坑的过程整理一遍给正在被单卡放不下、多卡还是放不下折磨的人一个参考。FSDP适合谁如果你在训练或微调大模型显存是硬瓶颈模型权重加到梯度再加上优化器状态后一张卡根本装不下但你又不想放弃多卡的环境重新折腾复杂的框架那FSDP是最值得先试的方案。它不改造你的模型结构只改你的DDP启动方式偶尔加几个参数就能跑起来。下面从头讲。1. 单卡训不动大模型瓶颈不在算力在显存1.1 先算一笔显存账7B模型单卡到底差多少很多人一开始不理解为什么大模型训练难难在什么地方。其实单卡放一个7B参数的模型做推理也许凑合但做训练几乎不可能。原因不是算力不够而是训练过程中的显存占用远大于模型本身的大小。拿一个7B模型来算一笔账。模型参数如果按FP32精度存每个参数占4个字节7乘以十亿再乘以四约等于28GB显存。这只是干巴巴的权重。训练还要反向传播算梯度梯度在混合精度下按FP32存也是一份接近28GB占用的张量。更重的是优化器状态以Adam为例需要维护一阶动量momentum和二阶动量variance这两个都是FP32每个参数要额外占8个字节也就是约56GB。把这三项加起来权重28GB、梯度28GB、优化器状态56GB合计112GB还只是模型状态部分没算激活值、中间变量和通信缓冲区。一张A100也就是80GB连模型状态都装不下更别提跑batch了。这就是单卡训不动大模型最本质的原因——不是你的代码不够好是显存账根本没算平。我见过不少人在这一步踩坑看到模型文件在磁盘上只有14GB觉得一张40GB的卡肯定够。14GB这个数字通常是FP16权重的大小训练时还要翻倍转成FP32再加梯度和优化器状态40GB就只剩下零头了。所以如果你是第一次做大模型训练建议先把上面这套公式自己算一遍心理有数了再选方案。1.2 DDP的局限replica模式下显存是加法不是除法大部分人接触多卡训练第一个学的就是DDPDistributedDataParallelPyTorch自带的分布式数据并行接口封装得很好写起来几乎是无感的。但DDP在显存上是做加法。DDP的思路是每张卡上都放一份完整的模型副本然后每个rank进程分配不同的数据batch去算算完之后通过AllReduce同步梯度保证每张卡上的权重更新一致。它的确能加速训练因为四张卡一起算四个不同batch吞吐上去了。但显存方面毫无变化——模型权重、梯度、优化器状态依然是每张卡各存一份四张卡等于存了四份。如果单卡放不下模型状态多卡DDP大概率还是放不下。那DDP能配合梯度累积来解决吗可以缓解batch大小的问题但模型状态这一份大头仍旧在无法靠攒几步再更新来回避。因为gradient accumulation只是把每个step的梯度攒起来权重并没有变小优化器状态也没有变小。真正想让显存降下来必须让那些每张卡都重复存一份的东西不要再重复。FSDP的切入点就是这里。它采用分片sharding思路把权重、梯度、优化器状态切碎了均匀分布到每张卡上。单卡不再持有完整副本而是持有全局参数的1/N。四卡跑7B每卡只存大约四分之一份模型状态显存压力成倍下降。2. FSDP的分片哲学拿通信换显存这笔账怎么算2.1 ZeRO-3的思路把状态切成片用完再拼回来FSDP的底层思想来自DeepSpeed提出的ZeROZero Redundancy Optimizer系列严格来说是ZeRO阶段三。ZeRO-1是把优化器状态分片ZeRO-2把优化器状态和梯度分片ZeRO-3则是把参数、梯度、优化器状态全部分片。FSDP实现的是类似ZeRO-3的完整分片。完整分片的意思很直观假设你有4张卡模型有7B参数。正常DDP下每张卡都有完整的7B参数。FSDP下每张卡只保留1/4的参数片需要用到某一层的权重时再通过通信把这一层的完整权重聚合到当前卡上用完立刻释放自己拿到的临时全量副本。这背后做了大量精细的调度但站在使用者的视角本质就是用通信换显存。如果觉得抽象可以套一个生活场景。想象一个团队四人合写一本厚书传统做法是每人买一本完整的书放在手边任何一个人要用某一章都直接翻。问题是书太厚了一个人桌上放不下。FSDP的做法是四个人每个人只保留全书四分之一的内容但任何一个人要读某一章时四人立刻把那一章各自保存的部分拼起来凑成完整章递给他读读完马上把完整版拆散收回各自的柜子里。代价是每次看书都多了拼装并拆散这个过程也就是通信开销。2.2 FSDP参数生命周期AllGather和ReduceScatter各司其职要真正理解FSDP为什么会比DDP省显存还没慢到离谱需要看它在一次训练迭代里到底对参数做了什么。FSDP把模型按层划分成多个分片单元FSDP instance每个训练step中有四个关键操作AllGather前向传播开始算某一层时先把散在各卡上的该层参数片收集起来组成完整权重放到当前卡的计算区域。这个操作之后当前卡拿到了计算所需的完整参数副本。前向计算拿完整权重做线性层、注意力计算等操作输出传给下一层。反向传播反向算到该层梯度时同样需要该层的完整参数来计算梯度。FSDP会在反向传播过程中再次AllGather完整权重。ReduceScatter算完该层梯度之后把梯度按分片规则各卡合并最终每张卡只保留自己负责的那片梯度。一次step结束后梯度分片更新优化器状态本地的参数分片也随之更新。下次step开始时再重复AllGather-计算-ReduceScatter的生命周期。这里的关键收益在于某个FSDP单元的参数在计算完之后立刻释放完整副本只在需要计算时短暂出现。用到的临时完整权重也只存在于当前卡不是所有卡都固有一份。这样模型状态就从每卡全量变成了每卡分片临时全量。2.3 分片粒度决定通信效率auto_wrap_policy的含义FSDP不是把整个模型切成一块就完了里面的门道在于怎么切。切得太粗每次通信的颗粒大临时全量副本存在的时间长显存省得不彻底切得太细通信次数暴增训练效率直线下降。PyTorch官方提供了auto_wrap_policy参数来指定分片粒度。常用的两种transformer_auto_wrap_policy按Transformer Block维度切。每个block包成一个FSDP实例block之间独立分片。这是目前大模型训练中最多人用的配置也是我认为最接近默认最优的选项。size_based_auto_wrap_policy按参数大小自动切。参数数量超过min_num_params阈值的子模块就包成FSDP实例。适合非Transformer结构的模型比如CNN或者自定义module。分片粒度的直觉是这样的。Transformer block内部有self-attention和MLP每次前向计算都需要整块参数所以把一个block整体作为一个分片单元正好匹配计算需求。切到block级别之后训练某个block时只AllGather这个block的权重其余block的分片继续待在本地显存里等待轮到它们时再聚齐。这是一个很优雅的按需取用模式。用表格对比不同粒度的取舍分片粒度通信频率临时全量显存峰值典型场景整个模型一块较低高整模型全量不推荐训练大模型按Transformer Block切中等单block权重大小大模型LLM训练/微调首选按子模块/参数大小切较高更小更受控非Transformer结构、模型层次灵活我自己在实操中的选择是Transformer结构一律用transformer_auto_wrap_policy省心且效果稳定。自定义结构先算清每个子模块的参数量再订min_num_params阈值通常从5,000,000起步调。3. 从DDP到FSDP一份可以直接改的工程配置3.1 最小改动构造器参数的几个关键位FSDP在PyTorch里的接口是torch.distributed.fsdp.FullyShardedDataParallel从PyTorch 2.0开始逐步稳定到2.1之后已经相当成熟。如果你已经有一个DDP训练脚本改成FSDP的成本其实不高。标准流程是初始化进程组把模型用FSDP包一层然后把原来model.parameters()传给optimizer剩下的循环逻辑几乎不用动。下面是个最小示例import torch import torch.distributed as dist from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy from transformers import AutoConfig, AutoModelForCausalLM def setup_distributed(): dist.init_process_group(backendnccl) torch.cuda.set_device(dist.get_rank()) def build_model_and_wrap(device): config AutoConfig.from_pretrained(your-llm-config) model AutoModelForCausalLM.from_config(config) # 关键位一auto_wrap_policy按Transformer Block切分 wrap_policy transformer_auto_wrap_policy( transformer_layer_cls{ YourTransformerBlockClass } ) # 关键位二sharding_strategy默认FULL_SHARD对应ZeRO-3 model FSDP( model, auto_wrap_policywrap_policy, sharding_strategytorch.distributed.fsdp.ShardingStrategy.FULL_SHARD, device_idtorch.cuda.current_device(), ) return model需要注意transformer_layer_cls要填你模型实际使用的Transformer Block类不同模型的类名不一样。比如模型是LlamaModelblock类大概率是LlamaDecoderLayer如果是自建的模块填你的DecoderLayer类即可。填错了一个常见后果是整个模型没被拆分FSDP等于白包显存依旧爆。3.2 sharding_strategy三种切法到底怎么选FSDP的sharding_strategy参数决定分片到什么程度直接影响显存节省幅度和通信开销。官方提供三种策略FULL_SHARD参数、梯度、优化器状态全部切分。对应ZeRO-3最省显存但通信量最大。SHARD_GRAD_OP优化器状态和梯度分片每张卡保留完整参数。对应ZeRO-2显存省得少一些但通信也少一些。NO_SHARD退化成DDP模式所有状态每卡完整副本。这个基本不做讨论除非你要对比实验。我实测的感受是绝大多数场景直接默认FULL_SHARD就行。特别是模型本身已经超出单卡能力边界的时候选SHARD_GRAD_OP治标不治本因为完整参数副本仍然卡在每卡显存里batch一大照样爆。但如果你的模型只是略微超过单卡、又不想承担额外的AllGather成本SHARD_GRAD_OP可以作为一个中间选项。3.3 混合精度、CPU offload与开启顺序FSDP支持把部分状态挪到CPU上进一步腾出GPU显存。通过cpu_offloadCPUOffload(offload_paramsTrue)设置后模型参数分片可以放在CPU内存中计算时再临时搬到GPU。和完全GPU分片相比它能大幅降低GPU显存峰值但会引入CPU-GPU之间的PCIe传输开销。我的建议是先试纯GPU分片显存还紧张再加CPU offload。CPU offload不是银弹它会把通信瓶颈从GPU互联转移一部分到CPU和PCIe带宽尤其是在数据加载、checkpoint保存的时候CPU侧很容易成为新瓶颈。混合精度方面FSDP有自己的mixed_precision参数也可以用torch.autocast在训练循环里控制。推荐在FSDP层面配置MixedPrecision而不是在外部随便套autocast因为FSDP需要知道梯度是否以FP32存储、通信时是否做精度转换这直接影响ReduceScatter的正确性。一个稳定的配置是参数用FP16通信、权重在计算时转BF16、梯度和优化器状态保持FP32from torch.distributed.fsdp import MixedPrecision fp16_policy MixedPrecision( param_dtypetorch.float16, # 分片参数存储精度 reduce_dtypetorch.float16, # 梯度通信精度 buffer_dtypetorch.float16, # buffer张量精度 )如果你的GPU支持BF16比如A100、H100、V100我更推荐把param_dtype设成torch.bfloat16。BF16的指数位和FP32一样训练稳定性比FP16好很多不会出现Loss突然变成NaN那种最常见的精度灾难。FP16经常需要动态损失缩放BF16在大多数情况下可以省掉这一步。3.4 activation checkpointing的正确姿势FSDP解决了模型状态的分片但激活值activation依然是训练时的显存大户。特别在长序列场景下每个Transformer Block前向计算都会产生大量中间激活层数一深累积起来非常可怕。减少激活内存的标准做法是activation checkpointing也叫梯度检查点。在FSDP框架下建议通过FSDP自带的API来开启而不是手工重写forward逻辑from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import ( checkpoint_wrapper, CheckpointImpl, apply_activation_checkpointing, ) def apply_fsdp_checkpointing(model): # 对每个FSDP实例包一层checkpoint wrapper non_reentrant_wrapper functools.partial( checkpoint_wrapper, offload_to_cpuFalse, checkpoint_implCheckpointImpl.NO_REENTRANT, ) apply_activation_checkpointing( model, checkpoint_wrapper_fnnon_reentrant_wrapper, auto_wrap_policyfunctools.partial( transformer_auto_wrap_policy, transformer_layer_cls{YourTransformerBlockClass}, ) )两个方案可以同时开启FSDP负责切分模型状态activation checkpointing负责削减前向激活。我见过不少只用FSDP不配checkpointing的案例效果还是有天花板。配合使用之后同一个batch大小的峰值显存能再降30%以上代价是前向多算一遍训练时长大约增加20%到30%。在显存有限的情况下这个交换是完全值得的。4. 深入优化通信、峰值显存与收敛稳定性4.1 通信瓶颈为什么小模型用FSDP反而会更慢FSDP不是免费的它的核心代价是通信量上升。DDP每个step做一次AllReduce通信的是完整梯度而FSDP的FULL_SHARD每个分片单元都要做一次AllGather和一次ReduceScatter。模型分片单元越多通信次数越多。对小模型来说显存本来够用用FSDP纯属多此一举——模型小到一张卡装得下分片带来的显存收益可忽略通信开销却实实在在增加了训练时间。我自己习惯的经验值是单卡能装下完整模型状态的两倍以上就老老实实用DDP或单卡训练没必要上FSDP。只有当模型状态超过单卡显存容量约60%、或者说你需要在更大的batch上追求吞吐时FSDP的收益才明显。如果你的集群GPU互联能力弱比如用千兆以太网而不用NVLink/InfiniBand通信开销会被进一步放大。此时可以考虑SHARD_GRAD_OP策略减少一部分AllGather通信或者在batch size和梯度累积之间做权衡。4.2 optimizer与梯度裁剪的FSDP版本常规优化器构造方式在FSDP下需要稍加注意。FSDP模型取出参数时拿到的是分片后的参数视图理论上你可以直接model.parameters()传给optimizer。但有坑FSDP的参数是支持use_orig_paramsTrue和默认模式两种。默认模式下model.parameters()会暴露FlatParameter的视图某些优化器尤其是需要按参数名或按特定shape处理参数的可能会行为异常。如果需要原生的参数列表在构造FSDP时设置use_orig_paramsTrue。这个参数在PyTorch 2.0之后存在训练LLM微调时和某些第三方库配合更丝滑。不过开启后FSDP内部通信优化略有削弱主要体现在参数被回写到原始视图需要额外处理。如果只是常规AdamW默认模式问题不大但涉及梯度裁剪就要谨慎。FSDP官方建议的梯度裁剪路径是使用torch.distributed.fsdp.fully_sharded_data_parallel.FSDP.clip_grad_norm_也可以简单地先调用model.clip_grad_norm_(max_norm)FSDP内部会统一收集各分片梯度再算全局范数。如果你用torch.nn.utils.clip_grad_norm_它对FSDP的FlatParameter视图可能算不出正确的全局梯度范数训练稳定性会出现莫名其妙的波动。这个坑藏得很深一旦发现Loss曲线乱跳先检查梯度裁剪是不是走对了FSDP的API。4.3 反向传播优化forward prefetch与通信和计算重叠FSDP默认在反向传播时才去AllGather需要的参数但同步通信的时候GPU计算单元是空闲的。为了把通信时间藏进计算时间里FSDP提供了forward_prefetchTrue参数。开启后FSDP会预取下一个FSDP单元的参数在前向还没结束时提前发起AllGather让通信与计算重叠。还有一个容易忽略的参数是limit_all_gathersTrue。这个参数用来限制同时进行的AllGather数量避免显存瞬间被多个临时全量副本挤爆。开启后FSDP会等前一个AllGather对应的参数计算完、释放掉完整副本再发起新的AllGather以时间换空间。如果遇到训练启动后峰值显存超限、但均值显存并不高的情况考虑开启这个参数。4.4 峰值显存控制的排查顺序用FSDP仍然碰到显存不足OOM大概率不是模型状态的问题而是激活值或者通信缓冲区的问题。我整理一个排查顺序遇到OOM照着走第一步查看nvidia-smi的显存曲线确认峰值出现在前向还是反向阶段。前向高通常激活值多开activation checkpointing反向高可能是AllGather出的完整参数副本累积。第二步把batch size减半简单验证是否batch相关。如果减半后能跑说明模型状态分片没问题重点是激活值和batch维度的张量。第三步检查limit_all_gathers是否开启没开就开。第四步检查梯度累积步数。FSDP下梯度累积的值和DDP类似但累积时梯度会保留在分片里如果临时张量没释放干净显存会稳步爬升。第五步关掉CPU offload再观察。有些情况下offload和FSDP通信叠加反而把临时张量撑大。这套顺序我用了很多次基本都能在两三步内定位问题。5. 踩坑实录训练中断、速度奇慢和诡异报错5.1 NCCL超时和慢路径警告FSDP跑分布式训练最常碰到的问题之一就是NCCL通信超时。原因多数不是通信带宽不够而是某个rank的计算和通信节奏不一致个别rank还在算loss另一个rank已经开始等AllGather结果了。解决思路通常有两个方向。一是把NCCL超时时间调大通过环境变量NCCL_TIMEOUT或torch.distributed.init_process_group(timeout...)设置。但这只是缓解症状。二是检查数据加载是否均衡比如有的rank数据加载太慢导致计算延迟用DataLoader的num_workers和prefetch_factor做调整。我在实践中发现FSDP对大batch和长序列特别敏感稍有负载不均就出现超时反而是数据加载的均衡性比GPU算力更重要。另一个高频告警是PyTorch打印的FSDP slow path warning意思是某个操作走了效率较低的路径。常见原因是forward里对FSDP模型输出的张量做了不必要的.detach()、.item()或强制的.cpu()这些操作会打乱FSDP的通信调度。如果看到类似Using slow path for FSDP的日志优先检查训练循环里有没有把中间张量反复搬出GPU的操作。5.2 gradient accumulation与scheduler step的配合失误FSDP下的梯度累积是个容易出细节bug的地方。常规做法是累积n个micro-batch的梯度后再统一更新参数但优化器每n步才step一次学习率scheduler必须在真正的参数更新时step而不是每个micro-batch都step。有人会把scheduler.step()放在每步都执行的位置导致学习率衰减速度比预期快n倍训练到后期Loss再也降不下去。我建议在梯度累积时把优化器和scheduler封装成这么一种节奏每个micro-batch结束只loss.backward()不step累计步数到达n之后先optimizer.step()再optimizer.zero_grad(set_to_noneTrue)最后才scheduler.step()。顺序上step和zero_grad的先后以你原来的脚本为准但scheduler一定不能在累积过程中被反复推进。这个问题在单卡上也会出现只是FSDP下累积步数多了之后影响被放大了。5.3 断点续训与state_dict的坑FSDP的state_dict和DDP完全不同。因为参数是分片存储的直接torch.save(model.state_dict())出来的内容可能已经是分片视图。如果每张卡各存一份最终得到的是N份残缺权重而不是一个完整模型。正确做法是用model.state_dict()配合torch.distributed.fsdp.FullyShardedDataParallel.state_dict_type设置。我常用的是FullStateDictConfig让rank0汇聚完整权重再保存from torch.distributed.fsdp import FullStateDictConfig, StateDictType def save_checkpoint(model, optimizer, epoch, path): dist.barrier() full_state_dict_config FullStateDictConfig(offload_to_cpuTrue, rank0_onlyTrue) with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, full_state_dict_config): model_state model.state_dict() if dist.get_rank() 0: torch.save({ epoch: epoch, model_state: model_state, optimizer_state: optimizer.state_dict(), }, path)加载时对应使用StateDictType.LOCAL_STATE_DICT或直接load_state_dict处理分片。新手最常遇到的问题是保存出来的模型单独加载做推理时缺参数就是因为没有走FULL_STATE_DICT。5.4 一个单机多卡的实测数据最后放一组我在实际项目里的对比数据环境是单机8卡A100 80GB模型是约13B参数的Decoder架构batch size每卡16序列长度2048优化器AdamW配置每卡显存峰值吞吐samples/s能否稳定训练单卡无FSDPOOM无法开始否DDP 8卡OOM无法开始否FSDP FULL_SHARD 8卡约62GB11.7是FSDP FULL_SHARD activation checkpointing约43GB8.9是FSDP CPU offload约30GB4.6是但CPU mem暴涨这个结果很有参考价值。FSDP让本来完全跑不起来的模型变得可训练加了activation checkpointing之后峰值显存又下一大截代价是吞吐下降CPU offload确实最省显存但吞吐几乎腰斩。实际项目里我最后选择了FSDP FULL_SHARD activation checkpointing的组合这也是一线大模型训练最常用的一套组合拳。6. 一些训练过程中的细节补充FSDP毕竟是一套工业级方案细节决定了能不能稳定收敛。这里补几个我自己容易忽略的零碎点。第一个是bfloat16和float16的选择。BF16精度下指数范围大基本不会像FP16那样出现梯度下溢训练稳定得多。FSDP内部梯度通信时如果用FP16还需要配合损失缩放否则小梯度直接变成0。我的建议是能上BF16就上BF16而FP16就额外加动态缩放。第二个是学习率的初始化。FSDP分片后每个参数的本地有效batch size和DDP不同吗其实是相同的因为数据并行部分没有变。但FSDP的通信调度会让每个step的实际时间分布更不均匀所以学习率调度如果用余弦衰减之类的策略要留意warmup步数不要设得太短至少几百步起步否则前期Loss容易炸。第三个是EMA或者模型平均。如果你习惯用指数移动平均EMA跟踪模型权重在FSDP下要注意EMA参数通常是额外维护一组完整权重这等于把优化器状态占用的显存又加了一份。我的处理方式是把EMA权重放在CPU上每个epoch从rank0的完整state_dict里算出EMA再搬回去不让EMA参与GPU上的计算图。第四个是企业里常见的多机多卡场景。FSDP在多机上的通信效率和单机不完全一样跨机网络带宽往往是瓶颈。给一个实用建议单机多卡优先考虑分片粒度细一点多机多卡可以适当调粗分片粒度减少跨机通信次数。具体来说多机情况下把多个Transformer Block包成一个FSDP实例虽然单次通信的数据量变大但通信频次降低跨机等待时间反而更少。7. 多机部署、显存调优和实际部署建议多机多卡部署FSDP没有想象中那么复杂但也不像单机那样改个参数就行。我看到不少人被多机场景坑过这里单独聊一下。多机场景下每台机器内部用NVLink连接机器之间走InfiniBand或者高速以太网。FSDP的AllGather和ReduceScatter首先要覆盖同一台机器内的GPU然后跨机器通信这个路径就决定了性能瓶颈通常在跨机带宽上。我的调整建议是如果跨机带宽只有25Gbps甚至10Gbps一定要配合梯度累积把通信次数降下来而不是一味增加batch size。把limit_all_gathers也打开避免跨机AllGather同时爆发。实测下来多机场景里的稳定性和吞吐很大程度取决于数据加载是否均匀和跨机通信是否避峰这两点比单纯调FSDP参数更管用。显存调优方面再补充一个细节FSDP的显存节省主要来自模型状态分片但torch.cuda.max_memory_allocated()记录的是峰值这个峰值往往来自某个时刻同时存在多个FSDP实例的临时全量副本。想要压住峰值除了limit_all_gathers还可以把forward_prefetch关掉因为预取会让下一个block的参数在当前block还没释放时提前AllGather显存峰值会更高。关掉prefetch换一些时间在显存极度紧张时是划算的。部署建议上如果你只是微调一个开源大模型优先用Hugging Face的Trainer集成FSDP它封装了绝大部分细节参数集中在一个JSON配置文件里。但如果是自研模型或者要精细控制训练流程那手动配置FSDP是绕不开的。手动配置的核心就三件事选对sharding strategy选好auto_wrap_policy把state_dict的保存类型设置正确。这三件事做好FSDP基本就跑通了。回到最初的问题多卡并行为什么显存还是爆因为你用的是DDP而DDP在显存上是加法。换成FSDP把模型状态切碎了分散到每张卡上才真正把多卡的显存统筹起来。我自己现在遇到超过单卡能力的新模型第一反应已经不是纠结要不要换方案而是直接按FSDP FULL_SHARD transformer_auto_wrap_policy activation checkpointing这套组合去搭状态字典保存和梯度裁剪提前写好跑起来省心很多。如果这篇文章能帮你少走一点弯路那就够了。