日期:2025年9月22日
关键词:PyTorch、模型并行、数据并行、FSDP、DDP、显存优化、分布式训练、多卡推理


引言

随着深度学习模型规模持续增长,从亿级到千亿参数的模型已成为常态,单 GPU 显存已难以满足训练与推理需求。PyTorch 提供了多种并行策略以应对这一挑战,但不同策略在显存效率、计算性能和实现复杂度方面差异显著。

本文系统梳理 PyTorch 中的四种主流并行范式:

  1. DataParallel(DP)
  2. DistributedDataParallel(DDP)
  3. Fully Sharded Data Parallel(FSDP)
  4. 手动模型并行(Model Parallelism)

我们将从原理、实现方式、优缺点和适用场景四个维度进行深入分析,并提供实践建议与优化技巧。


一、DataParallel(DP)

原理

DataParallel 是 PyTorch 早期提供的多 GPU 支持方案,采用单进程、多线程架构。

  • 模型在主 GPU(默认为 cuda:0)上构建并复制到其他设备。
  • 输入数据沿 batch 维度拆分,分发至各 GPU 进行前向计算。
  • 所有梯度汇总至主 GPU,由主进程完成参数更新。

实现方式

python

model = nn.DataParallel(model).cuda()
output = model(input)

优点

  • 实现简单,仅需一行代码封装。
  • 适合快速原型验证。

缺点

  • 单进程瓶颈:所有计算和通信集中在主进程,易导致 CPU 成为性能瓶颈。
  • 显存浪费:主 GPU 需存储完整模型、梯度和优化器状态,显存压力最大。
  • 高通信开销:每轮迭代需广播参数、收集梯度,通信成本高。
  • 与混合精度不兼容:与 torch.cuda.amp 集成存在限制,易出错。

适用场景

  • 已不推荐用于生产环境。
  • 仅适用于小模型、单机、快速验证等临时场景。

二、DistributedDataParallel(DDP)

原理

DDP 是当前分布式训练的工业标准,采用多进程并行架构。

  • 每个 GPU 对应一个独立进程,各自维护完整的模型副本。
  • 前向计算在本地完成,反向传播后通过 AllReduce 操作同步梯度(通常使用 NCCL 后端)。
  • 参数更新在各进程本地独立完成。

实现方式

python

import torch.distributed as dist

dist.init_process_group("nccl")
model = DDP(model.to(rank), device_ids=[rank])

优点

  • 高性能:无主卡瓶颈,支持多机多卡扩展。
  • 显存利用率高:梯度同步后立即释放,减少显存占用。
  • 支持混合精度:与 torch.cuda.amp 完美集成。
  • 可扩展性强:支持数千卡集群,适用于大规模训练任务。

缺点

  • 不节省显存:每个进程仍需存储完整模型,显存需求随 GPU 数线性增长。
  • 依赖 batch size > 1:数据并行依赖 batch 拆分,对 batch size = 1 的任务无效。
  • 初始化复杂:需通过 torchrun 或 mp.spawn 启动多进程,配置较繁琐。

适用场景

  • 中等规模模型训练(如 ResNet、ViT-Base)。
  • 多 batch 输入的监督学习任务。
  • 具备分布式训练基础设施的团队。

三、Fully Sharded Data Parallel(FSDP)

原理

FSDP 由 Meta 提出,核心思想是“分片一切”,将模型参数、梯度和优化器状态均进行分片,分布到多个设备上。

  • 前向计算时,通过 AllGather 动态加载所需参数。
  • 反向传播后,通过 ReduceScatter 归并梯度并更新分片参数。
  • 支持多种分片策略(如 SHARDHYBRID_SHARDNO_SHARD)。

实现方式

python

from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp.wrap import size_based_auto_wrap_policy

fsdp_policy = size_based_auto_wrap_policy(min_num_params=1e8)
model = FSDP(model, auto_wrap_policy=fsdp_policy)

优点

  • 显著节省显存:可减少 3~4 倍显存占用,支持百亿级参数模型训练。
  • 支持超大模型:适用于 LLM、大规模视觉模型等场景。
  • 与 DDP 兼容:可构建混合并行架构(如 FSDP + DDP)。

缺点

  • 通信开销大:频繁的 AllGather 和 ReduceScatter 操作可能成为瓶颈。
  • 调试困难:模型结构被动态包装,难以跟踪中间状态。
  • 可能降低训练速度:通信与计算难以完全重叠,整体吞吐可能下降。

适用场景

  • 大语言模型(LLM)训练。
  • 显存受限的超大规模视觉模型。
  • 无法使用模型并行的黑盒或预训练模型。

四、手动模型并行(Model Parallelism)

原理

手动模型并行是指将模型的不同层或模块显式分配到不同设备上,前向传播时通过 .to(device) 移动中间激活值。

  • 无全局通信机制,仅依赖设备间张量移动(P2P)。
  • 可灵活控制模型拆分粒度。

实现方式

python

class ModelParallelTransformer(nn.Module):
    def __init__(self, num_layers, devices):
        super().__init__()
        self.devices = devices
        self.layers = nn.ModuleList([
            TransformerBlock().to(devices[i % len(devices)])
            for i in range(num_layers)
        ])

    def forward(self, x):
        for layer in self.layers:
            device = next(layer.parameters()).device
            x = x.to(device)
            x = layer(x)
        return x

优点

  • 显存分摊:每张卡仅存储部分模型,降低单卡显存压力。
  • 通信开销低:仅需设备间张量移动,无全局同步。
  • 完全可控:可根据模型结构定制拆分策略。
  • 适合推理:无需梯度同步,适合低延迟推理场景。

缺点

  • 实现复杂:需手动管理设备分配和数据移动。
  • 负载不均:若拆分不合理,可能导致某些设备成为瓶颈。
  • 不支持自动梯度:需谨慎处理 no_grad 和设备上下文。

适用场景

  • 自定义大模型推理系统。
  • 流水线式处理架构(如 encoder-decoder)。
  • 多设备异构系统(如 CPU + GPU + TPU 混合部署)。

五、四种并行策略对比

特性DataParallelDDPFSDP手动模型并行
多进程
显存节省
计算加速中等(有开销)
通信开销
实现难度简单中等中等复杂
支持推理
支持 batch_size=1

六、高级优化技巧

1. torch.compile

python

model = torch.compile(model, mode="max-autotune")
  • 在 Ampere 及以上架构 GPU 上可提升 20%~50% 推理/训练速度。
  • 支持 fullgraph=True 以减少 kernel 启动次数。

2. 混合精度训练(AMP)

python

with torch.cuda.amp.autocast():
    output = model(input)
  • 使用 float16 或 bfloat16 可减少显存占用 50%。
  • 推荐与 GradScaler 配合使用。

3. 激活值检查点(Gradient Checkpointing)

python

from torch.utils.checkpoint import checkpoint
output = checkpoint(custom_forward, x)
  • 用计算时间换取显存,适用于深层网络。
  • 可减少激活值存储 60% 以上。

4. NVLink 与 P2P 访问

  • 启用 P2P 可加速跨卡数据移动:
    python
    torch.cuda.set_device(0)
    if torch.cuda.can_device_access_peer(0, 1):
        torch.cuda.enable_peer_access(1)

七、选型建议

需求场景推荐方案
快速验证小模型DataParallel(临时使用)
标准训练任务(batch > 1)DDP + AMP + torch.compile
超大模型训练FSDP + 激活值检查点
大模型推理手动模型并行 + AMP
多机训练DDP + FSDP 混合
显存极度受限FSDP + CPU Offload

八、总结

并行方式定位推荐使用场景
DataParallel过时方案避免使用
DDP分布式训练标准多 batch 训练任务
FSDP大模型训练利器显存受限的大模型训练
手动模型并行灵活控制方案大模型推理、自定义拆分

参考资料

  1. PyTorch 官方文档:Distributed Communication
  2. FSDP 教程:https://pytorch.org/tutorials/intermediate/FSDP_tutorial.html
  3. "ZeRO: Memory Optimizations Toward Training Trillion Parameter Models" (arXiv:1910.02054)
  4. NVIDIA 多 GPU 训练指南

欢迎收藏、分享与讨论。如在实际应用中遇到并行策略相关问题,欢迎留言交流实践经验。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐