新闻详情

新闻详情

首页 / 资讯中心 / 详情

PyTorch Lightning 张量并行与 2D 并行实战:用 `ModelParallelStrategy` 训练 Llama 3 7B

发布时间:2026/9/20 23:33:41来源:尧图网络
PyTorch Lightning 张量并行与 2D 并行实战:用 `ModelParallelStrategy` 训练 Llama 3 7B
人工智能深度学习机器学习预训练分布式训练微调【免费下载链接】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点击查看免费下载导读本文以仓库 examples/pytorch/tensor_parallel/README.md 为主线完整讲解如何在 PyTorch Lightning 中通过ModelParallelStrategy为 70 亿参数规模的 Llama 3 模型同时启用**张量并行Tensor Parallelism, TP**与FSDP 数据并行组合即 2D 并行。你将掌握ModelParallelStrategy的完整配置参数、configure_model钩子的并行化写法、基于torch.distributed.tensor.parallel的 Colwise/RoWwise/Sequence 切分方案、FSDP2 的fully_shard分片逻辑以及分布式 checkpoint 的保存与恢复机制并能在 4 卡环境下复现整个训练流程。示例概览一个开箱即用的 Llama 3 并行训练骨架该示例位于 examples/pytorch/tensor_parallel/目录结构如下train.py —— 入口训练脚本定义Llama3LightningModule 与 Trainer 配置model.py —— Llama 3 7B 的纯 PyTorch 实现RMSNorm、RoPE、GQA 注意力、SwiGLU FFNparallelism.py —— 核心的并行化函数parallelize()负责把 TP、FSDP2 与激活检查点施加到模型上data.py —— 带固定种子的随机 token 玩具数据集。示例演示的能力正是仓库项目描述所强调的场景对任意规模的模型在不修改模型源代码的前提下通过策略层组合多种并行维度TP FSDP2让同一份模型代码可以在单卡、单机多卡乃至多节点上无缝扩展。环境要求与安装根据 README运行本示例需要满足两个硬性条件PyTorch 2.3 及以上版本README 给出的安装命令为pip install torch2.3。需要说明的是ModelParallelStrategy使用的 FSDP2 APItorch.distributed._composable.fsdp在 PyTorch 2.4 及以上才完全稳定策略源码的 docstring 也明确指出 Requires PyTorch 2.4 or newer见 src/lightning/pytorch/strategies/model_parallel.py#L73。建议直接使用 2.4。至少 4 张 GPU每张显存不低于 24 GB。训练脚本入口处有显式断言train.py#L79assert torch.cuda.device_count() 4, This example requires at least 4 GPUs with 24 GB of memory each.此外依赖torch.distributed.tensor.parallel与torch.distributed.device_mesh它们随 PyTorch 发行版自带无需额外安装第三方包。运行示例与输出解读README 给出的启动方式如下cd examples/pytorch/tensor_parallel python train.py在 4 卡 CUDA 环境下预期输出如下GPU available: True (cuda), used: True TPU available: False, using: 0 TPU cores HPU available: False, using: 0 HPUs Number of model parameters: 6.7 B Starting training ... Initializing distributed: GLOBAL_RANK: 0, MEMBER: 1/4 Initializing distributed: GLOBAL_RANK: 1, MEMBER: 2/4 Initializing distributed: GLOBAL_RANK: 3, MEMBER: 4/4 Initializing distributed: GLOBAL_RANK: 2, MEMBER: 3/4 ---------------------------------------------------------------------------------------------------- distributed_backendnccl All distributed processes registered. Starting with 4 processes ---------------------------------------------------------------------------------------------------- LOCAL_RANK: 0 - CUDA_VISIBLE_DEVICES: [0,1,2,3] LOCAL_RANK: 3 - CUDA_VISIBLE_DEVICES: [0,1,2,3] LOCAL_RANK: 1 - CUDA_VISIBLE_DEVICES: [0,1,2,3] LOCAL_RANK: 2 - CUDA_VISIBLE_DEVICES: [0,1,2,3] Epoch 0: 100%|█████████████████████████████████████████████| 10/10 [01:4900:00, 0.09it/s, v_num2] Trainer.fit stopped: max_epochs1 reached. Saving a (distributed) checkpoint ... Training successfully completed! Peak memory usage: 36.73 GB这段输出反映了几个关键事实模型规模Number of model parameters: 6.7 B表明这是一个约 67 亿参数的模型。由于采用了trainer.init_module(empty_initTrue)在 meta 设备上初始化train.py#L66-L67模型参数并不真正分配显存因此可以仅凭参数计数来统计规模而实际的权重内存由后续的 TP/FSDP 切分承担。分布式初始化Lightning 自动启动 4 个进程MEMBER: 1/4 ... 4/4使用 nccl 后端并打印各进程的LOCAL_RANK与CUDA_VISIBLE_DEVICES。训练配置limit_train_batches10、max_epochs1使示例在约 2 分钟内跑完便于快速验证。分布式 checkpoint输出中的Saving a (distributed) checkpoint ...对应ModelParallelStrategy默认的save_distributed_checkpointTrue行为——每个 rank 把自身的分片权重与优化器状态保存为一个目录型 checkpoint。峰值显存Peak memory usage: 36.73 GB是在 4×24 GB 环境下的实测值说明 TPFSDP 组合确实将 6.7B 参数模型的显存占用压到了单卡 24 GB 以内若不用任何并行策略6.7B 的 BF16 权重本身就需要约 13.4 GB 参数存储加上优化器状态与激活值将远超单卡 24 GB。ModelParallelStrategy参数详解在 train.py#L50-L55 中策略被配置为 2D 并行strategy ModelParallelStrategy( # Define the size of the 2D parallelism # Set to auto to apply TP intra-node and DP inter-node data_parallel_size2, tensor_parallel_size2, )对照策略类的构造函数src/lightning/pytorch/strategies/model_parallel.py#L86-L101可配置参数如下参数默认值含义data_parallel_sizeauto每个数据并行组内的设备数。auto时取集群中的节点数即 FSDP 跨节点切分。tensor_parallel_sizeauto每个张量并行组内的设备数。auto时取单节点内的 GPU 数即 TP 在节点内切分。save_distributed_checkpointTrue为True时每个 rank 将权重/优化器分片保存到目录型 checkpoint文件数与 world size 相同为False时在 rank 0 上聚合完整权重后保存为单文件。process_group_backendNone分布式进程组后端None时根据设备自动选择CUDA 下为 nccl。timeout默认进程组超时分布式通信超时时间。auto 的语义与 DeviceMesh 的构建在setup_environment()中src/lightning/pytorch/strategies/model_parallel.py#L147-L159策略会将auto展开为具体数值然后构建二维DeviceMeshif self._data_parallel_size auto: self._data_parallel_size self.num_nodes if self._tensor_parallel_size auto: self._tensor_parallel_size self.num_processes self._device_mesh _setup_device_mesh( self._data_parallel_size, self._tensor_parallel_size, self.world_size, self.root_device ) # Users can access device mesh in LightningModule.configure_model() self.lightning_module._device_mesh self._device_mesh也就是说4 卡单机 data_parallel_size2, tensor_parallel_size2时device mesh 将 GPU 划分为 2 个 TP 组、每组 2 张卡GPU 0-1 一组、GPU 2-3 一组官方文档 docs/source-pytorch/advanced/model_parallel/tp_fsdp.rst#L102-L104 有同样说明。多节点训练时推荐直接使用autoTP 局限在节点内NCCL 带宽高、延迟低FSDP 跨节点通信可被计算隐藏/预取从而最小化延迟与网络带宽占用。这个TP 节点内、DP 节点间的最佳实践在 docs/source-pytorch/advanced/model_parallel/tp_fsdp.rst#L209-L213 中有专门阐述。策略还会在setup()阶段强制要求 LightningModule 重写configure_model()钩子否则抛出TypeError同时会拒绝检测到旧版torch.distributed.fsdp.FullyShardedDataParallel包装的模型src/lightning/pytorch/strategies/model_parallel.py#L169-L178因为该策略只支持 PyTorch ≥ 2.4 的 FSDP2 API。用configure_model钩子承载并行化逻辑示例中Llama3模块的关键设计是不在__init__里硬编码并行逻辑而是放到 Lightning 专属钩子configure_model中train.py#L19-L22def configure_model(self): # User-defined function that applies the desired parallelizations specific to the model # (TP, FSDP2, activation checkpointing, ...) parallelize(self.model, device_meshself.device_mesh)这里的self.device_mesh正是ModelParallelStrategy在setup_environment()阶段挂载到模块上的二维DeviceMesh。把并行化放进钩子而不是构造函数好处是模型源代码保持干净同样的模型可以按需选择不同的并行配置纯 TP、纯 FSDP 或 2D这正是零代码改动扩展模型规模的核心机制。初始化权重的工作被安排在on_train_starttrain.py#L24-L25因为并行切分完成后权重才真正被物化到设备上。整个流程为trainer.init_module(empty_initTrue)在 meta 设备创建模型 →configure_model完成并行切分 → 策略在setup()中通过_materialize_distributed_module将模型物化到 root devicesrc/lightning/pytorch/strategies/model_parallel.py#L180。并行化的 lossloss_parallelTP 切分最后一个输出层时logits 的类别维度被切到多张卡上因此交叉熵必须在并行上下文中计算。示例在 train.py#L32-L37 用torch.distributed.tensor.parallel.loss_parallel包裹前向损失与backwardwith loss_parallel(): return F.cross_entropy(output.reshape(-1, output.size(-1)), labels.reshape(-1)) def backward(self, *args, **kwargs): with loss_parallel(): super().backward(*args, **kwargs)这要求输出层的并行方案设置use_local_outputFalse与output_layoutsShard(-1)见下文 parallelism.py 分析让每张卡持有类别维的本地分片再由loss_parallel负责在梯度计算时做必要的跨卡规约。parallelize()深度拆解TP 切分方案parallelism.py 中的parallelize()函数取自并修改自 torchtitan 的 Llama 并行方案是整份示例的技术核心。它先取出两个子网格dp_mesh device_mesh[data_parallel] tp_mesh device_mesh[tensor_parallel]第一层并行嵌入、输出与根归一化层当tp_mesh.size() 1时首先对模型骨架应用一个并行计划parallelism.py#L36-L52plan { tok_embeddings: RowwiseParallel(input_layoutsReplicate()), output: ColwiseParallel( input_layoutsShard(1), output_layoutsShard(-1), # 类别维分片配合 loss_parallel use_local_outputFalse, ), norm: SequenceParallel(), layers.0: PrepareModuleInput( input_layouts(Replicate(), None), desired_input_layouts(Shard(1), None), use_local_outputTrue, ), } model parallelize_module(model, tp_mesh, plan)要点tok_embeddings用RowwiseParallel(input_layoutsReplicate())词表维度按行切分输入 token 在 TP 组内保持复制output用ColwiseParallel且output_layoutsShard(-1)类别维度分片配合前面loss_parallel实现并行 lossnorm用SequenceParallel在序列维度上并行计算 RMSNormlayers.0用PrepareModuleInput把第一个 transformer block 的输入从 Replicate 布局转换为序列分片布局为后续 block 的 SequenceParallel 做准备。第二层并行每个 Transformer block随后对每个TransformerBlock施加更细粒度的计划parallelism.py#L55-L82plan { attention: PrepareModuleInput( input_layouts(Shard(1), None), desired_input_layouts(Replicate(), None), ), attention.wq: ColwiseParallel(), attention.wk: ColwiseParallel(), attention.wv: ColwiseParallel(), attention.wo: RowwiseParallel(output_layoutsShard(1)), attention_norm: SequenceParallel(), feed_forward: PrepareModuleInput( input_layouts(Shard(1),), desired_input_layouts(Replicate(),), ), feed_forward.w1: ColwiseParallel(), feed_forward.w2: RowwiseParallel(output_layoutsShard(1)), feed_forward.w3: ColwiseParallel(), ffn_norm: SequenceParallel(), }注意力部分wq/wk/wv三个投影做列并行ColwiseParallel输出特征维被切分每卡只持有部分 head 的权重wo做行并行RowwiseParallel聚合各卡的注意力输出attention模块的输入先转为复制布局保证每张卡看到完整的序列数据后各自处理本地 head。FFN 部分SwiGLU 的w1/w3做列并行、w2做行并行这是标准 Megatron 式 FFN 切分。归一化层全部使用SequenceParallel。同时必须手动把注意力头数除以 TP 组大小parallelism.py#L76-L79attn_layer.n_heads attn_layer.n_heads // tp_mesh.size() attn_layer.n_kv_heads attn_layer.n_kv_heads // tp_mesh.size()因为wq被列切分后每张卡持有的实际是本地 head的权重这正是 model.py 中Attention使用self.n_heads计算 head_dim 的原因head 数必须与切分后的本地维度一致前向才不会越界。这一点体现了模型代码需与 TP 语义配合模型本身model.py是标准 Llama 结构但n_heads这种与并行度耦合的字段由并行化函数在运行时调整。FSDP2 数据并行与激活检查点当dp_mesh.size() 1时parallelize()使用 FSDP2 的fully_shard在数据并行维度分片parallelism.py#L84-L104mp_policy MixedPrecisionPolicy(param_dtypetorch.bfloat16, reduce_dtypetorch.float32) fsdp_config {mesh: dp_mesh, mp_policy: mp_policy} for layer_id, transformer_block in model.layers.items(): transformer_block checkpoint_wrapper(transformer_block) reshard_after_forward int(layer_id) len(model.layers) - 1 fully_shard( transformer_block, **fsdp_config, reshard_after_forwardreshard_after_forward, ) model.layers[layer_id] transformer_block model fully_shard(model, **fsdp_config)几个实现要点精度策略MixedPrecisionPolicy(param_dtypetorch.bfloat16, reduce_dtypetorch.float32)让参数以 BF16 存储、梯度以 FP32 规约。注释特别提醒目前ModelParallelStrategy尚未完全尊重Fabric(precision...)/Trainer 的全部精度设置FSDP 的mp_policy需要用户在此处手动管理parallelism.py#L87-L89。激活检查点通过torch.distributed.algorithms._checkpoint.checkpoint_wrapper对每个 block 开启显著降低激活值显存是 6.7B 模型在 4×24 GB 上可训的关键手段之一。reshard 优化除最后一个 block 外都设置reshard_after_forwardTrue最后一个 block 不重分片因为 FSDP 紧接着就会预取它避免不必要的通信。嵌套分片先对每个 block 调用fully_shard最后对整个模型再调用一次fully_shard形成分层分片按 block 粒度分片兼顾显存节省与通信效率。官方文档将这种组合概括为TP 提供计算可扩展性、FSDP 提供显存效率两者互补docs/source-pytorch/advanced/model_parallel/tp_fsdp.rst#L5-L6。数据加载的并行语义2D 并行对数据有特殊要求同一个 TP 组内的各 GPU 必须看到相同的 batch而数据并行维度之间必须看到不同的 batch详见 docs/source-pytorch/advanced/model_parallel/tp_fsdp.rst#L233-L240。示例的RandomTokenDatasetdata.py本身用固定种子 42 生成数据保证各 rank 上数据一致data.py#L9-L15然后由 Trainer 的分布式采样器完成分组def train_dataloader(self): dataset RandomTokenDataset(vocab_sizeself.model_args.vocab_size, seq_length128) # Trainer configures the sampler automatically for you such that # all batches in a tensor-parallel group are identical return DataLoader(dataset, batch_size8, num_workers4)其底层依据是策略覆写的distributed_sampler_kwargssrc/lightning/pytorch/strategies/model_parallel.py#L119-L124property def distributed_sampler_kwargs(self) - dict[str, Any]: data_parallel_mesh self.device_mesh[data_parallel] return {num_replicas: data_parallel_mesh.size(), rank: data_parallel_mesh.get_local_rank()}即采样器的 replica 数量 数据并行组大小、rank 组内局部 rank因此 TP 组内共享同一批数据、DP 组间切分数据。若你的数据集/增强存在随机性shuffle、数据增强等必须自行固定 seed否则同一个 TP 组内的数据会不一致导致并行结果错误。分布式 Checkpoint 的保存与加载README 输出中的Saving a (distributed) checkpoint ...对应策略默认save_distributed_checkpointTrue的行为。查看 src/lightning/pytorch/strategies/model_parallel.py#L301-L330 的实现save_distributed_checkpointTrue时使用torch.distributed.checkpoint将每个 rank 的权重分片与优化器分片写入目录型 checkpoint目录下文件数与 world size 相同模型与优化器 state 按optimizer_0, optimizer_1, ...命名存放元数据单独原子落盘save_distributed_checkpointFalse时各 rank 把完整状态聚合到 rank 0通过get_model_state_dict(..., full_state_dictTrue)与FSDP.rekey_optim_state_dict重排优化器键名后保存为单个文件src/lightning/pytorch/strategies/model_parallel.py#L255-L294。加载路径由load_checkpoint实现src/lightning/pytorch/strategies/model_parallel.py#L332-L347它直接把分片状态恢复到当前模型与优化器上load_model_state_dict/load_optimizer_state_dict被覆写为空操作因为状态在load_checkpoint中已经完成加载并支持strict_loading与weights_only选项。注意save_checkpoint明确不支持storage_options不经过CheckpointIO传参会抛出TypeError。实验性说明与后续方向README 特别用 NOTE 提醒ModelParallelStrategy目前是实验性 API随时可能变化。这与策略 docstring 中的 warning 一致src/lightning/pytorch/strategies/model_parallel.py#L68PyTorch 侧的 DTensor/TP/FSDP2 相关 API 同样处于实验阶段。因此在生产环境采用前应锁定 Lightning 与 PyTorch 版本并关注两个仓库的 API 演进。从官方文档 docs/source-pytorch/advanced/model_parallel/tp_fsdp.rst 可以了解到2D 并行最典型的落地场景是多节点大模型训练TP 因含阻塞式集合通信而更适合节点内高速互联环境FSDP 则能通过层预取把跨节点通信与计算重叠二者组合可在保持吞吐的同时把模型规模推到远超纯 FSDP 的上限。若只想先尝试纯 TP 一维并行可参考 docs/source-pytorch/advanced/model_parallel/tp.rst 中更简化的 FeedForward 示例若需要把torch.compile也纳入并行化流程可以研究文档 docs/source-pytorch/advanced/compile.rst#L147-L158 中在configure_model内组合编译与并行的写法。把本示例的data_parallel_size/tensor_parallel_size调整为auto并部署到多节点集群即可无缝扩展到大集群训练。赞分享人工智能深度学习机器学习预训练分布式训练微调【免费下载链接】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点击查看免费下载相关推荐使用 Lightning Fabric 的 ModelParallelStrategy 训练 7B 级大模型张量并行与 2D 并行TP FSDP实战使用 Lightning Fabric 的 ModelParallelStrategy 训练 7B 级大模型张量并行与 2D 并行TP FSDP实战人工智能深度学习机器学习预训练分布式训练微调使用 PyTorch Lightning 的 ModelParallelStrategy 实现张量并行Tensor Parallelism训练大型模型使用 PyTorch Lightning 的 ModelParallelStrategy 实现张量并行Tensor Parallelism训练大型模型 本文人工智能深度学习机器学习预训练分布式训练微调Lightning Fabric 张量并行Tensor Parallelism完全指南线性层切分原理、ModelParallelStrategy 配置与实战训练Lightning Fabric 张量并行Tensor Parallelism完全指南线性层切分原理、ModelParallelStrategy 配置与实人工智能深度学习机器学习预训练分布式训练微调上一篇OpenMontage 中的 FLUX.1 模型家族提示词实战指南FLUX1.1 pro、Kontext 与 Fill 的选择与使用规范下一篇Ghidra 快速上手指南平台要求、安装部署、GUI/无头/PyGhidra 多种运行模式与升级流程创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

耐克风格运动鞋商城网站模板:从布局到合规的完整设计指南 2026/9/21 0:15:48

耐克风格运动鞋商城网站模板:从布局到合规的完整设计指南

简介:这是一份“耐克品牌运动鞋商城网站模板”,定位为电商前端整站模板,适合前端学习者、中小商家或设计人员参考,可快速搭建包含商品陈列、分类浏览、购物车与结账流程的在线鞋店。包体共78个文件,大小1.98MB&#xf…

阅读更多 →
Sails 中 res.forbidden() 响应方法:403 权限拒绝的标准出口与自定义实战 2026/9/21 0:15:48

Sails 中 res.forbidden() 响应方法:403 权限拒绝的标准出口与自定义实战

后端 【免费下载链接】sails Realtime MVC Framework for Node.js 项目地址: https://gitcode.com/gh_mirrors/sa/sails 点击查看 免费下载 导读 res.forbidden() 是 Sails(Realtime MVC Framework for Node.js)内置的响应方法之一&#xf…

阅读更多 →
ANSYS Workbench静力分析入门:从零走通悬臂梁结构分析全流程 2026/9/21 0:15:48

ANSYS Workbench静力分析入门:从零走通悬臂梁结构分析全流程

1. 为什么我不建议你继续背菜单刚入行那会儿,我也干过背菜单这种事。把ANSYS Workbench里每一个菜单项抄在本子上,Toolbox里每个分析系统叫什么、右键菜单里每个选项在哪,背得滚瓜烂熟。结果第一次独立做一个支架的静力分析,打开软…

阅读更多 →
GraalVM Native Image PGO 常见问题解答:Profiling 实践、跨平台复用与 Profile 质量维护 2026/9/21 0:15:48

GraalVM Native Image PGO 常见问题解答:Profiling 实践、跨平台复用与 Profile 质量维护

编译器JIT编译语言运行时高性能计算内存管理 【免费下载链接】graal GraalVM compiles applications into native executables that start instantly, scale fast, and use fewer compute resources 🚀 项目地址: https://gitcode.com/gh_mirrors/gr/gra…

阅读更多 →
书店小说阅读App首页模板改造:从源码到毕设答辩的全流程指南 2026/9/21 0:15:48

书店小说阅读App首页模板改造:从源码到毕设答辩的全流程指南

简介:书店小说阅读应用手机首页模板是一份面向学校实训与毕业设计的商业源码包,专门帮助计算机专业学生和移动端开发者快速搭建具有真实项目质感的应用首页。zip压缩包共包含20个文件,以HTML/CSS/JavaScript前端代码为核心,14张PN…

阅读更多 →
Atlas 300V 24G推理加速卡实战:从环境部署到YOLOv5全流程跑通 2026/9/21 0:12:47

Atlas 300V 24G推理加速卡实战:从环境部署到YOLOv5全流程跑通

最近好几个做视觉落地的朋友都在问同一件事:Atlas 300V 24G 到底是不是运算加速卡?能不能拿它来部署 YOLO?说实话,第一次看到“运算加速卡”这个叫法时我也愣了一下,因为这个说法容易让人往通用 GPU 或训练卡上靠。但 …

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞