Lightning Fabric 模型并行实战:用 FSDP、Tensor Parallel 与 2D 并行训练十亿参数模型
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 ParallelFSDP、Tensor ParallelTP以及二者组合的 2D Parallel从显存构成与并行原理出发给出可直接运行的完整训练代码、核心参数调优与源码级实现佐证。读完本文你将掌握在 Fabric 中仅用少量代码改动把超大规模模型分布式训练跑起来、并针对显存与吞吐做系统调优的完整方法。为什么单卡放不下大模型训练显存构成的五个部分当前十亿乃至百亿参数级别的大模型通常需要多台机器上的大量 GPU 并行训练。一个直观的参照是即便使用 80 GB 显存的 H100 GPU当下单卡显存最大的型号之一也无法直接训练一个 30B 参数模型——哪怕 batch size 为 1、使用 16 位精度。原因在于训练过程中的显存消耗远不止模型参数本身它由以下五部分构成模型参数model parameters权重本身前向激活值layer activations前向传播中每层产生的中间结果反向传播时用于计算梯度反向梯度gradients反向传播中计算的梯度优化器状态optimizer states例如 Adam 优化器为每个参数额外维护两份指数滑动平均模型输出与损失model outputs and loss。当这五部分的总和超过单张 GPU 的显存时常规的数据并行训练DDP便无法继续使用——DDP 要求模型的权重、优化器状态、激活值与梯度能整体放入单张 GPU。要突破这一限制就需要引入模型并行Model Parallelism。什么是模型并行三种主流方式的原理与权衡模型并行并非单一技术而是多种并行策略的统称每种策略各有其显存收益与通信代价。Fully Sharded Data ParallelFSDPFSDP 将模型参数与优化器状态同时切分shard到多张 GPU 上显著降低单卡显存占用。它的优点是显存效率极高、无需改动模型代码缺点是在前向/反向过程中需要频繁地跨 GPU 收集all-gather与规约引入通信开销与实现复杂度。当显存是首要瓶颈、且集群具备高带宽互联时FSDP 是最佳选择。Tensor ParallelTPTP 将单个张量如线性层的权重矩阵拆分到多张 GPU 上实现计算与显存的细粒度分布。它在大规模 GPU 集群上扩展性好但每次运算后都需要同步张量切片产生通信开销。TP 对包含大量线性层的模型尤其是 LLM效果最好能在显存分布与计算效率之间取得平衡。Pipeline ParallelPPPP 把模型的层按段划分不同 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 ParallelFSDP TP四类方案的特性对照归纳如下特性DDPFSDPTP2D ParallelFSDP 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 时设置strategyfsdpfabric L.Fabric(acceleratorcuda, devices2, strategyfsdp)如需进一步配置后续章节会用到大量可调参数改为显式传入策略对象from lightning.fabric.strategies import FSDPStrategy fabric L.Fabric(acceleratorcuda, devices2, strategyFSDPStrategy())完整可运行的 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(acceleratorcuda, devices2, strategyFSDPStrategy()) fabric.launch() fabric.seed_everything(42) with fabric.rank_zero_first(): dataset WikiText2() # 1B parameters model Transformer(vocab_sizedataset.vocab_size, nlayers32, nhid4096, ninp1024, nhead64) model fabric.setup(model) optimizer torch.optim.Adam(model.parameters(), lr0.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_policypolicy) fabric L.Fabric(..., strategystrategy)对于 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_params10000) # 3. Pass it to the FSDPStrategy object strategy FSDPStrategy(auto_wrap_policypolicy)PyTorch 在torch.distributed.fsdp.wrap下提供了多种这类函数式策略可供选用。从源码看Fabric 的 FSDPStrategy 会把auto_wrap_policy处理进底层FullyShardedDataParallel的构造参数_auto_wrap_policy_kwargs并默认设置use_orig_paramsTrue以支持优化器与模型的联合 setup、多参数组以及torch.compile()。验证 FSDP 是否生效对比上文示例中torch.cuda.memory_summary()打印的峰值显存与常规 DDP 训练的差异。官方在 A100 40GB GPU、Lightning 2.1、PyTorch 2.1 环境下的基准数据如下注意这是在特定硬件/版本下的参考值实际数值随模型与集群变化指标DDPFSDP显存MB26,95311,578单次迭代时间sec0.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_sizedataset.vocab_size) # Fast: Creates the model on the GPU directly with fabric.init_module(): model Transformer(vocab_sizedataset.vocab_size) # Recommended for FSDP: with fabric.init_module(empty_initTrue): model Transformer(vocab_sizedataset.vocab_size)对于 FSDP官方推荐设置empty_initTrue它会创建不分配任何显存的假参数meta device 上的参数真正的初始化被推迟到fabric.setup()——此时 FSDP 已完成分片并重新创建真实参数从而可以初始化更大的模型。从源码看FSDPStrategy.module_init_context 在empty_initTrue时会将模块创建置于 meta device 上下文并在setup_module阶段物化。更多empty_initTrue的使用场景见 模型初始化指南。切分策略用显存换速度默认情况下FSDP 会自动切分 1) 模型权重、2) 反向传播中的梯度、3) 优化器状态。通过sharding_strategy可以调整切分范围以权衡显存与速度strategy FSDPStrategy( # Default: Shard weights, gradients, optimizer state (1 2 3) sharding_strategyFULL_SHARD, # Shard gradients, optimizer state (2 3) sharding_strategySHARD_GRAD_OP, # Full-shard within a machine, replicate across machines sharding_strategyHYBRID_SHARD, # Dont shard anything (similar to DDP) sharding_strategyNO_SHARD, ) fabric L.Fabric(..., strategystrategy)选择切分策略的推荐顺序Recipe先用默认设置FULL_SHARD。它最省显存但最慢尝试SHARD_GRAD_OP。如果显存不足回退到默认的FULL_SHARD否则通常会看到迭代速度提升跨多机训练时尝试HYBRID_SHARD机器内全切分、机器间复制。官方在 A100 40GB、Lightning 2.1、PyTorch 2.1 下的基准对比指标DDPNO_SHARDSHARD_GRAD_OPFULL_SHARD显存MB26,95323,18111,81511,578迭代时间sec0.260.300.310.36数据清晰展示了切分越多、越省显存、越慢的权衡曲线从NO_SHARD到FULL_SHARD显存从 23GB 降到 11.5GB迭代时间从 0.30s 升到 0.36s。用显存换速度激活检查点与 CPU 卸载训练 10B 参数模型或需要极大 batch size 时可考虑以速度为代价换取更多显存两条途径分别是激活检查点与 CPU 卸载。激活检查点Activation checkpointing激活值前向中各层的中间输出在反向传播计算梯度时需要用到默认会贯穿整个前向被存储。启用激活检查点后可以选择丢弃部分层的激活值、在反向需要时动态重算。这会略微降低训练速度但显著降低显存占用腾出的显存可用于增大模型容量或 batch sizestrategy FSDPStrategy( # Enable activation checkpointing on these layers activation_checkpointing_policy{ nn.TransformerEncoderLayer, nn.TransformerDecoderLayer, }, ) fabric L.Fabric(..., strategystrategy)典型实践是把activation_checkpointing_policy设为与auto_wrap_policy相同通常就是你的 transformer block包含 attention 与 feed-forward。CPU 卸载CPU offload最激进的显存节省手段是把参数卸载到 CPU 内存# Set cpu_offloadTrue strategy FSDPStrategy(..., cpu_offloadTrue) fabric L.Fabric(..., strategystrategy)代价是训练速度大幅下降——每个前向都需要在 CPU 与 GPU 之间传输参数。仅当 CPU 内存充足、且其他扩展手段无法提供足够显存节省时才应使用。官方基准A100 40GB、Lightning 2.1、PyTorch 2.1显示 CPU 卸载带来约 4 倍显存节省但迭代时间增加约 10 倍指标DDPFSDPFSDP CPU offload显存MB26,95311,5782,825迭代时间sec0.260.363.24保存与加载大模型检查点大模型训练成本高昂务必将检查点逻辑纳入训练循环。Fabric 提供了高效保存大型检查点的方法——直接把模型/优化器等对象放进 state 字典而不是手动序列化 state dict# 1. Define model, optimizer, and other training loop state state {model: model, optimizer: optimizer, iter: iteration} # DONT do this (inefficient): # state {model: model.state_dict(), optimizer: optimizer.state_dict(), ...} # 2. Save using Fabrics method fabric.save(path/to/checkpoint/file, state) # DONT 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_typesharded) # Save a single, consolidated checkpoint file strategy FSDPStrategy(state_dict_typefull)如何选择检查点格式state_dict_typesharded适用于预训练超大规模模型保存快、占用显存少但可移植性差需要额外步骤把分片检查点转换为常规检查点见 分布式检查点转换指南state_dict_typefull适用于预训练中小规模模型10B 参数、微调以及需要可移植性的场景。加载检查点同样简单且 Fabric 会自动识别路径中是full还是sharded格式# 1. Define model, optimizer, and other training loop state state {model: model, optimizer: optimizer, iter: iteration} # 2. Load using Fabrics method fabric.load(path/to/checkpoint/file, state) # DONT do this (inefficient): # model.load_state_dict(torch.load(path/to/checkpoint/file))需要注意full格式的检查点可以被所有策略加载而sharded格式只能被 FSDP 加载。更多特性见 检查点指南。进阶性能优化技巧关闭优化器的 foreachPyTorch 常用优化器的foreachTrue选项会加速参数与状态更新但可能带来轻微显存峰值模型越大越明显。若出现不希望的显存模式可关闭optimizer torch.optim.AdamW(model.parameters(), foreachFalse)限制 all-gatherlimit_all_gathers当训练接近单卡显存上限时可能出现 CUDA malloc retriesGPU 显存即将耗尽、崩溃前尝试释放缓存内存的现象频繁发生时对速度影响显著。常规做法是略微减小 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_gathersFalse, # Enable this if you are close to the max. GPU memory usage limit_all_gathersTrue, ) fabric L.Fabric(..., strategystrategy)可以在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, biasFalse) self.w2 nn.Linear(hidden_dim, dim, biasFalse) self.w3 nn.Linear(dim, hidden_dim, biasFalse) def forward(self, x): return self.w2(F.silu(self.w1(x)) * self.w3(x))该模型有三个线性层w1与w3的输出随后被逐元素相乘结果喂给w2。因此w1与w3适合列并行其输出可轻松与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 PyTorchs distributed tensor APIs to parallelize the model plan { w1: ColwiseParallel(), w2: RowwiseParallel(), w3: ColwiseParallel(), } parallelize_module(model, tp_mesh, plan) return model然后配置 Fabric 的ModelParallelStrategyimport lightning as L from lightning.fabric.strategies import ModelParallelStrategy # 1. Pass the parallelization function to the strategy strategy ModelParallelStrategy(parallelize_fnparallelize_feedforward) # 2. Configure devices and set the strategy in Fabric fabric L.Fabric(acceleratorcuda, devices2, strategystrategy) 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 张 GPUimport 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, biasFalse) self.w2 nn.Linear(hidden_dim, dim, biasFalse) self.w3 nn.Linear(dim, hidden_dim, biasFalse) 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 PyTorchs 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_fnparallelize_feedforward) fabric L.Fabric(acceleratorcuda, devices2, strategystrategy) fabric.launch() # Initialize the model model FeedForward(8192, 8192) model fabric.setup(model) # Define the optimizer optimizer torch.optim.AdamW(model.parameters(), lr3e-3) optimizer fabric.setup_optimizers(optimizer) # Define dataset/dataloader dataset RandomDataset(8192, 64) dataloader torch.utils.data.DataLoader(dataset, batch_size8) 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(fIteration {i} complete) fabric.print(fPeak memory usage: {torch.cuda.max_memory_allocated() / 1e9:.02f} GB)官方基准显示随着 GPU 数量翻倍单卡峰值显存近似减半配置1 GPU无 TP2 GPUs4 GPUs8 GPUs每卡显存4.04 GB2.03 GB1.02 GB0.60 GBTP 的数据加载注意事项在张量并行的模型中参与同一 TP 组的每张 GPU 必须收到完全相同的输入否则训练无法收敛。因此在数据集/dataloader 中做 shuffle、或应用随机变换/数据增强时必须正确设置随机种子。这也意味着 TP 下的全局 batch size 受限于单卡显存。要扩大 batch size 并加速训练需要把 TP 与数据并行尤其是 FSDP组合使用——这正是下一节的 2D 并行。实战三2D 并行TP FSDP扩展到数百张 GPU2D 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 PyTorchs 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 PyTorchs FSDP2 APIs to parallelize the model fully_shard(model.w1, meshdp_mesh) fully_shard(model.w2, meshdp_mesh) fully_shard(model.w3, meshdp_mesh) fully_shard(model, meshdp_mesh) return model函数必须把model作为第一个参数、DeviceMesh作为第二个参数。随后把函数传给ModelParallelStrategy并指定数据并行与张量并行的规模import lightning as L from lightning.fabric.strategies import ModelParallelStrategy strategy ModelParallelStrategy( parallelize_fnparallelize_feedforward, # Define the size of the 2D parallelism # Set these to auto (default) to apply TP intra-node and FSDP inter-node data_parallel_size2, tensor_parallel_size2, ) fabric L.Fabric(acceleratorcuda, devices4, strategystrategy) fabric.launch()device mesh 的划分逻辑在上述 4 卡示例中Fabric 创建的 device mesh 会把 GPU 0-1 与 GPU 2-3 各分为一组因为data_parallel_size2每组 2 张 GPU 对应tensor_parallel_size2。随后调用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 张 GPUimport 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, biasFalse) self.w2 nn.Linear(hidden_dim, dim, biasFalse) self.w3 nn.Linear(dim, hidden_dim, biasFalse) 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 PyTorchs 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 PyTorchs FSDP2 APIs to parallelize the model fully_shard(model.w1, meshdp_mesh) fully_shard(model.w2, meshdp_mesh) fully_shard(model.w3, meshdp_mesh) fully_shard(model, meshdp_mesh) return model strategy ModelParallelStrategy( parallelize_fnparallelize_feedforward, data_parallel_size2, tensor_parallel_size2, ) fabric L.Fabric(acceleratorcuda, devices4, strategystrategy) fabric.launch() # Initialize the model model FeedForward(8192, 8192) model fabric.setup(model) # Define the optimizer optimizer torch.optim.AdamW(model.parameters(), lr3e-3) optimizer fabric.setup_optimizers(optimizer) # Define dataset/dataloader dataset RandomDataset(8192, 128) dataloader torch.utils.data.DataLoader(dataset, batch_size8) 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(fIteration {i} complete) fabric.print(fPeak 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_sizeauto, tensor_parallel_sizeauto, )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(seed42) dataloader DataLoader(dataset, batch_size8, shuffleTrue) # 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),仅供参考

相关新闻

Gazebo仿真中Unitree A1机器人控制器优化实践

Gazebo仿真中Unitree A1机器人控制器优化实践

1. 项目背景与核心需求在机器人仿真开发领域,Gazebo作为一款功能强大的物理仿真引擎,被广泛应用于各类机器人控制算法的验证环节。宇树(Unitree)作为四足机器人领域的知名厂商,其官方提供的控制器在Gazebo仿真环境中表…

2026/9/20 15:40:48 阅读更多 →
VMware 安装 Ubuntu 24.04 LTS 完整教程:从虚拟机创建到开发环境搭建

VMware 安装 Ubuntu 24.04 LTS 完整教程:从虚拟机创建到开发环境搭建

1. 为什么还要用虚拟机跑 Ubuntu1.1 虚拟机方案在当下的真实价值很多人第一反应是:现在云服务器这么便宜,WSL 也这么好用,为什么还要在本地装虚拟机跑 Ubuntu?我用了七八年虚拟机,也用过云主机和 WSL,说句实…

2026/9/22 1:08:02 阅读更多 →
漏磁图像缺陷检测:Laplace锐化+分水岭+KNN分类实战

漏磁图像缺陷检测:Laplace锐化+分水岭+KNN分类实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

2026/9/22 15:28:03 阅读更多 →

最新新闻

KMeans聚类在宿舍分配中的实战:特征工程到K值选择

KMeans聚类在宿舍分配中的实战:特征工程到K值选择

简介:针对高校宿舍分配场景,这份基于KMeans聚类算法的Python源码包提供了从数据预处理、模型训练到结果可视化的完整实现,适合需要将无监督学习落地到实际管理问题的数据科学初学者或高校信息管理相关技术人员。压缩包共13个文件,…

2026/9/23 18:38:49 阅读更多 →
fpm 构建 Solaris SRV4 软件包(solaris 输出格式)完全指南

fpm 构建 Solaris SRV4 软件包(solaris 输出格式)完全指南

fpm 构建 Solaris SRV4 软件包(solaris 输出格式)完全指南 【免费下载链接】fpm Effing package management! Build packages for multiple platforms (deb, rpm, etc) with great ease and sanity. 项目地址: https://gitcode.com/gh_mirrors/fp/fpm …

2026/9/23 18:38:49 阅读更多 →
Java Swing数独游戏工程级实现与难度控制

Java Swing数独游戏工程级实现与难度控制

简介:本资源是一份面向Java初学者与课程设计实践者的完整数独小游戏开发项目,适用于高校Java程序设计、GUI编程或软件工程类课程作业参考。项目基于Swing构建图形界面,代码结构清晰,涵盖游戏逻辑、难度生成、用户交互及资源管理等…

2026/9/23 18:38:49 阅读更多 →
Fedora开发环境避坑指南:保姆级教程解决常见报错

Fedora开发环境避坑指南:保姆级教程解决常见报错

Fedora开发环境避坑指南:保姆级教程解决常见报错 盯着屏幕上一片红色的StackTrace,是不是感觉脑子瞬间宕机?刚把Fedora装好,连个Python环境都跑不通,报错信息长得像天书,根本不知道从哪下手。别慌,这份保姆级教程就是为你…

2026/9/23 18:38:49 阅读更多 →
基于 TVM 编译栈的 WebAssembly 独立深度学习推理:wasm-standalone 项目实战解析

基于 TVM 编译栈的 WebAssembly 独立深度学习推理:wasm-standalone 项目实战解析

编译器深度学习模型优化 【免费下载链接】tvm Open deep learning compiler stack for cpu, gpu and specialized accelerators 项目地址: https://gitcode.com/gh_mirrors/tvm7/tvm 点击查看 免费下载 本文围绕仓库中的 apps/wasm-standalone 实验性项目&#xff…

2026/9/23 18:38:48 阅读更多 →
2026最新怎么注册营业执照,程序员如何搭建个人开发环境

2026最新怎么注册营业执照,程序员如何搭建个人开发环境

2026最新怎么注册营业执照,程序员如何搭建个人开发环境 刚学会Python语法,打开VS Code却不知从何下手?这是90%新手最真实的困境。2026最新的技术栈迭代很快,但基础项目搭建逻辑没变。很多教程只讲“怎么写代码”,却忽略了“怎么…

2026/9/23 18:37:48 阅读更多 →

日新闻

3招搞定手机怎么下载微信面试难题实战项目解析

3招搞定手机怎么下载微信面试难题实战项目解析

3招搞定手机怎么下载微信面试难题实战项目解析 面试被问“手机怎么下载微信”背后的原理,90%的人答不上来。别笑,这看似弱智的问题,实则是考察你对移动应用分发机制、安全校验及网络协议理解的试金石。我带过不少校招新人,他们背了八股文,却连一个A…

2026/9/23 0:00:23 阅读更多 →
2k显示屏性能优化踩坑:版本升级后API全变了,这份源码解析救了我

2k显示屏性能优化踩坑:版本升级后API全变了,这份源码解析救了我

2k显示屏性能优化踩坑:版本升级后API全变了,这份源码解析救了我 刚把开发环境的显示器从1080P换到2K,跑老项目直接报错,版本升级后 API…

2026/9/23 0:01:25 阅读更多 →
3步搞定美眉图实战项目,告别官方文档抓不住重点

3步搞定美眉图实战项目,告别官方文档抓不住重点

3步搞定美眉图实战项目,告别官方文档抓不住重点 官方文档翻了三遍还是云里雾里?别急,美眉图在实战项目中常被用来做数据可视化,但它的原理比你想的简单。今天咱们直接上手,用一个完整的小项目把美眉图跑通,不再死磕那些冗长的理论说明。…

2026/9/23 0:01:25 阅读更多 →

周新闻

Flutter for OpenHarmony游戏卡片渐变背景实战:从原理到性能优化

Flutter for OpenHarmony游戏卡片渐变背景实战:从原理到性能优化

直接铺开项目本身吧。这几个月我一直在折腾一件事:用Flutter给OpenHarmony做一款游戏集合类的App,说白了就是把若干小游戏塞进一个壳里,用统一入口分发。这个方向本身不算新鲜,真正让我花了不少心思的,是首页那堆游戏卡…

2026/9/23 4:55:02 阅读更多 →
Word表格编号全攻略:从列表编号到题注交叉引用

Word表格编号全攻略:从列表编号到题注交叉引用

写Word文档,最让人头疼的往往是那些“看起来不起眼”的小问题。比如表格编号这事:今天在表后面多加了两个空白行,明天给客户交稿前发现整个章节的编号全部错位,光是挨个改序号就能耗掉大半个下午。我前阵子帮人整理一份上百页的技…

2026/9/23 4:49:06 阅读更多 →
从第一个站到第二个站:独立开发者的静态网站选型与落地实践

从第一个站到第二个站:独立开发者的静态网站选型与落地实践

1. 项目概述1.1 核心需求解析做独立开发者这几年,说实话,第一个网站上线的那天晚上我兴奋得没睡着。但等它跑了半年,流量惨淡、功能臃肿、代码自己都懒得看第二遍之后,我才慢慢琢磨明白一个道理:第一个网站是练手&…

2026/9/23 9:53:41 阅读更多 →

月新闻

持续集成 流水线自动化与 声明式交付 实践:原型怎样变成可用功能

持续集成 流水线自动化与 声明式交付 实践:原型怎样变成可用功能

持续集成 流水线自动化与 声明式交付 实践:原型怎样变成可用功能分类:[AI/大模型]细分主题:AI 增强型 CI/CD 流水线自动化与 GitOps 实践:Agent 工作流、工具调用与任务拆解:从原型到生产的验收清单很多团队在尝试用大…

2026/9/23 9:53:40 阅读更多 →
容器编排 生产环境运维与排障实战:复盘记录怎样真正派上用场

容器编排 生产环境运维与排障实战:复盘记录怎样真正派上用场

容器编排 生产环境运维与排障实战:复盘记录怎样真正派上用场分类:[工程技术]细分主题:Kubernetes 生产环境运维与排障实战:可复制的项目复盘模板与决策记录大部分团队的事故复盘报告,最后都变成了躺在 Confluence 或钉…

2026/9/23 9:53:40 阅读更多 →
容器 容器化技术与镜像安全管理:核心链路应该先拆哪一步

容器 容器化技术与镜像安全管理:核心链路应该先拆哪一步

容器 容器化技术与镜像安全管理:核心链路应该先拆哪一步分类:[工程技术]细分主题:Docker 容器化技术与镜像安全管理:核心链路的逐步实现与关键代码取舍面对一个积累了五六年历史包袱的单体架构应用(包含 Web 接口、后台…

2026/9/23 9:53:40 阅读更多 →