Lightning Fabric 模型并行实战:用 FSDP、Tensor Parallel 与 2D 并行训练十亿参数模型
2026/9/19 8:40:43 网站建设 项目流程

Lightning Fabric 模型并行实战:用 FSDP、Tensor Parallel 与 2D 并行训练十亿参数模型

【免费下载链接】pytorch-lightningPretrain, finetune ANY AI model of ANY size on 1 or 10,000+ GPUs with zero code changes.项目地址: https://gitcode.com/gh_mirrors/py/pytorch-lightning

导读

当模型规模达到数十亿参数时,单张 GPU 的显存(即使是最新一代 H100 的 80 GB)也不再够用,常规数据并行(DDP)会因显存溢出而失效。本文以 docs/source-fabric/advanced/model_parallel/index.rst 为核心骨架,系统讲解 Lightning Fabric 原生支持的 Fully Sharded Data Parallel(FSDP)、Tensor Parallel(TP)以及二者组合的 2D Parallel:从显存构成与并行原理出发,给出可直接运行的完整训练代码、核心参数调优与源码级实现佐证。读完本文,你将掌握在 Fabric 中仅用少量代码改动把超大规模模型分布式训练跑起来、并针对显存与吞吐做系统调优的完整方法。


为什么单卡放不下大模型:训练显存构成的五个部分

当前十亿乃至百亿参数级别的大模型,通常需要多台机器上的大量 GPU 并行训练。一个直观的参照是:即便使用 80 GB 显存的 H100 GPU(当下单卡显存最大的型号之一),也无法直接训练一个 30B 参数模型——哪怕 batch size 为 1、使用 16 位精度。原因在于训练过程中的显存消耗远不止模型参数本身,它由以下五部分构成:

  1. 模型参数(model parameters):权重本身;
  2. 前向激活值(layer activations):前向传播中每层产生的中间结果,反向传播时用于计算梯度;
  3. 反向梯度(gradients):反向传播中计算的梯度;
  4. 优化器状态(optimizer states):例如 Adam 优化器为每个参数额外维护两份指数滑动平均;
  5. 模型输出与损失(model outputs and loss)

当这五部分的总和超过单张 GPU 的显存时,常规的数据并行训练(DDP)便无法继续使用——DDP 要求模型的权重、优化器状态、激活值与梯度能整体放入单张 GPU。要突破这一限制,就需要引入模型并行(Model Parallelism)


什么是模型并行:三种主流方式的原理与权衡

模型并行并非单一技术,而是多种并行策略的统称,每种策略各有其显存收益与通信代价。

Fully Sharded Data Parallel(FSDP)

FSDP 将模型参数与优化器状态同时切分(shard)到多张 GPU 上,显著降低单卡显存占用。它的优点是显存效率极高、无需改动模型代码;缺点是在前向/反向过程中需要频繁地跨 GPU 收集(all-gather)与规约,引入通信开销与实现复杂度。当显存是首要瓶颈、且集群具备高带宽互联时,FSDP 是最佳选择。

Tensor Parallel(TP)

TP 将单个张量(如线性层的权重矩阵)拆分到多张 GPU 上,实现计算与显存的细粒度分布。它在大规模 GPU 集群上扩展性好,但每次运算后都需要同步张量切片,产生通信开销。TP 对包含大量线性层的模型(尤其是 LLM)效果最好,能在显存分布与计算效率之间取得平衡。

Pipeline Parallel(PP)

PP 把模型的层按段划分,不同 GPU 各自处理不同的层段,将 GPU 间通信压缩到流水线阶段边界。它显著降低通信量,但会引入"流水线气泡"(pipeline bubbles)——部分 GPU 处于空闲等待状态,导致效率损失。PP 适合层数深、结构顺序化(如 LLM)的模型,但需要精细管理以最小化空闲时间。

选择模型并行方式的现实原则:需要综合考虑模型架构、硬件互联与训练效率。实际工程中,混合方案(Hybrid)——组合 FSDP、TP 与 PP——往往能取长补短,是最常用的做法。

重要前提:Lightning Fabric 通过 PyTorch原生支持上述全部并行方式(FSDP、TP、2D Parallel),但流水线并行(PP)目前尚未支持(详见 docs/source-fabric/advanced/model_parallel/index.rst)。


各并行方式横向对比

原文索引页给出了 DDP、FSDP、TP 与 2D Parallel(FSDP + TP)四类方案的特性对照,归纳如下:

特性DDPFSDPTP2D Parallel(FSDP + TP)
模型代码改动无需改动无需改动需要改动需要改动
全局 batch size随 GPU 数量线性扩展随 GPU 数量线性扩展固定,不随 GPU 数扩展沿数据并行维度扩展
权重/优化器状态分布每卡一份完整副本分布到所有 GPU分布到所有 GPU分布到所有 GPU
超大单层并行计算不支持不支持(单个 FSDP 层 gather 后仍需适配单卡)支持支持
配置门槛中(需了解模型架构设置自动包装策略)高(需深入理解模型架构)
主要瓶颈显存网络(多节点时 GPU 间传输常成瓶颈,需高速网络)传输与计算不重叠TP 限机器内、FSDP 跨机器,缓解传输瓶颈

几个关键差异点值得展开:

  • FSDP 的剩余约束:单个 FSDP 层在 forward/backward 期间被收集(gather)到单卡时,其显存占用必须能被单张 GPU 容纳——这是 FSDP 使用时必须记住的硬约束;
  • TP 的 batch size 特性:由于 TP 组内每张 GPU 必须接收完全相同的输入,全局 batch size 受限于单卡显存,不会随 GPU 数量增长;
  • 2D 并行是最佳组合:把 TP 限制在机器内部、把 FSDP 用于跨机器,可同时获得 FSDP 的显存效率与 TP 的计算扩展性,并显著降低跨节点数据传输瓶颈。

实战一:用 FSDP 训练十亿参数模型

FSDP 的完整指南位于 docs/source-fabric/advanced/model_parallel/fsdp.rst,以下按实战路径逐步展开。

使用 FSDP 的前置检查清单

在切换到 FSDP 之前,请确认以下三条:

  • ✅ 拥有多张 GPU;
  • ✅ 已尝试 batch size 为 1 的常规 DDP 训练但仍显存溢出(OOM);
  • ✅ 已安装 PyTorch 2.0 或更新版本。

单行启用 FSDP 策略

启用 FSDP 只需在创建 Fabric 时设置strategy="fsdp"

fabric = L.Fabric(accelerator="cuda", devices=2, strategy="fsdp")

如需进一步配置(后续章节会用到大量可调参数),改为显式传入策略对象:

from lightning.fabric.strategies import FSDPStrategy fabric = L.Fabric(accelerator="cuda", devices=2, strategy=FSDPStrategy())

完整可运行的 FSDP 训练示例

以下代码直接来自原文档,使用 1B 参数的 Transformer 模型(32 层、隐藏维度 4096)在 2 张 GPU 上以 FSDP 训练:

import torch import torch.nn as nn import torch.nn.functional as F import lightning as L from lightning.fabric.strategies import FSDPStrategy from lightning.pytorch.demos import Transformer, WikiText2 fabric = L.Fabric(accelerator="cuda", devices=2, strategy=FSDPStrategy()) fabric.launch() fabric.seed_everything(42) with fabric.rank_zero_first(): dataset = WikiText2() # 1B parameters model = Transformer(vocab_size=dataset.vocab_size, nlayers=32, nhid=4096, ninp=1024, nhead=64) model = fabric.setup(model) optimizer = torch.optim.Adam(model.parameters(), lr=0.1) optimizer = fabric.setup_optimizers(optimizer) for i in range(10): input, target = fabric.to_device(dataset[i]) output = model(input.unsqueeze(0), target.unsqueeze(0)) loss = F.nll_loss(output, target.view(-1)) fabric.backward(loss) optimizer.step() optimizer.zero_grad() fabric.print(loss.item()) fabric.print(torch.cuda.memory_summary())

注意训练循环中的几个 Fabric 惯用法:fabric.launch()负责启动分布式进程;fabric.setup(model)在策略内部完成 FSDP 包装;fabric.backward(loss)会正确处理 FSDP 下的梯度同步;fabric.print只在 rank 0 上打印。

识别大层并用 auto_wrap_policy 指定包装粒度

FSDP 收益最大的模型是包含大量 >100M 参数大层的结构(LLM、ViT 中的线性层),这些层的参数、激活值与优化器状态可被均匀地切分到所有 GPU 上。反之,只有几千参数的小层不应被切分——通信开销会主导并拖慢训练。

通过**包装策略(wrapping policy)**指定 FSDP 应管理哪些层。Fabric 2.1+ 支持直接传入层类集合:

# 1. Define a set of layers that FSDP should manage # Here we are choosing the large encoder and decoder layers policy = {nn.TransformerEncoderLayer, nn.TransformerDecoderLayer} # 2. Pass the policy to the FSDPStrategy object strategy = FSDPStrategy(auto_wrap_policy=policy) fabric = L.Fabric(..., strategy=strategy)

对于 Lightning < 2.1 的老版本,auto_wrap_policy也接受 PyTorch 的函数式策略,例如按参数量自动包装:

from functools import partial # 1. Import a suiting wrapping policy from PyTorch from torch.distributed.fsdp.wrap import size_based_auto_wrap_policy # 2. Configure the policy policy = partial(size_based_auto_wrap_policy, min_num_params=10000) # 3. Pass it to the FSDPStrategy object strategy = FSDPStrategy(auto_wrap_policy=policy)

PyTorch 在torch.distributed.fsdp.wrap下提供了多种这类函数式策略可供选用。从源码看,Fabric 的 FSDPStrategy 会把auto_wrap_policy处理进底层FullyShardedDataParallel的构造参数(_auto_wrap_policy_kwargs),并默认设置use_orig_params=True以支持优化器与模型的联合 setup、多参数组以及torch.compile()

验证 FSDP 是否生效:对比上文示例中torch.cuda.memory_summary()打印的峰值显存与常规 DDP 训练的差异。官方在 A100 40GB GPU、Lightning 2.1、PyTorch 2.1 环境下的基准数据如下(注意这是在特定硬件/版本下的参考值,实际数值随模型与集群变化):

指标DDPFSDP
显存(MB)26,95311,578
单次迭代时间(sec)0.260.36

可以看到 FSDP 把显存占用降低了约 57%,代价是迭代时间略有增加(0.26s → 0.36s)——这正是显存与速度权衡的直接体现。

用 init_module 加速模型初始化

PyTorch 的标准做法是先把全部参数放到 CPU 内存、再在第二步搬到 GPU,模型越大这两步耗时越长。Fabric 提供fabric.init_module()上下文管理器,可以直接在 GPU 上创建模型并降低初始化显存峰值:

# Slow: Places the model on CPU first model = Transformer(vocab_size=dataset.vocab_size) # Fast: Creates the model on the GPU directly with fabric.init_module(): model = Transformer(vocab_size=dataset.vocab_size) # Recommended for FSDP: with fabric.init_module(empty_init=True): model = Transformer(vocab_size=dataset.vocab_size)

对于 FSDP,官方推荐设置empty_init=True:它会创建不分配任何显存的"假参数"(meta device 上的参数),真正的初始化被推迟到fabric.setup()——此时 FSDP 已完成分片并重新创建真实参数,从而可以初始化更大的模型。从源码看,FSDPStrategy.module_init_context 在empty_init=True时会将模块创建置于 meta device 上下文,并在setup_module阶段物化。

更多empty_init=True的使用场景见 模型初始化指南。

切分策略:用显存换速度

默认情况下,FSDP 会自动切分 1) 模型权重、2) 反向传播中的梯度、3) 优化器状态。通过sharding_strategy可以调整切分范围以权衡显存与速度:

strategy = FSDPStrategy( # Default: Shard weights, gradients, optimizer state (1 + 2 + 3) sharding_strategy="FULL_SHARD", # Shard gradients, optimizer state (2 + 3) sharding_strategy="SHARD_GRAD_OP", # Full-shard within a machine, replicate across machines sharding_strategy="HYBRID_SHARD", # Don't shard anything (similar to DDP) sharding_strategy="NO_SHARD", ) fabric = L.Fabric(..., strategy=strategy)

选择切分策略的推荐顺序(Recipe)

  1. 先用默认设置(FULL_SHARD)。它最省显存但最慢;
  2. 尝试SHARD_GRAD_OP。如果显存不足,回退到默认的FULL_SHARD;否则通常会看到迭代速度提升;
  3. 跨多机训练时,尝试HYBRID_SHARD(机器内全切分、机器间复制)。

官方在 A100 40GB、Lightning 2.1、PyTorch 2.1 下的基准对比:

指标DDPNO_SHARDSHARD_GRAD_OPFULL_SHARD
显存(MB)26,95323,18111,81511,578
迭代时间(sec)0.260.300.310.36

数据清晰展示了"切分越多、越省显存、越慢"的权衡曲线:从NO_SHARDFULL_SHARD,显存从 23GB 降到 11.5GB,迭代时间从 0.30s 升到 0.36s。

用显存换速度:激活检查点与 CPU 卸载

训练 10B+ 参数模型或需要极大 batch size 时,可考虑以速度为代价换取更多显存,两条途径分别是激活检查点与 CPU 卸载。

激活检查点(Activation checkpointing):激活值(前向中各层的中间输出)在反向传播计算梯度时需要用到,默认会贯穿整个前向被存储。启用激活检查点后,可以选择丢弃部分层的激活值、在反向需要时动态重算。这会略微降低训练速度,但显著降低显存占用,腾出的显存可用于增大模型容量或 batch size:

strategy = FSDPStrategy( # Enable activation checkpointing on these layers activation_checkpointing_policy={ nn.TransformerEncoderLayer, nn.TransformerDecoderLayer, }, ) fabric = L.Fabric(..., strategy=strategy)

典型实践是把activation_checkpointing_policy设为与auto_wrap_policy相同(通常就是你的 transformer block,包含 attention 与 feed-forward)。

CPU 卸载(CPU offload):最激进的显存节省手段是把参数卸载到 CPU 内存:

# Set `cpu_offload=True` strategy = FSDPStrategy(..., cpu_offload=True) fabric = L.Fabric(..., strategy=strategy)

代价是训练速度大幅下降——每个前向都需要在 CPU 与 GPU 之间传输参数。仅当 CPU 内存充足、且其他扩展手段无法提供足够显存节省时才应使用。官方基准(A100 40GB、Lightning 2.1、PyTorch 2.1)显示 CPU 卸载带来约 4 倍显存节省,但迭代时间增加约 10 倍:

指标DDPFSDPFSDP + CPU offload
显存(MB)26,95311,5782,825
迭代时间(sec)0.260.363.24

保存与加载大模型检查点

大模型训练成本高昂,务必将检查点逻辑纳入训练循环。Fabric 提供了高效保存大型检查点的方法——直接把模型/优化器等对象放进 state 字典,而不是手动序列化 state dict:

# 1. Define model, optimizer, and other training loop state state = {"model": model, "optimizer": optimizer, "iter": iteration} # DON'T do this (inefficient): # state = {"model": model.state_dict(), "optimizer": optimizer.state_dict(), ...} # 2. Save using Fabric's method fabric.save("path/to/checkpoint/file", state) # DON'T do this (inefficient): # torch.save("path/to/checkpoint/file", state)

为降低显存峰值并加快落盘,默认情况下每个进程/GPU 会把各自的分片保存到指定路径的文件夹中,形成如下结构:

path/to/checkpoint/file ├── .metadata ├── __0_0.distcp ├── __1_0.distcp ... └── meta.pt

这种"分片检查点"(sharded checkpoint)格式在 Fabric 中保存与加载效率最高。若希望得到单一合并文件,可通过state_dict_type切换:

# Default: Save individual files with state from each process strategy = FSDPStrategy(state_dict_type="sharded") # Save a single, consolidated checkpoint file strategy = FSDPStrategy(state_dict_type="full")

如何选择检查点格式

  • state_dict_type="sharded":适用于预训练超大规模模型,保存快、占用显存少,但可移植性差,需要额外步骤把分片检查点转换为常规检查点,见 分布式检查点转换指南;
  • state_dict_type="full":适用于预训练中小规模模型(<10B 参数)、微调以及需要可移植性的场景。

加载检查点同样简单,且 Fabric 会自动识别路径中是full还是sharded格式:

# 1. Define model, optimizer, and other training loop state state = {"model": model, "optimizer": optimizer, "iter": iteration} # 2. Load using Fabric's method fabric.load("path/to/checkpoint/file", state) # DON'T do this (inefficient): # model.load_state_dict(torch.load("path/to/checkpoint/file"))

需要注意:full格式的检查点可以被所有策略加载,而sharded格式只能被 FSDP 加载。更多特性见 检查点指南。

进阶性能优化技巧

关闭优化器的 foreach:PyTorch 常用优化器的foreach=True选项会加速参数与状态更新,但可能带来轻微显存峰值,模型越大越明显。若出现不希望的显存模式,可关闭:

optimizer = torch.optim.AdamW(model.parameters(), foreach=False)

限制 all-gather(limit_all_gathers):当训练接近单卡显存上限时,可能出现 CUDA malloc retries(GPU 显存即将耗尽、崩溃前尝试释放缓存内存的现象),频繁发生时对速度影响显著。常规做法是略微减小 batch size,而 FSDP 额外提供了limit_all_gathers旋钮:

strategy = FSDPStrategy( # Default: The CPU will schedule the transfer of weights between GPUs # at will, sometimes too aggressively limit_all_gathers=False, # Enable this if you are close to the max. GPU memory usage limit_all_gathers=True, ) fabric = L.Fabric(..., strategy=strategy)

可以在torch.cuda.memory_summary()的输出或 PyTorch profiler 中监控 CUDA malloc retries 的次数。


实战二:Tensor Parallel 切分线性层

Tensor Parallel 的完整指南位于 docs/source-fabric/advanced/model_parallel/tp.rst。它是一种把层分布到多设备上训练大模型的技术,通过减少设备间通信改善内存管理与效率;但对小模型而言,通信开销可能超过收益,最适用于包含超大层的模型。

注意:Tensor Parallelism 在 Lightning Fabric 与 PyTorch 中均为实验性特性,API 未来可能变更。

原理:线性层的两种切分方式

张量并行的核心思想是把一个线性层的计算拆分到多张 GPU 上,每张 GPU 只需持有权重矩阵的一部分。线性层可按两种方式切分:

列并行(Column-wise Parallel):权重矩阵沿列维度均匀切分。每张 GPU 收到相同输入,用自己的权重子矩阵做常规矩阵乘法,最后把各 GPU 输出拼接(concatenate)成完整输出。

行并行(Row-wise Parallel):权重矩阵沿行维度均匀切分,输入也沿内维(对应权重矩阵行数变少)同步切分。每张 GPU 用各自的权重子矩阵与输入子矩阵做常规矩阵乘法,最后对各 GPU 输出做逐元素求和(all-reduce)得到最终输出。

列并行与行并行组合:当多个线性层顺序出现(如 MLP 或 Transformer)时,组合两种方式效果最佳——列并行层的输出不必拼接,直接喂给行并行层,从而避免 GPU 间昂贵的数据传输。层间的激活函数因为是逐元素运算,无需额外通信即可应用。

用 ModelParallelStrategy 对模型应用 TP

将 TP 应用于模型前,需要充分理解模型架构,决定在哪些层使用哪种并行方式。以一个三线性层 MLP 玩具模型为例:

import torch.nn as nn import torch.nn.functional as F class FeedForward(nn.Module): def __init__(self, dim, hidden_dim): super().__init__() self.w1 = nn.Linear(dim, hidden_dim, bias=False) self.w2 = nn.Linear(hidden_dim, dim, bias=False) self.w3 = nn.Linear(dim, hidden_dim, bias=False) def forward(self, x): return self.w2(F.silu(self.w1(x)) * self.w3(x))

该模型有三个线性层:w1w3的输出随后被逐元素相乘,结果喂给w2。因此w1w3适合列并行(其输出可轻松与w2的行并行组合)。

在 Fabric 中,把并行逻辑写成独立函数(保持模型源码整洁、可维护):

from torch.distributed.tensor.parallel import ColwiseParallel, RowwiseParallel from torch.distributed.tensor.parallel import parallelize_module def parallelize_feedforward(model, device_mesh): # Lightning will set up a device mesh for you tp_mesh = device_mesh["tensor_parallel"] # Use PyTorch's distributed tensor APIs to parallelize the model plan = { "w1": ColwiseParallel(), "w2": RowwiseParallel(), "w3": ColwiseParallel(), } parallelize_module(model, tp_mesh, plan) return model

然后配置 Fabric 的ModelParallelStrategy

import lightning as L from lightning.fabric.strategies import ModelParallelStrategy # 1. Pass the parallelization function to the strategy strategy = ModelParallelStrategy(parallelize_fn=parallelize_feedforward) # 2. Configure devices and set the strategy in Fabric fabric = L.Fabric(accelerator="cuda", devices=2, strategy=strategy) fabric.launch()

策略把自定义并行函数作为输入,训练代码其他部分无需改动——当后续调用fabric.setup(model)时,Fabric 会自动把parallelize_feedforward应用到模型上。这一点可以从源码得到印证:ModelParallelStrategy.setup_module 中直接调用self._parallelize_fn(module, self.device_mesh)并校验返回值必须是nn.Module实例,随后执行_materialize_distributed_module完成物化。

完整的 TP 训练示例(需至少 2 张 GPU):

import torch import torch.nn as nn import torch.nn.functional as F from torch.distributed.tensor.parallel import ColwiseParallel, RowwiseParallel from torch.distributed.tensor.parallel import parallelize_module import lightning as L from lightning.pytorch.demos.boring_classes import RandomDataset from lightning.fabric.strategies import ModelParallelStrategy class FeedForward(nn.Module): def __init__(self, dim, hidden_dim): super().__init__() self.w1 = nn.Linear(dim, hidden_dim, bias=False) self.w2 = nn.Linear(hidden_dim, dim, bias=False) self.w3 = nn.Linear(dim, hidden_dim, bias=False) def forward(self, x): return self.w2(F.silu(self.w1(x)) * self.w3(x)) def parallelize_feedforward(model, device_mesh): # Lightning will set up a device mesh for you tp_mesh = device_mesh["tensor_parallel"] # Use PyTorch's distributed tensor APIs to parallelize the model plan = { "w1": ColwiseParallel(), "w2": RowwiseParallel(), "w3": ColwiseParallel(), } parallelize_module(model, tp_mesh, plan) return model strategy = ModelParallelStrategy(parallelize_fn=parallelize_feedforward) fabric = L.Fabric(accelerator="cuda", devices=2, strategy=strategy) fabric.launch() # Initialize the model model = FeedForward(8192, 8192) model = fabric.setup(model) # Define the optimizer optimizer = torch.optim.AdamW(model.parameters(), lr=3e-3) optimizer = fabric.setup_optimizers(optimizer) # Define dataset/dataloader dataset = RandomDataset(8192, 64) dataloader = torch.utils.data.DataLoader(dataset, batch_size=8) dataloader = fabric.setup_dataloaders(dataloader) # Simplified training loop for i, batch in enumerate(dataloader): output = model(batch) loss = output.sum() fabric.backward(loss) optimizer.step() optimizer.zero_grad() fabric.print(f"Iteration {i} complete") fabric.print(f"Peak memory usage: {torch.cuda.max_memory_allocated() / 1e9:.02f} GB")

官方基准显示,随着 GPU 数量翻倍,单卡峰值显存近似减半:

配置1 GPU(无 TP)2 GPUs4 GPUs8 GPUs
每卡显存4.04 GB2.03 GB1.02 GB0.60 GB

TP 的数据加载注意事项

在张量并行的模型中,参与同一 TP 组的每张 GPU 必须收到完全相同的输入,否则训练无法收敛。因此在数据集/dataloader 中做 shuffle、或应用随机变换/数据增强时,必须正确设置随机种子。

这也意味着 TP 下的全局 batch size 受限于单卡显存。要扩大 batch size 并加速训练,需要把 TP 与数据并行(尤其是 FSDP)组合使用——这正是下一节的 2D 并行。


实战三:2D 并行(TP + FSDP)扩展到数百张 GPU

2D Parallel 的完整指南位于 docs/source-fabric/advanced/model_parallel/tp_fsdp.rst。它组合 TP 与 FSDP,兼顾 FSDP 的显存效率与 TP 的计算扩展性,通过平衡各自取舍、优化显存并最小化通信开销,实现在大规模 GPU 集群上训练超大模型。本教程以 Tensor Parallelism 文档 与 FSDP 基础知识为前提。

注意:2D Parallelism 在 Lightning Fabric 与 PyTorch 中均为实验性特性,API 未来可能变更。

启用 2D 并行:device mesh 与 parallelize 函数

沿用上一节的 FeedForward 模型。并行函数除了做 TP 切分,还沿数据并行维度用 FSDP2 的fully_shard切分参数:

import torch.nn as nn import torch.nn.functional as F from torch.distributed.tensor.parallel import ColwiseParallel, RowwiseParallel from torch.distributed.tensor.parallel import parallelize_module from torch.distributed._composable.fsdp.fully_shard import fully_shard def parallelize_feedforward(model, device_mesh): # Lightning will set up a device mesh for you # Here, it is 2-dimensional tp_mesh = device_mesh["tensor_parallel"] dp_mesh = device_mesh["data_parallel"] if tp_mesh.size() > 1: # Use PyTorch's distributed tensor APIs to parallelize the model plan = { "w1": ColwiseParallel(), "w2": RowwiseParallel(), "w3": ColwiseParallel(), } parallelize_module(model, tp_mesh, plan) if dp_mesh.size() > 1: # Use PyTorch's FSDP2 APIs to parallelize the model fully_shard(model.w1, mesh=dp_mesh) fully_shard(model.w2, mesh=dp_mesh) fully_shard(model.w3, mesh=dp_mesh) fully_shard(model, mesh=dp_mesh) return model

函数必须把model作为第一个参数、DeviceMesh作为第二个参数。随后把函数传给ModelParallelStrategy,并指定数据并行与张量并行的规模:

import lightning as L from lightning.fabric.strategies import ModelParallelStrategy strategy = ModelParallelStrategy( parallelize_fn=parallelize_feedforward, # Define the size of the 2D parallelism # Set these to "auto" (default) to apply TP intra-node and FSDP inter-node data_parallel_size=2, tensor_parallel_size=2, ) fabric = L.Fabric(accelerator="cuda", devices=4, strategy=strategy) fabric.launch()

device mesh 的划分逻辑:在上述 4 卡示例中,Fabric 创建的 device mesh 会把 GPU 0-1 与 GPU 2-3 各分为一组(因为data_parallel_size=2,每组 2 张 GPU 对应tensor_parallel_size=2)。随后调用fabric.setup(model)时,每个用fully_shard包装的层会被切成两份分片(对应 GPU 0-1 组与 GPU 2-3 组),再在每组内部应用 TP,把分片后的张量进一步切到组内各 GPU 上。

从源码可以验证 mesh 的构建规则:ModelParallelStrategy.setup_environment 中,"auto"会被解析为data_parallel_size = 节点数tensor_parallel_size = 每节点 GPU 数_setup_device_mesh则强制校验data_parallel_size * tensor_parallel_size == world_size,否则抛出RuntimeError,然后通过init_device_mesh(..., mesh_dim_names=("data_parallel", "tensor_parallel"))构建二维 mesh。

完整训练示例(需至少 4 张 GPU):

import torch import torch.nn as nn import torch.nn.functional as F from torch.distributed.tensor.parallel import ColwiseParallel, RowwiseParallel from torch.distributed.tensor.parallel import parallelize_module from torch.distributed._composable.fsdp.fully_shard import fully_shard import lightning as L from lightning.pytorch.demos.boring_classes import RandomDataset from lightning.fabric.strategies import ModelParallelStrategy class FeedForward(nn.Module): def __init__(self, dim, hidden_dim): super().__init__() self.w1 = nn.Linear(dim, hidden_dim, bias=False) self.w2 = nn.Linear(hidden_dim, dim, bias=False) self.w3 = nn.Linear(dim, hidden_dim, bias=False) def forward(self, x): return self.w2(F.silu(self.w1(x)) * self.w3(x)) def parallelize_feedforward(model, device_mesh): # Lightning will set up a device mesh for you # Here, it is 2-dimensional tp_mesh = device_mesh["tensor_parallel"] dp_mesh = device_mesh["data_parallel"] if tp_mesh.size() > 1: # Use PyTorch's distributed tensor APIs to parallelize the model plan = { "w1": ColwiseParallel(), "w2": RowwiseParallel(), "w3": ColwiseParallel(), } parallelize_module(model, tp_mesh, plan) if dp_mesh.size() > 1: # Use PyTorch's FSDP2 APIs to parallelize the model fully_shard(model.w1, mesh=dp_mesh) fully_shard(model.w2, mesh=dp_mesh) fully_shard(model.w3, mesh=dp_mesh) fully_shard(model, mesh=dp_mesh) return model strategy = ModelParallelStrategy( parallelize_fn=parallelize_feedforward, data_parallel_size=2, tensor_parallel_size=2, ) fabric = L.Fabric(accelerator="cuda", devices=4, strategy=strategy) fabric.launch() # Initialize the model model = FeedForward(8192, 8192) model = fabric.setup(model) # Define the optimizer optimizer = torch.optim.AdamW(model.parameters(), lr=3e-3) optimizer = fabric.setup_optimizers(optimizer) # Define dataset/dataloader dataset = RandomDataset(8192, 128) dataloader = torch.utils.data.DataLoader(dataset, batch_size=8) dataloader = fabric.setup_dataloaders(dataloader) # Simplified training loop for i, batch in enumerate(dataloader): output = model(batch) loss = output.sum() fabric.backward(loss) optimizer.step() optimizer.zero_grad() fabric.print(f"Iteration {i} complete") fabric.print(f"Peak memory usage: {torch.cuda.max_memory_allocated() / 1e9:.02f} GB")

2D 并行的典型使用场景:TP 限机器内、FSDP 跨机器

上述玩具示例把并行配置在同一台机器的多张 GPU 上,但 2D 并行真正的主战场是多节点训练。核心工程判断是:

  • TP 应限制在机器内部:张量并行的集体通信是阻塞式的,需要极快的 GPU 数据传输才能保持高吞吐;
  • FSDP 适合跨机器:FSDP 天然可以把 GPU 数据传输与计算重叠(例如预取层),通信效率高。

因此"机器内用 TP、机器间用 FSDP"通常是同时最小化延迟与网络带宽占用的最佳策略,能扩展到远超单独使用 FSDP 的模型规模。实现上只需把两个维度都设为"auto"(默认值):

from lightning.fabric.strategies import ModelParallelStrategy strategy = ModelParallelStrategy( # Default is "auto" # Applies TP intra-node and DP inter-node data_parallel_size="auto", tensor_parallel_size="auto", )

2D 并行的数据加载语义

在 2D 并行下,数据加载的语义需要精确理解:参与同一 TP 组的 GPU 必须收到相同输入,而跨数据并行维度的输入必须不同。也就是说,如果 TP 在节点内、FSDP 跨节点,那么每个节点收到不同 batch,而节点内每张 GPU 收到同一份 batch。

使用 PyTorch dataloader 并经fabric.setup_dataloaders设置后,Fabric 会通过配置分布式 sampler 自动处理这一语义。从源码看,ModelParallelStrategy.distributed_sampler_kwargs 返回{"num_replicas": data_parallel_mesh.size(), "rank": data_parallel_mesh.get_local_rank()}——即采样器只沿数据并行维度切分数据集,从而保证 TP 组内各 GPU 的 batch 一致。但请注意:数据集中的 shuffle 或随机增强仍须自行固定随机种子

import lightning as L fabric = L.Fabric(...) # Define dataset/dataloader # If there is randomness/augmentation in the dataset, fix the seed dataset = MyDataset(seed=42) dataloader = DataLoader(dataset, batch_size=8, shuffle=True) # Fabric configures the sampler automatically for you such that # all batches in a tensor-parallel group are identical, # while still sharding the dataset across the contenteditable="false">【免费下载链接】pytorch-lightningPretrain, finetune ANY AI model of ANY size on 1 or 10,000+ GPUs with zero code changes.项目地址: https://gitcode.com/gh_mirrors/py/pytorch-lightning

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询