ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

大模型预训练与微调实战:MindSpore Transformers 并行策略与显存优化指南

大模型预训练与微调实战:MindSpore Transformers 并行策略与显存优化指南 1. 为什么大模型预训练与微调值得用 MindSpore 重做一遍大语言模型从 2023 年火到现在真正动手训过模型的人都知道训练一个像样的模型瓶颈从来不在算法本身而在工程。显存不够、通信太慢、并行策略选错、checkpoint 存不下来这些问题随便一个都能让你在集群前面耗上一整周。MindSpore Transformers 这套组合是我最近半年在国产算力环境里反复折腾后觉得值得认真聊一聊的方案。它解决的核心问题很直接在算力受限、显存吃紧的条件下怎么把一个大语言模型的预训练和微调跑起来并且跑得稳、跑得快、跑得省。这篇文章适合谁看如果你已经用 PyTorch 训过小模型想迁移到 MindSpore 生态或者你手上有昇腾卡、想搞清楚分布式并行到底怎么配再或者你只是想在自己的机器上微调一个 7B 级别的模型但总是 OOM那这篇内容应该能帮你省下不少试错时间。我会从整体设计思路讲到具体的并行配置、显存优化手段、实操步骤再到踩过的坑和排查方法尽量把每个决策背后的“为什么”说清楚。先说一个基本认知大语言模型的训练本质上是一个“显存换时间、通信换显存”的博弈过程。你的总显存就那么多模型参数、梯度、优化器状态、激活值这四座大山压在那里单卡放不下就必须切分。切分的方式决定了通信开销通信开销又决定了训练效率。MindSpore Transformers 在这件事上给出的答案是多维混合并行加上细粒度的重计算和异构内存管理。听起来很唬人但拆开来看每一块都有明确的适用场景和配置方法。我见过太多人一上来就照搬别人的并行配置结果要么通信瓶颈卡死要么显存还是爆。原因很简单并行策略和你的集群拓扑、模型结构、序列长度是强耦合的。8 卡 NVLink 全互联和 8 卡跨机走网络最优策略完全不同。所以这篇文章不会给你一个“万能配置”而是帮你建立一套判断逻辑让你能根据自己的硬件条件推导出合适的方案。2. 整体设计思路与并行方案选型2.1 大模型训练的显存账本怎么算在聊并行之前先把显存账算清楚。很多人 OOM 的时候第一反应是“模型太大了”但其实模型参数只是显存占用的一部分。以一个 7B 参数的模型为例用 FP16 存储参数本身需要 14GB梯度又是 14GB如果用的是 Adam 优化器优化器状态一阶矩和二阶矩在 FP32 下需要 56GB再加上激活值、临时缓冲区单卡 80GB 的卡都未必够。这就是为什么数据并行在单卡放不下模型时直接失效——每张卡都要存一份完整的模型副本。所以显存优化的第一原则是先想办法让单卡能放下模型状态再考虑怎么加速。切分模型状态的手段主要有三种张量并行Tensor ParallelismTP把单个权重矩阵按维度切开分到不同卡上流水线并行Pipeline ParallelismPP把不同的层分到不同卡上数据并行Data ParallelismDP每张卡存完整模型但处理不同数据。这三者组合起来就是所谓的 3D 并行。MindSpore Transformers 在这三种并行的基础上还支持序列并行Sequence Parallelism和优化器并行Optimizer Parallelism。序列并行是把长序列维度切开特别适合处理超长文本优化器并行则是把优化器状态切片存储减少每张卡的显存压力。这几个手段的组合方式决定了你最终能不能在给定硬件上跑起来。2.2 为什么选择多维混合并行而不是单一策略单一并行策略的问题在于扩展性差。纯数据并行在模型大到单卡放不下时就废了纯张量并行的通信量随卡数平方增长超过 8 卡基本不可用纯流水线并行则有气泡问题卡越多气泡越大。实际生产中几乎没有人用单一策略都是混合着来。MindSpore Transformers 的并行配置通过parallel_config来定义核心参数包括data_parallel、model_parallel、pipeline_stage、micro_batch_num等。我一般建议的配置逻辑是这样的先确定单卡能承受的最大模型状态据此决定 TP 和 PP 的切分粒度然后用剩余的卡做 DP 来提升吞吐。举个例子64 张卡训一个 70B 模型可以配成 TP8、PP4、DP2这样每张卡上的模型状态大约是 70B 的 1/32显存压力就小很多了。这里有个经验公式可以参考单卡显存需求 ≈ (参数量 × 精度字节数 × (1 梯度系数 优化器系数)) / (TP × PP) 激活值开销。梯度系数在 FP16 下约为 1优化器系数在 Adam FP32 下约为 4。激活值开销跟 batch size、序列长度、重计算策略有关后面会细讲。2.3 通信开销的隐藏成本并行策略选完之后真正的挑战才刚开始通信。张量并行每层都要做 AllReduce流水线并行在阶段边界要传激活值数据并行在梯度聚合时要 AllReduce。这些通信如果走的是慢速网络训练速度会断崖式下跌。我的实测经验是TP 尽量放在同一台机器内走 NVLink 或 HCCS 这种高速互联PP 可以跨机因为通信频率低DP 的梯度聚合可以跟计算重叠影响相对小。MindSpore 在这方面做了通信计算重叠的优化但前提是你的配置得合理。如果你把 TP8 配到跨机的 8 张卡上那通信延迟会让你怀疑人生。还有一个容易被忽略的点micro_batch_num 的设置。流水线并行中一个 global batch 会被切成多个 micro batch 在流水线中流动。micro_batch_num 越大气泡越小但显存占用也越大。一般建议 micro_batch_num 至少是 pipeline_stage 的 2 到 4 倍具体要看显存余量。3. 核心细节解析与实操要点3.1 环境准备与依赖安装MindSpore Transformers 的环境搭建有几个关键点。首先是 MindSpore 本身的版本要和 CANN 版本匹配这个在官方文档里有对应表装错了会各种报错。我一般用 conda 建一个独立环境然后按顺序装CANN → MindSpore → MindSpore Transformers。conda create -n ms_train python3.9 conda activate ms_train pip install mindspore2.3.0 pip install mindformers1.0.0 pip install sentencepiece transformers datasets注意transformers这个库的版本要控制好太新的版本可能和 mindformers 有接口冲突。我遇到过aimv2 is already used by a transformers config这类报错就是因为 transformers 版本太新配置名冲突了。解决办法是降级到 4.35 左右的版本或者手动改配置名。另外如果你在昇腾环境上跑要确认ASCEND_HOME环境变量指向正确的 CANN 路径并且npu-smi info能正常看到卡。这些基础检查不做后面出问题很难定位。3.2 模型配置文件的拆解MindSpore Transformers 用 YAML 文件来管理模型配置这个文件里的每一项都直接影响训练行为。以 Llama 7B 为例核心配置项包括model: model_config: type: LlamaConfig vocab_size: 32000 hidden_size: 4096 num_layers: 32 num_heads: 32 seq_length: 2048 arch: type: LlamaForCausalLM parallel_config: data_parallel: 2 model_parallel: 4 pipeline_stage: 2 micro_batch_num: 8 gradient_aggregation_group: 4 optimizer: type: AdamWeightDecay learning_rate: 1e-4 weight_decay: 0.01seq_length这个参数特别关键它直接决定激活值显存占用。2048 和 8192 的显存差距可能是好几倍。如果你的卡显存紧张先把 seq_length 降下来跑通了再往上加。gradient_aggregation_group是梯度聚合的组大小影响通信效率。一般设成 data_parallel 的大小或者其因数。设太小通信频繁设太大显存占用高。3.3 重计算策略的选择重计算Recompute是显存优化的杀手锏。原理很简单前向传播时不保存中间激活值反向传播时重新算一遍。代价是计算量增加约 30%收益是激活值显存大幅下降。MindSpore Transformers 支持多种重计算粒度full整层重计算省显存最多计算开销最大select选择性重计算只重算注意力部分none不重计算我的建议是如果显存够用优先none如果差一点用select如果差很多才上full。因为重计算本质上是拿时间换空间在算力受限的环境下这个交换未必划算。配置方式是在 YAML 里加recompute_config: recompute: select select_recompute: true parallel_optimizer_comm_recompute: trueparallel_optimizer_comm_recompute这个选项是重算优化器通信的能进一步省显存但会增加通信量要权衡。3.4 异构内存与 offload 机制当显存实在不够时offload 是最后的救命稻草。MindSpore 支持把优化器状态甚至部分参数卸载到主机内存需要时再加载回来。这个机制在微调场景下特别有用因为微调的 batch size 通常不大对吞吐要求没那么高。配置 offload 需要在 YAML 里设置optimizer: type: AdamWeightDecay optimizer_offload: true但要注意offload 会显著增加主机内存占用和 PCIe 通信量。如果你的主机内存不够大或者 PCIe 带宽是瓶颈offload 反而会让训练慢到无法接受。我一般只在微调 13B 以上模型且显存严重不足时才用。4. 实操过程与核心环节实现4.1 数据准备与预处理大模型训练的数据质量直接决定最终效果。MindSpore Transformers 支持多种数据格式最常用的是 bin 格式和 mindrecord 格式。bin 格式适合大规模预训练读写效率高mindrecord 格式适合小规模微调加载方便。数据预处理的核心步骤是分词和打包。分词用 tokenizer打包是把多条短文本拼成固定长度的序列减少 padding 浪费。我一般用 HuggingFace 的 datasets 库做预处理然后转成 MindSpore 能读的格式。from datasets import load_dataset from transformers import AutoTokenizer tokenizer AutoTokenizer.from_pretrained(llama_tokenizer) dataset load_dataset(json, data_filestrain.jsonl) def tokenize_function(examples): return tokenizer(examples[text], truncationTrue, max_length2048) tokenized dataset.map(tokenize_function, batchedTrue, num_proc8) tokenized.save_to_disk(tokenized_data)这里有个坑tokenizer 的pad_token要设好很多开源 tokenizer 默认没有 pad_token不设的话训练时会报错。一般把 pad_token 设成 eos_token 就行。4.2 启动分布式训练MindSpore 的分布式启动方式有几种最常用的是msrun和mpirun。msrun是 MindSpore 自带的启动工具配置简单推荐优先用。msrun --worker_num8 --local_worker_num8 \ --master_port8118 --node_rank0 \ --log_dir./logs \ run_train.py \ --configconfigs/llama7b.yaml \ --train_dataset./data/train.bin \ --run_modetrainworker_num是总卡数local_worker_num是单机卡数。如果是多机每台机器都要起一个 msrun 进程node_rank依次递增master_port指向主节点。启动后第一件事是看日志确认并行策略生效了。日志里会打印每张卡的 rank、分配的模型层数、显存占用情况。如果发现某张卡显存明显偏高说明切分不均匀要回去检查 pipeline_stage 和 layer 的分配。4.3 训练过程中的监控与调优训练跑起来之后要盯几个关键指标loss 曲线、吞吐量tokens/s、显存占用、通信耗时占比。loss 曲线如果震荡厉害先检查学习率是不是太大或者 warmup 步数不够。大模型训练一般需要 2000 步左右的 warmup学习率从 0 线性增加到峰值再余弦衰减。吞吐量上不去大概率是通信瓶颈。可以用 MindSpore 的 profiler 工具抓一下 timeline看 AllReduce 和 AllGather 占了多长时间。如果通信占比超过 30%就要考虑调整并行策略了。显存占用要留 10% 到 20% 的余量跑满容易 OOM。如果显存一直很紧张可以适当降低 micro_batch_num 或者开启重计算。4.4 微调场景的特殊处理微调和预训练的最大区别是数据量小、训练步数少所以更容易过拟合。我一般会做几件事降低学习率到 1e-5 到 2e-5增加 weight_decay用 LoRA 或者 QLoRA 减少可训练参数量。MindSpore Transformers 支持 LoRA 微调配置方式是在 YAML 里加pet_config: pet_type: lora lora_rank: 8 lora_alpha: 16 lora_dropout: 0.1 target_modules: [q_proj, v_proj]LoRA 的好处是可训练参数只有原模型的 1% 左右显存占用大幅下降单卡就能微调 7B 模型。缺点是效果可能不如全量微调特别是在数据分布差异大的场景下。5. 常见问题与排查技巧实录5.1 显存 OOM 的排查路径OOM 是最常见的问题排查要按顺序来。先看是哪个阶段 OOM前向、反向还是优化器更新。前向 OOM 通常是 seq_length 或 batch_size 太大反向 OOM 可能是激活值没释放优化器更新 OOM 则是优化器状态太大。排查工具推荐用npu-smi info实时看显存或者用 MindSpore 的mindspore.ops里的显存分析接口。我一般会在代码里插几个显存打印点定位到具体是哪一层爆的。解决手段按优先级排降 batch_size → 降 seq_length → 开重计算 → 开优化器并行 → 开 offload。前面几个代价小后面几个代价大尽量先用前面的。5.2 通信超时与卡死分布式训练卡死十有八九是通信问题。常见原因有端口被占用、网络不通、某张卡挂了。排查步骤是先用telnet测主节点端口通不通再看日志里最后一条通信记录是哪张卡发的。MindSpore 有个HCCL_EXEC_TIMEOUT环境变量可以设通信超时时间默认好像是 1800 秒。如果网络不稳定可以适当调大但治标不治本根本还是要解决网络问题。还有一个坑是master_port冲突。如果同一台机器上跑多个训练任务端口要错开。我一般用 8118、8119、8120 这样递增。5.3 loss 不下降或出现 NaNloss 不下降的原因很多按概率排学习率太大、数据有问题、并行配置导致梯度错误、初始化有问题。先检查数据把几条样本打印出来看看是不是乱码或者全 padding。然后检查学习率大模型一般从 1e-4 开始试太大就降到 1e-5。如果还不行检查梯度是不是有 NaN用mindspore.ops.isnan查一下。并行配置导致梯度错误的情况比较隐蔽通常是 TP 和 PP 的切分方式跟模型结构不匹配。比如某些层有跨层连接PP 切分时要注意别把连接切断。5.4 常见问题速查表问题现象可能原因排查方法解决手段启动即 OOMbatch_size 或 seq_length 太大看日志确认 OOM 阶段降低 batch_size 或 seq_length训练中途 OOM激活值累积未释放显存监控开启重计算通信超时网络不通或端口冲突telnet 测端口检查网络、换端口loss 震荡学习率太大看 loss 曲线降低学习率、增加 warmuploss 为 NaN梯度爆炸检查梯度加梯度裁剪、降学习率吞吐量低通信瓶颈profiler 抓 timeline调整并行策略某卡显存偏高切分不均匀看各卡显存调整 pipeline_stage5.5 几个我踩过的坑第一个坑是transformers版本冲突。前面提过aimv2 is already used by a transformers config这个报错本质是新版 transformers 加了新配置跟 mindformers 里的配置名撞了。解决办法要么降版本要么手动改配置名。我一般选降版本省事。第二个坑是 checkpoint 保存失败。大模型 checkpoint 动辄几十 GB保存时如果磁盘空间不够或者权限不对会静默失败。建议训练前先确认磁盘空间并且把 checkpoint 路径设成绝对路径。第三个坑是多机训练时时钟不同步。分布式训练依赖时钟同步来做日志对齐和超时判断如果机器之间时间差太大会出现莫名其妙的超时。建议开 NTP 服务。第四个坑是数据加载成为瓶颈。如果数据预处理太慢GPU/NPU 会等数据利用率上不去。解决办法是用多进程加载或者提前把数据转成高效的二进制格式。6. 性能调优的进阶手段6.1 算子融合与图优化MindSpore 的图编译模式会自动做算子融合把多个小算子合并成一个大算子减少 kernel launch 开销。这个在训练大模型时效果很明显特别是 LayerNorm、GeLU 这些逐元素操作。要确保图优化生效得用GRAPH_MODE而不是PYNATIVE_MODE。虽然 PYNATIVE 调试方便但性能差不少。我一般调试用 PYNATIVE正式训练切 GRAPH。import mindspore as ms ms.set_context(modems.GRAPH_MODE, device_targetAscend)6.2 混合精度训练的细节混合精度是标配但细节决定成败。MindSpore 支持自动混合精度AMP也支持手动指定某些层用 FP32。我一般用 AMP 的O2级别大部分算子用 FP16关键算子如 softmax、loss用 FP32。from mindspore import amp model amp.auto_mixed_precision(model, levelO2)要注意的是混合精度下梯度容易溢出需要配合 loss scaling。MindSpore 的 AMP 会自动处理 loss scaling但如果你手动改了什么要确认 scaling 还在生效。6.3 梯度累积与等效 batch size显存不够时梯度累积是提升等效 batch size 的好办法。原理是多次前向反向累积梯度再统一更新参数。这样显存占用不变但等效 batch size 变大了。配置方式是在 YAML 里设gradient_accumulation_steps。比如设成 4就是累积 4 步再更新一次。等效 batch size micro_batch_size × gradient_accumulation_steps × data_parallel。但要注意梯度累积会改变优化器的更新频率学习率可能需要相应调整。一般梯度累积步数增加时学习率可以适当调大一点。6.4 不同硬件配置下的策略选择最后聊聊硬件适配。如果你用的是 8 卡单机TP8 是最优解因为卡间走高速互联通信开销小。如果是 16 卡双机可以考虑 TP8、PP2把 TP 放在机内PP 跨机。如果是 32 卡四机TP8、PP4 或者 TP4、PP4、DP2 都可以试。关键原则是通信量大的并行维度放在机内通信量小的放跨机。TP 通信量最大必须机内PP 通信量中等可以跨机DP 通信量最小跨机无所谓。另外昇腾卡的 HCCS 带宽和 NVLink 有差异具体配置时要参考实际硬件的互联拓扑。可以用npu-smi info -t topo看卡间互联关系据此决定 TP 的切分方式。这套东西说起来复杂但核心逻辑就一条让每张卡都忙起来同时别让任何一张卡成为瓶颈。显存、通信、计算三者平衡好了训练效率自然就上去了。我在实际项目里从 8 卡调到 64 卡最大的体会是配置没有银弹每次换硬件都要重新调一遍。但只要理解了背后的原理调起来就有方向不会瞎试。
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进