先别急着讨论分布式并行怎么配、显存怎么省。我先说一个很多人在本地部署大语言模型时都会遇到的问题:模型权重下载下来才发现,一张卡根本装不下,就算勉强装下,跑一次推理慢到怀疑人生,更别提从头训练或微调了。真正把大语言模型从“能跑”推到“能训、能微调、能上线”,需要解决两件事:一是分布式并行的调度,二是显存优化的每一分抠索。这篇文章就是围绕这套实战记录写的,基于 MindSpore Transformers(也就是 mindformers 工具链)做大语言模型高效预训练与微调的完整过程,涵盖分布式并行策略选型、显存优化手段、脚本参数配置、常见故障排查。适合手里有 4 到 8 张 GPU 或昇腾 NPU 的团队,也适合想从单机快速跑到多机的个人开发者。
1. 项目思路拆解:预训练与微调到底难在哪
1.1 大语言模型训练的三个核心矛盾
先算一笔基础账。一个 7B 参数量的模型,单单把权重用 FP16 存下来就是 14GB。可训练阶段要保存的东西远不止权重:梯度一份、优化器状态两份(Adam 的一阶动量和二阶方差)、每层前向的激活值中间结果一份。按混合精度的常见配置,每个参数在训练过程中平均要占 16 到 20 字节,也就是说 7B 模型全参训练,光权重、梯度、优化器这三类就需要超过 110GB 显存。注意这还没算激活值,激活值会随序列长度、batch 大小和层数快速增长。一张 80GB 的卡根本装不下,两张也悬。这就是第一个矛盾:模型规模超过单卡物理上限。
第二个矛盾是通信和计算。多卡协同训练,每个 step 都要把梯度或中间结果同步给其他卡,通信时间占比会随着并行维度增加而上升。如果只开数据并行,8 卡训练和 1 卡训练相比,单 step 计算量不变,但多了一次全卡梯度 AllReduce,通信一旦成为瓶颈,加速比根本达不到 8 倍。所以并行策略不是越复杂越好,而是要匹配集群的带宽拓扑。
第三个矛盾是工期。预训练一个 7B 模型一般要跑几十万步,微调也要几千到几万步。训练速度慢一天,人力成本、电费、卡租都是真金白银。显存优化表面上省的是显存,实际上省的是 batch size 和训练时长——同样的卡,能跑更大的 batch、更长的序列,就能少走弯路。
1.2 为什么选 mindformers 这套工具链
熟悉 PyTorch 生态的人第一反应可能是 HuggingFace Transformers 加 DeepSpeed。这个组合当然成熟,但如果团队里既有 GPU 又有昇腾 NPU,或者公司规范里明确要求统一到 MindSpore 技术栈,那基于 MindSpore 的 Transformers 库就成了更合理的选择。mindformers 提供的是模型族、训练流程、并行配置、推理部署的一体化封装:模型定义自带分布式能力,配置以 YAML 为中心,预训练和微调共用一套训练器,切换模型和下游任务基本是换配置而不是改代码。
另一个实际理由是调试成本。大语言模型的分布式训练一旦报错,问题可能出在通信组初始化、张量切分、数据读取、优化器状态同步任何一环。mindformers 把模型加载、并行切分、优化器创建这些环节都收敛到 trainer 内部,错误信息的指向通常比手搓多进程代码更容易定位。对很多团队来说,少踩一次框架层面的坑,比省那几分钟训练时间更重要。
1.3 方案选型:省心为主,性能其次
我在选型时给自己定了三条原则:可复现、可观测、可回退。可复现指同一份配置在不同批次环境上必须跑出相同结果,所以锁版本、锁随机种子比盲目升级新特性重要;可观测指一定要有稳定的日志和监控,训练中途 loss 异常能及时发现;可回退指任何优化开关(重计算、offload、混合精度)先在小规模配置上验证再上全量,避免一次花屏式的大改导致整个集群空转。后文所有配置都是沿着这三条原则展开的。
2. 分布式并行:数据、张量、流水线、序列四个维度怎么组合
2.1 数据并行:最直观也最先上
数据并行最简单:每张卡持有一份完整模型,喂不同的数据批次,每轮反向传播结束后做梯度 AllReduce。mindformers 里通过 parallel_config 的 data_parallel 字段控制。它适合模型能塞进单卡、但训练数据量很大的场景,扩展性在同类并行里最好,因为通信量只和模型尺寸有关,和数据量无关。
但纯数据并行有个老问题:梯度同步频率固定,batch 一大就拖慢。后来社区普遍把梯度累积(gradient_accumulation_steps)和数据并行配合使用,用小 batch 算梯度、攒够多个微步再更新一次参数,通信频率下降,训练也稳定一些。实际使用中,我习惯先把 DP 设定到节点内卡数或者节点数的倍数,再往别的复杂度加。DP 是一切的底座,任何其他并行都是建立在这个分组之上的。
2.2 张量并行和流水线并行:为单卡装不下而生
当模型权重本身超过单卡显存,就必须把模型拆开。张量并行(也叫模型并行)是把单层内的矩阵按维度切开,分散到多张卡,前向时通过 AllGather/ReduceScatter 交换中间结果。mindformers 参数里的 model_parallel 控制的是这个切分份数。它的问题是单卡间通信非常密集,每个 Transformer 层都要做两次集合通信,所以只适合卡间带宽高的场景,典型是同一台 8 卡机器内部,跨机跑张量并行基本是灾难。
流水线并行则是按层切分模型:layer 0-7 放卡 0,layer 8-15 放卡 1,前向像流水线一样一节节传,反向再一节节传回来,卡之间传输的是激活值和梯度,而不是频繁的小张量集合通信。它的通信开销小,但对显存的均衡度敏感,需要靠 micro_batch_num 把微批次切细,才能缓解流水线气泡(bubble)时间。pipeline_stage 参数就是流水线的段数,一般设置成节点数或者层数可整除的数。
2.3 序列并行与通信开销的账
序列并行近两年逐渐普及。它的核心观察是:Transformer 里数据并行在 LayerNorm 和 Dropout 这类非张量并行区要额外同步权重,而注意力计算中 token 维度天然可以切分。序列并行把训练时的注意力计算按序列长度切开,配合张量并行使用,能进一步降低单卡激活值峰值。mindformers 在长序列训练场景下,这个开关的价值非常明显:序列长度从 2K 提到 8K 时,如果不做激活值管理,显存会直接翻几倍,而序列并行加激活重计算能把增长斜率压下来。
组合的关键是让“并行度乘积等于总卡数”,还要预留通信优化空间。常见公式是:worker_num = data_parallel × model_parallel × pipeline_stage。例如 8 卡跑 7B,可以配 DP=1、TP=4、PP=2,也可以 DP=2、TP=4、PP=1,差异取决于你的数据量和卡间带宽。开序列并行时,它通常附着在 model_parallel 维度上,不额外占卡数。
2.4 用 7B 模型算一笔账
纸上谈兵没感觉,我拿 7B 模型算过一次。不切并行、FP16 混合精度、序列长度 4096、micro batch 2,优化器状态加激活值轻松突破 120GB,单卡 80GB 直接 OOM。调整成 TP=4、PP=2、DP=1 之后,权重、梯度、优化器状态按模型维度切到 8 卡,每卡大约 15GB,激活值因为张量并行切分降到 30GB 上下,再开激活重计算,每卡峰值压到 50GB 以内,总算能稳定跑。这说明一个问题:不要上来就八卡 DP,先算清显存账,再决定并行组合。
3. 显存优化:把每一块 HBM 都抠出来
3.1 激活重计算:用时间换空间
激活值是训练显存里最容易被忽略的大头。Transformer 每一层的前向都要存下中间激活,反向时才能用来求梯度,几十层累计下来就是巨量显存。激活重计算(activation checkpointing)的思路是:前向时干脆不存中间结果,反向时当场重新算一遍。mindformers 在模型配置里打开 recompute 开关即可,也可以在 select_recompute 里指定只对部分层生效,把时间换空间的损失降到最低。
实测下来,重计算能让激活值显存下降 60% 到 80%,代价是 15% 到 30% 的训练吞吐下降。所以不要无脑全开:如果显存还够,优先开关键层;如果序列长、batch 大,就把重计算和梯度累积组合用。有一个小细节容易被忽略,重计算对象是前向计算,Dropout 的随机 mask 也要重新生成,同一份数据两次前向必须保持随机状态一致,mindformers 内部处理好了这件事,但如果自己改网络,别踩这个坑。
3.2 混合精度:FP16 和 BF16 怎么选
混合精度训练已经是标配。FP16 省显存,但动态范围窄,loss scale 管理不好就容易溢出;BF16 指数位和 FP32 一样,不需要 loss scaling,训练更省心,但尾数精度不足,在部分求和的场景会引入噪声。昇腾上 bf16 的支持这些年已经比较完善,GPU 上 A100 之后也是 bf16 更稳。
我的选择逻辑是:能开 BF16 就开 BF16,尤其预训练阶段;微调阶段如果发现小数据集上精度敏感,再退回 FP16 加 loss scaling。mindformers 的 mixed_precision 字段可以直接切。这里有个经验:混合精度不能只看训练轮数,还要每个 step 观察 loss 是否出现 NaN 或阶梯式跳变,一旦发现就查 loss scale 和数值溢出,别等到第三个 epoch 才发现模型已经毁了。
3.3 优化器状态切分与 CPU Offload
Adam 优化器每个参数要维护两阶动量,加上主权重副本,是训练显存大头。ZeRO 思路是沿数据并行维度把这些状态切分,每张卡只持有自己那份,通信时再做跨卡聚合。mindformers 通过优化器侧的 parallel_optimizer 或 zero 配置开启,效果等于 ZeRO-1/2,能把优化器状态显存除以 DP 卡数。这是目前性价比最高的一项优化,推荐优先配置。
如果显存还是不够,再考虑 CPU Offload:把优化器状态挪到内存,计算时取到卡上,更新完又放回去。它扩展了可训练模型的上限,但会增加 CPU 和 PCIe 的传输开销,训练吞吐会明显打折。我的建议是,Offload 是最后手段,不是第一选择。
3.4 微调场景的 LoRA 与全参微调取舍
微调阶段的显存画像和预训练完全不同。全参微调要维护完整梯度,虽然激活值相比预训练更小,但权重、梯度、优化器状态一样不少,7B 全参微调至少需要 80GB 级别显存。LoRA 把更新量压缩成极小的低秩矩阵,冻结原始权重,训练显存里最大的优化器开销几乎消失,等式变成“冻结权重 + 小学习率 + 低秩适配器”。
实践中,中等数据量的指令微调、领域适配用 LoRA 完全够用;数据量大、任务目标差异大的场景,全参微调上限更高。mindformers 的微调配置里可以选 lora 适配器类型和 target modules,也可以直接全参微调。这个选择不要交给直觉,应该由显存账和任务难度共同决定。
4. 实操过程:从环境准备到跑通一次训练
4.1 环境与版本匹配
我遇到的第一个大坑就是版本不匹配。mindformers 对 MindSpore 版本有明确依赖关系,装错版本会出现算子不兼容、甚至 import 就报错。正确做法是先查官方版本对应表,确定 MindSpore 版本再安装 mindformers,不建议随意装 latest。开发调试时,我习惯在 VS Code 里安装 MindSpore 内核,直接在 Jupyter 里跑小规模配置验证,比反复提交训练任务快得多。
多机场景还需要确认通信库就绪:节点之间要能免密互连,网卡名称、IP 段一致,防火墙放行通信端口。很多分布式问题最后都查出来是网络不通,而不是代码不通。
4.2 预训练配置逐项拆解
以 LLaMA-2 7B 预训练配置为例,核心是 model、parallel、optimizer、trainer 四块。model_config 里最常改的是 seq_length、hidden_size、num_layers;parallel_config 里是 data_parallel、model_parallel、pipeline_stage、micro_batch_num;optimizer 里关注 type 和并行开关;trainer 里是 batch_size、gradient_accumulation_steps、learning_rate 和 checkpoint 间隔。
| 配置项 | 作用 | 常见取值 |
|---|---|---|
| data_parallel | 数据并行度 | 节点数倍数 |
| model_parallel | 张量并行度 | 节点内卡数 |
| pipeline_stage | 流水线段数 | 2 / 4 / 8 |
| micro_batch_num | 流水线微批次数 | 8 ~ 32 |
| seq_length | 序列长度 | 4096 / 8192 |
| recompute | 激活重计算开关 | true / false |
一个容易错的地方是 global batch size 的计算公式:global_batch = micro_batch × micro_batch_num × gradient_accumulation_steps × data_parallel × pipeline_stage。改任何一项,global batch 都会跟着变,直接影响学习率曲线。我专门维护了一张配置参数含义表,记录每项参数的作用和生效条件,避免十天半月后自己都忘了当时为什么这样配。
4.3 微调的关键差异
预训练和微调在 mindformers 里共用 trainer,区别在数据、模型权重和优化器状态。微调要加载预训练 checkpoint,通常还需要把序列长度裁剪到目标长度,减少激活值。数据格式要转成对话模板或指令格式,tokenizer 也要保持一致。
微调阶段最容易被忽略的是学习率与训练轮数。预训练通常会用 warmup 后衰减的大学习率,微调的数据量小,学习率要降一个量级,轮数也不能照搬。我在指令微调时一般把 7B 的学习率设为 1e-5 到 3e-5,轮数 2 到 3 轮,过长反而学坏。如果数据是多任务拼接,还要注意样本长度统一,mindformers 的数据集配置里设好 max_length 和截断策略。
4.4 启动、日志与监控
训练启动用 msrun:
msrun --worker_num=8 --local_worker_num=8 --master_port=8218 --config=configs/llama2/run_llama2_7b.yaml这里 worker_num 要等于总卡数,local_worker_num 等于单机卡数,master_port 各进程保持一致。
启动之后不要只看 loss。我每轮会同步关注四件事:loss 均值与方差、吞吐量(tokens/s)、显存占用曲线、通信等待事件占比。前两个用日志文件统计,后两个用 nvidia-smi 或 MindInsight 查看。loss 曲线平滑下降不代表训练健康,如果吞吐掉了一半,多半是通信或者数据加载出了问题,早发现早止损。
5. 常见问题与排查实录
5.1 通信初始化报错
最常见的启动报错是 NCCL/HCCL 初始化失败,现象是 rank 进程随机卡死,过一会儿超时。先检查 worker_num 和实际卡数是否一致,再检查节点间 TCP 连通性,最后看通信端口的防火墙。如果单机多卡没问题、跨机必挂,99% 是网卡问题:InfiniBand 和 RoCE 配置不一致、网段不通、或默认走 TCP 而非 RDMA。我排查过最诡异的一次,是两台机器的主网卡名称不同,导致环境变量 RDMA 绑定失败,统一网卡命名后立刻恢复。
5.2 OOM 与显存碎片
OOM 不一定代表模型真的放不下。运行时显存碎片化也会导致申请大块连续显存失败。处理顺序:优先开优化器状态切分,再开激活重计算,最后检查数据加载流程是否把不必要的数据放在卡上。另外 MindSpore 的显存统计要看峰值而不只是占用量,峰值有尖刺说明某个环节瞬间申请了超大 Tensor,通常是长序列推理或某个算子实现问题。
5.3 配置注册名的重复冲突
我在改造配置系统时遇到过一个很典型的报错,大意是某个配置名 already used by a transformers config, pick another name。这是模型注册表里出现了同名配置,通常是复制修改配置时忘记改 name 字段,或者 import 了多个定义了相同注册名的模块。排查思路很简单:全局搜配置名,改掉重复定义,同时注意注册名是全局的,不能只在局部文件里重命名了事。我见过的常见诱因是多人协作时各分支都加了同名 config,合并后就撞了。
5.4 训练不收敛与 Loss 异常
loss 如果一开始就 NaN,优先检查学习率是否过大、loss scale 是否溢出、数据里是否有脏样本;如果 loss 一直不变,优先检查梯度是否被 mask 掉或者优化器状态没有正确加载;如果吞吐量稳定、loss 却周期性跳变,往往是梯度累积边界处理错误或者数据 shuffle 范围太小。排查 loss 问题先把 parallel、重计算全部关掉、单卡小 batch 复现,能复现就按上面的方向逐一排除,不能复现就把排查范围缩小到分布式同步逻辑。
5.5 几个容易被绕晕的术语
顺便回答一个新手经常问的问题:生成语言模型和大语言模型是不是一个东西。广义上,生成式语言模型是大语言模型的子集;说大语言模型的时候通常强调规模和通用能力,说生成式模型时强调输出方式是自回归生成。在工程上,这两类模型在很多框架里共用同一套训练流程,所以不必被术语绕住。mindformers 里切换模型族时,关注的是配置文件和权重格式,而不是“模型类型”这个名字本身。
6. 从训练到落地的最后一步
6.1 模型导出与本地部署
训练完不是终点。mindformers 训练产物通常是包含多份分片参数和优化器状态的 checkpoint,导出前需要做权重合并且统一到推理格式。这一步容易踩坑:并行分片的 checkpoint 如果不做合并,直接加载到单卡推理会提示维度不匹配;优化器状态应当剔除,避免权重文件无谓增大。
导出后我一般先在本地做一次单卡推理冒烟测试,确认生成效果和显存占用符合预期。本地部署大语言模型时,我倾向于把序列长度、max batch 等推理参数相对训练配置调小,而不是直接复用训练配置,理由是推理阶段激活值虽然不需要保存梯度,但 KV cache 会随序列长度线性增长,配置不当照样 OOM。
6.2 推理阶段的显存控制
推理显存由权重、KV cache、计算中间态三部分组成。权重大头可以用量化缩小,KV cache 要靠控制并发数和 max_seq_len 来管理。实践中,同一个 7B 模型,FP16 推理权重约 14GB,量化到 INT8 再减半;如果一次服务要求高并发长序列,就必须做 KV cache 的显存预留计算,而不是靠感觉分配。训练阶段抠出来的显存经验,在推理侧完全复用得上,这也是我一直建议先把训练账算明白的原因。
6.3 后续还可以扩展的方向
这套流程跑通以后,可以往三个方向扩展:第一是多模态,把视觉编码器和大语言模型桥接起来,现在很多工作在做跨模态对齐;第二是长序列训练,靠序列并行加高效注意力进一步拉长上下文;第三是稀疏化与量化训练,把训练阶段的低精度经验反哺到推理侧的极低比特量化。每往一个方向走,核心还是这篇文章里那套账:并行维度怎么组合、显存从哪里省。
我个人实际操作中最深的一条体会是:分布式训练排错的第一原则是先复现、后定位。不要在一个 8 卡集群上开着调试打日志,那会把人折磨疯。把并行度全部降为 1,单卡把网络结构跑通,再逐步加并行度和优化开关,每一步都跑一个极小的 smoke test,确认改动生效再继续。这个过程看着慢,实际省下来的时间远超预期。显存优化的本质是在吞吐、稳定性和显存之间做权衡,没有一个开关是白开的,也不存在银弹。希望这套实战记录能给正在折腾大语言模型训练的人一点直接的参考。