PyTorch 分布式工具选型,要比较数值与通信路径
PyTorch 分布式工具选型,要比较数值与通信路径
DDP、FSDP 或其他训练方案的参数表不能代替实验。下面按数值正确性、显存和通信路径拆开验证。
1. 先建立单卡数值基线
训练问题应拆成数值正确性、数据供给、显存使用和通信行为四部分。先以小规模、固定输入验证前向和反向结果,再观察多进程路径,避免把单一监控值当成整体结论。
先保存固定 batch 的损失、梯度摘要和模型状态,再扩到多进程。工具版本、混合精度策略或切分方式变化后,需要重新比对这些基线。
2. 按最小闭环验证
每次试验都应写清框架版本、设备类型、批量形状、随机种子和启动方式。发生偏差时,优先比较中间张量与梯度,而不是直接调整并行参数。
先断言单卡与分布式首个 step 的损失和梯度在约定容差内,再比较显存峰值与通信时间。若数值已经偏离,吞吐排名没有参考意义。
3. 参考实现与图示
下面的代码用于对照 DDP 与 FSDP 的初始化和执行路径。运行时应补充设备拓扑、PyTorch 版本以及混合精度配置。
import os import torch import torch.nn as nn import torch.distributed as dist from torch.distributed.fsdp import ( FullyShardedDataParallel as FSDP, MixedPrecision, BackwardPrefetch, ) from torch.distributed.fsdp.wrap import size_based_auto_wrap_policy class SimpleTransformerLayer(nn.Module): def __init__(self, hidden_dim: int): super().__init__() self.attn = nn.Linear(hidden_dim, hidden_dim) self.mlp = nn.Sequential( nn.Linear(hidden_dim, hidden_dim * 4), nn.GELU(), nn.Linear(hidden_dim * 4, hidden_dim) ) self.norm = nn.LayerNorm(hidden_dim) def forward(self, x: torch.Tensor) -> torch.Tensor: h = self.norm(x + self.attn(x)) return h + self.mlp(h) def setup_distributed(): """初始化分布式环境,严格校验环境变量""" dist.init_process_group(backend="nccl") local_rank = int(os.environ["LOCAL_RANK"]) torch.cuda.set_device(local_rank) return local_rank def create_fsdp_model(local_rank: int, hidden_dim: int = 4096) -> FSDP: # 建立网络结构 raw_model = nn.Sequential( SimpleTransformerLayer(hidden_dim), SimpleTransformerLayer(hidden_dim) ).to(local_rank) # 1. 混合精度策略配置 (优先采用 bfloat16 防止 Underflow) bf16_ready = torch.cuda.is_bf16_supported() mp_policy = MixedPrecision( param_dtype=torch.bfloat16 if bf16_ready else torch.float16, reduce_dtype=torch.bfloat16 if bf16_ready else torch.float16, buffer_dtype=torch.bfloat16 if bf16_ready else torch.float16, ) # 2. 自动 Wrap 策略 (基于参数量阀值,避免切碎小层造成通信碎片) auto_wrap_p = size_based_auto_wrap_policy(min_num_params=int(1e6)) # 3. 构建 FSDP 实例 fsdp_model = FSDP( raw_model, auto_wrap_policy=auto_wrap_p, mixed_precision=mp_policy, backward_prefetch=BackwardPrefetch.BACKWARD_PRE, # 预取下一层参数掩盖通信 device_id=local_rank ) return fsdp_model if __name__ == "__main__": # 模拟启动命令: torchrun --nproc_per_node=2 script.py if "LOCAL_RANK" in os.environ: rank = setup_distributed() model = create_fsdp_model(rank) # 伪造输入数据进行单步测试 inputs = torch.randn(4, 128, 4096, device=rank) outputs = model(inputs) loss = outputs.sum() loss.backward() if rank == 0: print(f"训练前向与反向单步成功,Loss: {loss.item():.4f}") dist.destroy_process_group()4. 复核清单
- 单卡与多卡是否使用相同样本顺序和损失定义。
- 梯度累积与全局 batch 的换算是否一致。
- 峰值显存和通信时间是否来自同一采样窗口。
- 节点退出时能否保留可诊断日志和检查点状态。
选择依据写进实验记录
工具名称不是结论。把数值差异、显存占用和通信开销分别记录,才知道方案适不适合当前模型与集群。
