ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

FSDP2+HSDP+meta初始化:torchtitan中的现代分布式训练实践

FSDP2+HSDP+meta初始化:torchtitan中的现代分布式训练实践 最近很多做大规模训练的朋友开始翻 torchtitan 这个仓库。说实话第一次看到这个项目名我也有点懵因为它不像 transformers 那样拿来就用而是一个把 FSDP2、HSDP、Tensor Parallel、meta 初始化、DTensor checkpoint 全部串起来的 PyTorch 官方参考实现。很多团队现在不是直接跑 torchtitan而是把它当成一套“现代分布式训练管线”的模板把里面的并行策略和初始化流程抄进自己的训练代码里。这篇文章就围绕 FSDP2 数据并行这条主线把 fully_shard 分片、HSDP 的通信域设计、meta 初始化加载权重这三块讲清楚。适合正在把模型从 1B 往 7B、13B 甚至更大规模推或者已经被 FSDP1 的隐式递归包装折磨过、想换个更清爽的并行方案的工程同学。代码以 torchtitan 的风格为主但我会把它拆成可以直接抄进自己项目的核心片段不是纯项目源码解读。1. FSDP2 到底改了什么东西从模块包装到每参数分片1.1 FSDP1 的“递归包装”为什么让人头疼先说 FSDP1。它的核心做法是把你传入的 module 递归地切成多个 FSDP 单元每个单元负责自己那段参数的 all-gather 和 reduce-scatter。听起来很美好但用起来有几个让人抓狂的点一是你得引入 auto_wrap_policy 或者手写_auto_wrap_policy告诉它多大的模块才值得包一层二是reshard_after_forward这个参数要自己斟酌设成 True 能省显存但会频繁通信设成 False 可以少点同步但激活内存立刻飙升三是你想在某个子模块前后插入自定义逻辑时经常会发现它已经被 FSDP 包装层包住了访问内部模块还得绕。这些“设计税”不是随便吐槽是实打实影响效率和调试体验的。尤其是在你对比同一个模型在 FSDP1 和 FSDP2 下的训练曲线时会发现 FSDP1 那种“大模块包小模块再包更小模块”的方式很容易让 torch.compile 在捕获图时碰壁编译优化经常被截断在某一个包装边界上。FSDP2 的出发点很简单不再去递归包装 module而是把每个参数直接做成一个 DTensor。分片信息挂在张量本身上而不是靠外部包装器记录。你说“这个参数切到哪几个 rank 上”它就对应一个Shard(0)的 placement代码里看得一清二楚。1.2 FSDP2 的 DTensor 化设计理解 FSDP2 的关键就三个字DTensor。DTensor 是 PyTorch 的分布式张量抽象它在全局逻辑张量之上附加了一个device_mesh和一组placements。一个参数是 DTensor 之后你在逻辑上仍然可以把它当成完整形状的 Tensor 写代码但实际显存只保存在本地分片上。fully_shard做的事就是把模型的每个参数从普通 Tensor 转成带分片信息的 DTensor再挂上通信相关的钩子。这个设计和 FSDP1 相比有一个非常大的优势你可以像写单机代码一样写前向逻辑。不用关心哪个模块在哪个 FSDP 单元里也不用担心包装层级影响子模块访问因为 FSDP2 根本没有新模块类型。所有分片和通信逻辑都通过 DTensor 的 dispatcher 在张量运算时自动触发模块结构保持原样。还有一个容易被忽略的好处是FSDP2 的参数分片信息可以和其他并行维度直接组合。比如你要同时用 Tensor Parallel 和 FSDPFSDP1 得小心调整包装顺序否则 TP 的通信域会被 FSDP 包住导致 collectives 互相嵌套出问题。而 FSDP2 每个参数可以按 TP 分一维、按 FSDP 分另一维DTensor 本身支持多维 placement组合起来逻辑自然得多。这也是 torchtitan 能同时支持 TP、FSDP、HSDP 多种策略而不爆雷的重要原因。1.3 参数分片之后的前反向通信过程FSDP2 的前向流程可以这样理解假设模型当前参数是分片状态每个 rank 只持有完整参数的 1/N。前向要计算某个模块时系统先把涉及到的参数张量从本地分片扩展成完整参数这一步用的是 all-gather。算完该模块的前向之后如果reshard_after_forwardTrue系统又会立刻释放完整参数回到分片状态给下一层腾显存。反向过程类似只是顺序反过来。计算完梯度之后FSDP2 会对梯度做 reduce-scatter把不同 rank 上同一分片的梯度累加好再存回对应的参数分片。所以每个 rank 更新参数时手里拿到的梯度就是所有数据并行副本聚合后的结果更新后参数保持分片状态直到下一次前向再被 all-gather。这里要注意FSDP2 默认对每层或每个指定模块做“前向前 gather、反向后 reduce-scatter”的细粒度通信而不是整个模型一把梭。这样通信和计算可以重叠网络等待时间被掩盖在矩阵乘法后面。你不需要手动调通信频率只需决定在哪个粒度上调用fully_shard粒度越大通信越少但显存越高粒度越小越省显存但同步开销越大。torchtitan 的默认策略是在每个 Transformer Block 上调用一次fully_shard相当于让每一层自己负责参数 gather 和梯度聚合这是一个在显存和通信上都很均衡的甜点位置。2. fully_shard 分片实战自动包装与手动包装怎么选2.1 最简用例一行代码做参数量分片如果你只是想把一个模型快速改成 FSDP2 数据并行代码可以短到离谱import torch from torch.distributed.fsdp import fully_shard model MyTransformer(...) model fully_shard(model)就这一行模型参数就会被自动分片到当前device_mesh的所有 rank 上。默认情况下fully_shard会把整个模型当成一个分片单元也就是传说中“一刀切”的方式。但实际训练里我基本不会这么写。原因很简单整个模型作为一个单元分片时前向传播要先 all-gather 所有参数反向结束后再 reduce-scatter 所有梯度通信量集中且和计算重叠度差。这就好比你把一整个仓库的货一次性搬到柜台再一次性搬走中间柜台完全闲着。而如果按模块粒度一步步 gather 和释放搬运和售卖就能同时进行。所以 torchtitan 里的做法是手动在关键模块上调用fully_shard通常每个 transformer block 一次最后再对 embedding 等剩余参数单独处理。代码结构大概长这样for layer in model.layers: fully_shard(layer, meshshard_mesh, reshard_after_forwardTrue) fully_shard(model, meshshard_mesh, reshard_after_forwardTrue)第一轮先把每一层分别分片最后一轮再处理剩下没分片的参数。注意这里reshard_after_forwardTrue是默认推荐的它能保证前向计算完成后立刻释放完整参数显著降低峰值显存。2.2 自动包装与手动包装的取舍很多人会纠结FSDP2 有没有类似 FSDP1 的自动包装策略其实是有的fully_shard也可以搭配模块名做自动策略。比如用module fully_shard(module, ...)时如果你不传具体模块它会整包分片如果你想自动按命名空间分片可以用 transform 或者遍历named_modules自己决定哪些模块调fully_shard。我倾向于手动指定理由有三个可读性好。别人看你的代码能直接知道哪些模块是分片边界不用猜。显存可控。只包装特定模块就能控制完整参数的存活时间。和 activation checkpoint、offload 等策略嵌套时不容易出错。自动包装只适合快速实验或者模型结构极度规整、每一层都能统一处理时。手动包装虽然要多写几行循环但它把“哪些层会做 all-gather”这个关键信息显式暴露出来了后面调性能、查 bug 会省很多时间。这里我还想强调一个容易踩的细节fully_shard是有返回值的你必须要接收返回值。它不是原地修改 module而是返回一个新的包装后的模块。如果你写成fully_shard(model)不接收返回值后面用原始 model 跑训练会发现参数完全没有分片显存直接被打爆。这个坑我亲眼见过好几个同事踩过排查了半天才发现是没接返回值。2.3 分片布局与显存收益的量化估算FSDP2 默认对参数的第 0 维做分片。也就是说一个形状为(C_out, C_in)的线性层权重会被沿C_out方向切成 N 份每个 rank 持有(C_out/N, C_in)的本地分片。对 attention 里的 QKV 权重、FFN 里的上下投影这类大块参数这种切法非常自然。我们用 7B 模型算笔账。假设模型权重是 BF16总共参数量 70 亿全量权重大小是 14GB。如果 8 卡 FSDP 分片每卡只持有 1.75GB 权重参数。这时候加上优化器状态如果用的是 AdamW每参数需要保存 fp32 的 momentum、variance通常还有一份 fp32 master weight一共 12 字节每参数。分片后每卡优化器状态是 70 亿乘 12 字节除以 8约 10.5GB。再加上激活、梯度、通信缓冲等单卡总显存占用量可能在 20GB 到 25GB 左右。对于 80GB 的 A100/H100 来说这就比较宽裕了。如果换成 13B 模型权重分片后每卡 3.25GB优化器状态约 19.5GB整体能达到 35GB 到 45GB 的占用仍然可以放进 80GB 单卡但已经需要关注 activation checkpoint 和通信缓冲的尺寸。这里的数字我按常规配置估算实际会因为reshard_after_forward设置、激活 checkpoint 策略、batch size 不同而浮动但数量级是有参考价值的。3. HSDP混合分片并行的通信域设计3.1 HSDP 解决的核心问题跨机通信瓶颈FSDP2 本身已经把所有参数分片到所有 rank 上。但跨机场景下你会遇到一个很现实的问题机器的网卡带宽远不如机内 NVLink。假设你有 4 台机器每台 8 卡总共 32 卡。如果按普通 FSDP 把参数分到 32 卡上那么每次 all-gather 完整参数都需要跨机传输。一次 7B 模型的完整参数 gather 就是 14GB 数据其中約 3/4 是要跨机传输的这个量对 IB 网络也是压力。尤其在训练频率很高时网络会成为吞吐瓶颈。HSDP 的思路很直白分片照做但只在一台机器内部的 8 卡之间分片。机器之间采用纯数据并行每台机器保存一份完整参数副本。这样 all-gather 只在机内发生走 NVLink速度飞快跨机只需要在梯度聚合时做数据并行的 all-reduce次数和体量都小得多。3.2 通信域划分sharding group 与 replication groupHSDP 的通信域分成两层一层是分片域一组 rank 共同切分参数另一层是复制域不同组各持一份完整参数副本。在 torch.distributed 里这两个维度通常用二维device_mesh来表达。维度顺序在上面的初始化代码里很关键。一般约定(replicate, shard)也就是第一维是数据并行复制域第二维是分片域。举个例子4 机 8 卡共 32 卡mesh 形状是(4, 8)含义是有 4 个数据并行副本每个副本内部有 8 个分片 rank。在 torchtitan 中这也对应命令行里的--dp-replicate 4 --dp-shard 8。通信域设计上要注意的是如果你把 mesh 建成(8, 4)含义就变成 8 个副本、每个副本内部 4 卡分片跨机通信的比例会完全不同。这不只是数字问题而是你的集群拓扑决定哪个维度的通信应该走高速链路。我在多机调试时见过最诡异的情况就是维度过反了机内做数据并行复制、机间做分片结果训练速度反而不如单机 8 卡因为 all-gather 每次都在跨机网络上跑。3.3 用 Torchtitan 风格配置 HSDP先看最核心的 mesh 创建from torch.distributed.device_mesh import init_device_mesh # 假设共 32 卡4 机 8 卡 # 维度和集群拓扑保持一致(复制域大小, 分片域大小) mesh init_device_mesh(cuda, (4, 8)) # 取子 mesh分片域 shard_mesh mesh[shard] # 取子 mesh复制域 rep_mesh mesh[rep]创建好 mesh 后应用 FSDP2 分片时就传入分片子 meshfrom torch.distributed.fsdp import fully_shard for layer in model.layers: fully_shard(layer, meshshard_mesh, reshard_after_forwardTrue) fully_shard(model, meshshard_mesh, reshard_after_forwardTrue)有意思的是这段代码和普通单机 FSDP2 几乎一样区别只在 mesh 的形状。如果你把 mesh 设成(8, 4)同样一份代码就从“单机 8 卡 FSDP”变成了“4 机 HSDP”。这就是 DeviceMesh 抽象的价值并行策略之间的切换很多时候只是换一个 mesh 维度划分模型代码基本不动。Torchtitan 在并行策略初始化上还有一个很有用的封装它把ParallelDims和DeviceMesh绑定成一个整体通过命令行参数控制每维大小。你在代码里看到的world_mesh[dp]、world_mesh[fsdp]都是从这个二维 mesh 上切出来的。这种写法的好处是你不需要在模型实现里硬编码任何并行维度全部由外部配置驱动改并行方案时只改配置不改模型代码。3.4 HSDP 的梯度同步流程HSDP 里梯度同步的完整流程我把它理解成两步第一步复制域之间先把梯度对齐。因为每个数据并行副本各自消费不同的 micro-batch算出来的梯度不一样需要先跨副本做一次聚合第二步聚合后的梯度再沿分片域做 reduce-scatter落到每个参数分片所属的 rank 上。具体执行顺序和通信后端有关但效果是明确的跨机只传一次合并后的梯度而不是把完整参数传来传去。比如 4 机 8 卡场景跨机传输的数据量相比全量 FSDP 能减少约 3/4这对集群规模越大、模型越大收益越明显。还有一个实践细节HSDP 下reshard_after_forward可以视情况设为 False。因为分片域被限制在机内完整参数在机内 gather 出来的拷贝即使多驻留一段时间占用的也主要是本机显存不会引发跨机频繁通信。这种做法适合你想进一步减少前向阶段同步次数的场景代价是机内显存峰值上升。具体怎么取舍取决于你是网络敏感还是显存敏感我一般建议先保持 True 跑通再根据显存余量尝试关闭。4. meta 初始化与真实权重加载让 7B 模型“无中生有”4.1 为什么需要 meta 初始化所谓 meta 初始化就是让模型参数先从一个不占实际内存的metadevice 上创建。torch.device(meta)上生成的 Tensor 只记录 shape 和 dtype不真正分配显存或内存。这样做的第一个好处是创建超大模型时你的 CPU 内存不会先被一个完整模型的 fp32 副本撑爆。很多人在单机 64GB 内存上加载 13B 模型 fp32 就快爆了如果创建过程还要叠加临时缓冲区基本必挂。meta 初始化把“创建结构”和“分配数据”彻底分开CPU 内存占用从“模型大小级别”降到“结构信息级别”。第二个好处和 FSDP 强相关你想在参数真正落地显存之前就把分片关系建立好。如果在普通 CPU 上创建完整模型再去做 FSDP 分片中间过程会多出一次“完整模型权重复制到各 rank”的开销。而 meta 初始化配合 FSDP可以直接让每个 rank 只分配自己那部分分片的真实显存全程不需要出现一份完整模型在单一 rank 上。这对 7B、13B 这种规模的模型几乎是必经之路。4.2 完整流程meta 建模到显存落位我一贯推荐的安全流程是五步走import torch from torch.distributed.fsdp import fully_shard # 1. 在 meta device 上创建模型结构 with torch.device(meta): model MyTransformer(...) # 2. 做 FSDP2 分片包装参数仍然是 meta但分片关系已经确定 mesh init_device_mesh(cuda, (1, 8)) fully_shard(model, meshmesh, reshard_after_forwardTrue) # 3. 参数初始化两种常见路线 # 路线 A使用模块自带的初始化逻辑 model.to_empty(devicecuda) # 从 meta 变为 cuda但不拷贝数据 # 之后手动调用 reset_parameters或直接加载真实权重注意to_empty这一步非常关键。它会把 DTensor 本地分片从metadevice 变成真正的cuda显存但内容是未定义的你必须紧接着做初始化或加载权重。如果你没有真实 checkpoint只是从零训练路线是model.to_empty(devicecuda) for module in model.modules(): if hasattr(module, reset_parameters): module.reset_parameters()如果你的模块初始化逻辑比较复杂比如 attention 里带有nn.init之外的 buffer 计算那么reset_parameters要在to_empty之前还是之后会影响顺序。我的经验是先在 meta 阶段就调用一遍reset_parameters再用to_empty。因为很多reset_parameters会基于当前 Tensor 的 device 做初始化meta device 上部分init操作也能执行但如果某些 op 在 meta device 上不可用就只好先to_empty再初始化。两种都测试过之后我建议优先采用“先to_empty到 cuda再逐个模块 reset”的路线兼容性最好只是需要多注意不同 rank 之间参数的随机种子一致性。4.3 加载 checkpoint 的几种路径加载真实权重时有一个非常大的教训不要直接torch.load整个 state_dict 然后load_state_dict。因为在torch.load过程中完整模型的权重副本已经全部落在 CPU 内存里了这等于绕过了 meta 初始化的初衷。对大模型就算 CPU 内存暂时够用加载时也会因为逐参数 copy 到显存而慢得离谱。更合理的路径是分片加载。如果 checkpoint 本身就是按 sharded 格式保存的比如 PyTorch 的 DTensor checkpoint 或者 safetensors 分片文件那么可以直接让每个 rank 只读取自己分片对应的那一部分权重再装载到对应位置。torchtitan 就是这个思路它把 checkpoint 的保存和加载都建立在 DTensor 级别不聚合到某个 rank 上避免“先聚合成完整模型再分片”的尴尬。如果只有一份完整 checkpoint没有分片文件我在实践中会先写一个离线脚本把权重预切成 N 份每份对应一个 sharded checkpoint 目录训练启动时按 rank 读取对应分片。这样虽然多一步预处理但大规模训练启动阶段可以稳定跑完不会出现几十个 rank 同时抢占 CPU 内存导致 OOM 的情况。加载时还有一个容易忽略的assign参数。load_state_dict的assignTrue会让参数对象直接替换成加载进来的 Tensor在某些场景下能省一次 copy。但注意如果参数是 FSDP2 的 DTensorassign 替换后要确保新的 Tensor 仍然带分片信息否则下一步训练直接崩。我的经验是对 FSDP2 模型默认不要开assignTrue用普通 copy 路径更安全数据量没那么大时性能差异可以忽略。4.4 与 FSDP2 结合时的注意事项meta 初始化和 FSDP2 结合时有几个很容易踩的细节我单独列一下buffer 同步。RoPE 的 cos/sin、BatchNorm 的 running_mean/running_var 这些 buffer不会因为参数分片自动保持一致。如果模型初始化时每个 rank 用自己的reset_parameters生成 bufferrank 之间的 buffer 可能不同。需要手动torch.distributed.broadcast对齐或者在每个 rank 上基于相同 seed 生成。to_empty只作用于当前 rank 的本地分片。不用担心它会突然把全量参数都分配进来它尊重 DTensor 的分片布局。不要在fully_shard包装之后再去修改模型的nn.Parameter列表。FSDP2 在包装时已经建立了参数分片元数据如果你后面替换、删除参数新的参数可能没有被正确分片表现为某层param是普通 Tensor 而不是 DTensor。如果确实需要改结构就重新走一遍包装。checkpoint 保存时要保存model.state_dict()它会自动输出 DTensor 的全局布局信息。加载时配合 torch.distributed.checkpoint 的分片逻辑才能在多 rank 下正确恢复。5. 实战中我踩过的坑与排查技巧5.1 meta 初始化的经典连环坑第一个坑忘了to_empty直接训练。如果模型参数还停留在 meta device一进前向就会报“tensor is on meta device”之类的错误。排查时你会看到模型参数is_meta是 True。解决方法是确保从 checkpoint 加载或to_empty之后的参数device是 cuda。第二个坑to_empty之后没有初始化就训练。这种情况更隐晦因为参数已经落在显存里前向也能跑但 loss 可能是 NaN 或直接发散。原因是显存里是未定义数据不是合理初始化。遇到这种情况检查有没有在to_empty之后执行reset_parameters或权重加载。第三个坑初始化顺序和 FSDP 包装顺序不对。我一开始写过先to_empty再fully_shard结果 FSDP 包装时重新分配了参数对象之前初始化全部白费。所以标准顺序必须是meta 创建 -fully_shard包装 -to_empty或加载初始化。这个顺序一旦乱了轻则初始化无效重则分片信息丢失导致参数通信时报 shape 不匹配。5.2 fully_shard 后参数形状与访问方式fully_shard之后你打印param.shape看到的是全量形状而不是分片后的形状。这是 DTensor 的障碍屏蔽特性逻辑上你在操作一个完整参数但底层实际显存只有分片。如果你确实想看本地分片是什么形状要用param.to_local().shape这个区别在手动调试和自定义 checkpoint 保存时特别重要。我见过有人用param.shape去计算本地分片大小结果每个 rank 都以为自己在持有一份完整参数内存统计和实际占用对不上。记住一条铁律DTensor 的逻辑 shape 是全局视角to_local()才是当前 rank 的真实数据。5.3 HSDP 配置与通信域顺序问题HSDP 出错时最典型的症状是训练能跑但速度慢得离谱。这类问题 80% 出在 mesh 维度顺序上。举个真实例子。一个 4 机 8 卡环境有人把 mesh 建成了(8, 4)本意可能是“8 卡数据并行4 卡分片”但网络拓扑里 8 卡数据并行是跨机的导致每次参数同步都在跨机网络上来回跑。观察到的现象是单卡算力利用率正常但每步训练时间比预期多了好几倍。排查方法很简单在初始化完成后打印 mesh 每个维度的rank_ids_on_dim确认第一维是哪些机器、第二维是哪些机器再对照你的网络拓扑。如果发现跨机维度恰好是分片域就调整 mesh 形状的顺序。这个检查最好写进网络初始化函数里每次启动都打印一次比跑半个小时后才发现速度不对劲要省太多时间。5.4 checkpoint 加载性能问题我试过在 8 卡环境加载一个 7B 模型的普通 checkpoint直接用torch.load加load_state_dict结果等了快十分钟才加载完。原因就是每个 rank 都在完整读取 CPU state_dict然后逐个 copy 到显存既占内存又慢。改用分片 checkpoint 方案之后加载时间从十分钟级别降到了一分钟以内。具体做法离线先把 checkpoint 转成按参数分片存储的格式训练启动时用torch.distributed.checkpoint.load加载。如果你的训练代码已经用了 FSDP2加载的 state_dict 天然是 DTensor 格式torch.distributed.checkpoint会自动根据当前 rank 的 DTensor 布局读取对应分片不需要你手动去算哪个参数在哪一块。另一个建议是优先使用 safetensors 格式保存 checkpoint。它比torch.save更快且支持零拷贝内存映射加载对大型 checkpoint 的启动速度提升非常明显。torchtitan 社区里很多人在讨论 checkpoint 方案时都明确推荐 safetensors 加分片存储的组合。5.5 显存估算与实际占用偏差我见过很多人用“参数量乘 2 字节”来估算显存这是严重低估。要记住模型训练显存里至少包含四块权重分片、梯度、优化器状态、激活。FSDP2 能把权重和优化器状态压到很低但激活这一块仍然吃得很凶尤其是长序列场景。我自己的估算法是先算权重分片和优化器状态这个是确定的再根据 batch size、序列长度、层数估算激活。如果激活太大就开启 activation checkpointing以少量计算换显存。torchtitan 在启动时也允许配置 activation checkpoint 的粒度我之前一直不开后来用 7B 模型跑 4096 序列长度时发现显存爆了开完 checkpoint 之后峰值显存降了 30% 以上吞吐损失只有 10% 上下性价比非常高。所以在排查 OOM 问题时不要只盯着 FSDP2 的参数分片策略先看激活占了多大。很多时候不是分片没做好是激活这块忘了算进去。最后闲聊一句我个人的体会FSDP2、HSDP、meta 初始化这些东西单独看都是一些 API 和参数但串起来之后你才真正摸到大规模训练的骨架。torchtitan 最大的价值不是提供一个开箱即用的训练器而是把这些现代并行技术的组合方式摆在明面上让你知道每一步是为了解决哪个具体瓶颈。真到自己搭训练框架时知道“什么时候该用 HSDP、什么时候该用 meta 初始化、怎么设计 communication mesh”比背多少 API 都管用。
返回列表