ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

MindSpore分布式训练实践:从数据并行到混合并行完整指南

MindSpore分布式训练实践:从数据并行到混合并行完整指南 既然是因为MindSpore 1.10版本之后分布式能力变化太大那这篇就把从单卡脚本迁移到数据并行、再升级到混合并行的完整路线、调参细节、踩坑记录一次说透给准备上手分布式训练的同学一份能直接照着做的实操笔记。1. 分布式并行选型先想清楚再动手1.1 为什么要上分布式以及数据并行的适用边界先聊一个扎心的事实很多人上分布式训练不是因为它“听起来高大上”而是单卡实在跑不动了。具体来说就两种情形一种是显存不够模型加载就OOM一种是算力不够一个epoch跑到天荒地老。数据并行解决的是第二种问题把一份大batch拆成多份分给多张卡每张卡持有完整模型副本各自算完梯度后做一次AllReduce同步然后更新参数。但数据并行有个明显的“斤斤计较”点通信开销。假设8张卡跑1万张图片的训练集batch_size256那么每张卡分到32张图。每轮迭代都要做一次全量梯度的AllReduce数据量等于模型参数量乘以4字节FP32。如果是BERT这样3亿参数量的模型一次AllReduce就要传输约1.2GB数据。8卡之间做这个同步PCIe和NVLink带宽好的时候还行带宽差一点的机器通信时间甚至能占到总迭代时间的40%以上。所以数据并行不是万能的它适合模型本身能塞进单卡、只是数据集太大训练太慢的场景。如果模型本身就放不进单卡数据并行就彻底无效了。因为每卡都要一份完整模型你卡放不下所有卡都放不下。这时候必须上混合并行——把模型本身切开来每个设备只存一部分这个思路跟数据并行的“全员完整副本”有本质区别。1.2 混合并行的本质把“大”拆成“小”混合并行也叫模型并行的一种工程化综合体的核心是“切分”把网络结构、算子、优化器状态、梯度都切到多张卡上每张卡只负责一部分计算。MindSpore 1.10把这件事用三种方式落地算子级并行、流水线并行、优化器并行。算子级并行是最细粒度的切分。拿矩阵乘法来说$Y X \times W$X的维度是[Batch, 768]W是[768, 768]这俩都可以沿着某个维度切。比如把W切成[384, 768]和[384, 768]两块分别放在两张卡上每张卡算出一半的输出再拼起来。这就是最简单的算子切分。流水线并行则是按层切Transformer的前几层放在卡0中间几层放卡1后几层放卡2。数据像流水线一样依次流过各卡。优化器并行针对的是Adam这类优化器保存两份状态m和v导致显存翻三倍的问题参数一份、m一份、v一份把这些状态也切分开每卡只维护一部分。这三种切分组合在一起才真正解决了“模型大到一张卡放不下”的问题。但组合的代价是通信模式更复杂算子级并行需要all-gather和reduce-scatter频繁交换中间结果流水线并行需要点对点通信传递embedding输出和梯度优化器并行还需要在更新参数前做一次全量汇聚。通信模式多了“踩坑点”也就指数级增多。1.3 选型决策什么时候用哪种方案直接给一张我实际项目里用的决策表不搞虚的场景模型能否装入单卡推荐方案通信开销级别配置复杂度小模型 海量数据能数据并行每轮一次AllReduce低大模型 数据量大不能混合并行算子流水线优化器多阶段多次通信高大模型 微调任务不能/勉强混合并行或重计算优先中中小数据 大模型不能混合并行数据并行部分甚至可以砍掉中高这里多说一句很多人一上来就奔着混合并行去觉得“更高级”。实际上如果模型能单卡放下数据并行的性价比远高于混合并行——代码改动量小、调参难度低、性能好预测。我在生产环境里见过最离谱的事故是某个团队拿着一个只有1.3亿参数量的BERT-tiny模型硬上8卡混合并行结果因为频繁的切分通信训练速度还不如单卡。工具是把双刃剑先算清楚账再动手别为了炫技而上复杂度。2. 核心API与关键配置MindSpore 1.10 分布式开发的地基2.1 初始化与上下文配置的正确姿势MindSpore分布式开发第一行要做的不是别的是初始化通信组。代码长这样import mindspore as ms from mindspore import context from mindspore.communication import init, get_rank, get_group_size # 必须在init之前配置 context.set_context(device_targetGPU) # 或 Ascend context.set_context(modecontext.GRAPH_MODE) # 分布式必须图模式 # 初始化通信组 init() rank_id get_rank() group_size get_group_size() print(f当前进程 rank{rank_id}, 总设备数{group_size})这里有几个细节值得说透。第一init()必须放在set_context之后否则某些后端会找不到设备上下文直接报错。第二分布式训练必须使用GRAPH_MODEPyNative模式下有些集合通信算子没办法正常执行这是MindSpore目前的实现限制别在这种地方硬杠。重点说set_auto_parallel_context这是分布式训练的“总开关”from mindspore.context import ParallelMode context.set_auto_parallel_context( parallel_modeParallelMode.DATA_PARALLEL, # 这里还能选 SEMI_AUTO_PARALLEL 或 AUTO_PARALLEL device_num8, gradients_meanTrue, # 梯度AllReduce后是否取平均多卡训练一般True parameter_broadcastTrue # 初始化时广播参数保证每卡起点一致 )gradients_meanTrue这个参数很多人不在意但实际影响很大。如果设成FalseAllReduce默认是求和那么等效batch_size直接翻了8倍学习率不跟着相应调大loss曲线很容易飞。设成True则是取平均等效batch_size不变更符合大多数人的心理预期。还有一个极容易被坑的device_num必须与实际参与训练的卡数一致多写了或者少写了部分后端在通信组初始化时会报错或者更恶心地——不报错但训练结果异常。实际踩过的场景是机器有8卡但某次只分配到了4卡代码里还写着device_num8训练前几十个step正常后面loss直接变NaN。排查了半天才发现是通信组大小和实际设备不匹配。2.2 数据并行的数据集切分不患寡而患不均数据并行的核心是“每张卡吃不同的数据”。MindSpore的数据集API天然支持分片关键是配置别出错。我常用的写法import mindspore.dataset as ds # 每个rank各取一份互不重叠的数据 dataset ds.ImageFolderDataset(data_path, num_shardsget_group_size(), shard_idget_rank(), shuffleTrue) dataset dataset.batch(batch_size, drop_remainderTrue)num_shards和shard_id必须成对出现含义就是“把数据集切成group_size份我是第rank_id份”。忘了加这俩参数你会看到每张卡都在吃同一份完整数据——训练loss毫无问题甚至还能正常下降但等效batch_size实际上没变训练速度没有本质提升因为每卡计算量没变好几天的算力就白烧了。另一个隐蔽问题drop_remainderTrue必须加上。分布式训练里如果最后一个batch的样本数量在各卡上不一致AllReduce的梯度形状对不上就会直接崩溃。之前遇到过一次训练到第999个step总共1000步突然报 shape 不匹配查了半天就是数据集的最后一个batch剩了17张图8卡分不均。drop掉反而最稳。数据加载还有一层容易被忽略的性能瓶颈num_parallel_workers。在我实际调优的项目里很多“训练慢”的case最终定位到的是数据加载线程数配置太低GPU/昇腾在等数据而不是在算模型。默认的并行worker数量往往只有4在大数据集、高分辨率图像或复杂预处理流水线下完全不够用。建议把这个参数调到816之间并配合num_parallel_workers和python_multiprocessingTrue一起用。如果发现CPU核数还有余量但训练上不去优先怀疑数据管道。2.3 混合并行的策略配置从粗到细的切分过程MindSpore 1.10里做混合并行通常用ParallelMode.SEMI_AUTO_PARALLEL或者ParallelMode.AUTO_PARALLEL。这俩名字看着相似实际用法完全不同SEMI允许你手动指定每一层的切分策略AUTO则交给框架的规划算法全自动搜索。在模型结构比较复杂或者你清楚知道哪里通信瓶颈的时候SEMI_AUTO更可控也是我日常项目里用得最多的模式。算子级切分的核心接口是shard()。下面这段代码演示了如何把一个Dense层按输入维度切到8张卡上from mindspore import nn, Tensor import numpy as np class MyNet(nn.Cell): def __init__(self): super().__init__() self.fc nn.Dense(in_channels768, out_channels3072) # 设置该算子的切分策略 self.fc.shard(strategy((1, 1), (8, 1), (1,)),)解释一下strategy的含义((1, 1), (8, 1), (1,))对应三个张量——输入、权重、偏置的切分方式。输入保持完整1表示不切权重按第一个维度切成8份每卡持有一部分偏置不切。这样每张卡只做了一个[768, 384]的矩阵乘法算完后通过通信拼回[768, 3072]的完整输出。流水线并行则是另一种玩法。给Cell设置stageMindSpore会把不同stage的计算调度到不同的设备上执行class PipelineNet(nn.Cell): def __init__(self): super().__init__() self.embedding nn.Embedding(vocab_size30000, embedding_size768) self.layer1 TransformerLayer(hidden_size768) self.layer2 TransformerLayer(hidden_size768) self.layer1.pipeline_stage 0 self.layer2.pipeline_stage 1 self.classifier nn.Dense(768, 1000).to_float(ms.float16)设备编号从0开始stage0的层在卡0/卡1如果组内2卡上算stage1的层在后几张卡上算。数据先过stage0算出来的中间激活通过通信传给stage1继续算。这里最关键的配置是batch大小必须能被stage数量整除不然最后一个micro batch会被卡在边界上报pipeline相关错误。2.4 重计算和混合精度白拿的显存优化混合并行的显存压力很大因为每卡的中间激活可能存在多个stage。MindSpore提供重计算(recompute)机制把部分前向激活缓存丢到反向重算一遍。这是一个典型的“空间换时间”显存能省不少但训练时间会增加。实际经验是只对显存消耗排名靠前的几个Transformer层开重计算而不是全量开收益最高。for layer in self.layers: layer.recompute()混合精度也是混合并行里必开的选项。MindSpore里一句话搞定from mindspore import dtype as mstype model.to_float(mstype.float16) # 大部分算子切到FP16但fp16有个经典的坑loss scale。如果不开动态loss scale小梯度会被直接清零训练根本走不动。MindSpore的Model训练接口默认会用DynamicLossScaleManager管理但如果手动写了训练循环就必须自己维护loss scale这是很多自定义训练脚本写崩的根源之一。3. 实操过程从单卡脚本迁移到多卡训练完整路线3.1 迁移前准备先给自己的代码做个体检不要一上来就改并行配置先花半小时确认几件事能省下后面两天的排查时间。先把单卡脚本跑通输出loss正常。这一步确保模型本身没有问题不然分布式环境下各种通信错误会把你的排查方向带偏。然后检查代码里有没有全局状态——比如用Python全局变量缓存了某个中间结果或者以文件形式保存了临时数据。分布式训练是多进程独立运行进程之间除了集合通信函数没有任何隐式的数据共享。之前见过一段代码把数据预处理结果缓存到/tmp目录命名只有时间戳8个进程同时写后启动的进程把先启动的进程预处理的文件直接覆盖了最后所有卡数据一模一样训练结果诡异到没法用。这类bug在单卡上根本不会出多卡环境一跑就现形。还有随机种子的问题。每张卡的初始化参数必须一致否则不同进程的模型起点不同步AllReduce梯度就没有意义。因此在模型初始化之后一定要设置好全局随机种子并建议开启parameter_broadcastTrue让rank0的参数在初始化后广播给其他节点双保险。3.2 数据并行实战8卡跑起来的完整示例直接看一个基于ResNet-50在ImageNet做数据并行的完整训练代码这是我实际项目里精简过的版本import mindspore as ms import mindspore.dataset as ds import mindspore.dataset.vision as vision import mindspore.dataset.transforms as C from mindspore import nn, context, Model from mindspore.context import ParallelMode from mindspore.communication import init, get_rank, get_group_size from mindspore.train.callback import LossMonitor, ModelCheckpoint, CheckpointConfig # 1. 并行上下文 context.set_context(device_targetGPU) context.set_context(modecontext.GRAPH_MODE) init() context.set_auto_parallel_context( parallel_modeParallelMode.DATA_PARALLEL, device_numget_group_size(), gradients_meanTrue, parameter_broadcastTrue ) rank_id get_rank() # 2. 数据集分片 def create_dataset(data_path, batch_size32): dataset ds.ImageFolderDataset(data_path, num_shardsget_group_size(), shard_idrank_id, shuffleTrue) image_ops [ vision.Decode(), vision.Resize((224, 224)), vision.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), vision.HWC2CHW() ] dataset dataset.map(operationsimage_ops, num_parallel_workers8) dataset dataset.batch(batch_size, drop_remainderTrue) return dataset # 3. 网络、优化器和Model封装 network resnet50(num_classes1000) loss_fn nn.CrossEntropyLoss() optimizer nn.Momentum(network.trainable_params(), learning_rate0.01, momentum0.9) model Model(network, loss_fn, optimizer) # 4. 回调每个rank只保存自己的checkpoint推荐rank0保存或统一汇总 ckpt_config CheckpointConfig(save_checkpoint_steps1000, keep_checkpoint_max5) ckpoint_cb ModelCheckpoint(prefixresnet50_8p, directoryf./ckpt_rank_{rank_id}, configckpt_config) # 5. 训练 model.train(epochs90, train_datasetcreate_dataset(./data), callbacks[LossMonitor(100), ckpoint_cb])启动方式用的是MindSpore推荐的mpirunmpirun -n 8 python train_resnet50.py跑起来之后重点看每张卡的日志里rank_id是不是0到7都出现了并且各卡loss数值接近。如果某张卡的loss明显偏离优先怀疑数据集shard分配有问题参考上面2.2的思路排查。这一步顺利的话数据并行基本就通了。训练中途用nvidia-smi看一眼GPU利用率。如果利用率低于90%说明通信或者数据加载有瓶颈去调num_parallel_workers或者检查网络带宽。如果利用率稳定在95%以上数据并行这part算真正达标了。3.3 混合并行实战当一个模型放不进单卡时假设我们要在一个12GB显存的卡上训练一个5亿参数量的GPT类模型。FP16下这种量级的模型光参数和优化器状态就要吃掉10GB以上加上中间激活肯定爆显存。这就必须上混合并行。我用一个简化版本展示关键配置思路import mindspore as ms from mindspore import context from mindspore.context import ParallelMode from mindspore.communication import init context.set_context(device_targetGPU, modecontext.GRAPH_MODE) init() context.set_auto_parallel_context( parallel_modeParallelMode.SEMI_AUTO_PARALLEL, device_num8, gradients_meanTrue, parameter_broadcastTrue, pipeline_stages2, # 启用2级流水线 optimizer_parallelTrue # 开启优化器并行 )这里的关键设置是pipeline_stages2它会把网络切成2个stage默认前一半层在设备0-3上算后一半在设备4-7上算。optimizer_parallelTrue则意味着Adam的m/v状态也被切分每卡只保留自己负责的那部分参数的优化器状态。在模型定义层面主要靠shard()和pipeline_stage两个关键词来切。实际项目里我会对最重的注意力SDPA矩阵乘设置切分策略class Attention(nn.Cell): def __init__(self, hidden_size, num_heads): super().__init__() self.num_heads num_heads self.hidden_size hidden_size # 三个QKV投影矩阵 self.q_proj nn.Dense(hidden_size, hidden_size) self.k_proj nn.Dense(hidden_size, hidden_size) self.v_proj nn.Dense(hidden_size, hidden_size) self.out_proj nn.Dense(hidden_size, hidden_size) # 把每个投影矩阵的权重按列切分到8卡上 self.q_proj.shard(strategy((1, 1), (1, 8), (1,))) self.k_proj.shard(strategy((1, 1), (1, 8), (1,))) self.v_proj.shard(strategy((1, 1), (1, 8), (1,))) self.out_proj.shard(strategy((8, 1), (1, 1), (1,)))((1, 1), (1, 8), (1,))表示权重矩阵的第2维切成8份——也就是每一卡只负责输出特征中的1/8。最终的out_proj则是把8份结果在特征维度拼回来。这套策略下来单卡显存占用从“放不下”变成“放得下三分之二”再加上重计算基本能跑。不过这里必须提醒混合并行不像数据并行配好了就完事。策略设计不好通信开销可能比省下来的计算量还大。经验法则是切分维度的通信数据量越小越好。比如Q/K/V的投影切输出维度每个token的中间输出要all-gather 768维的数据这个开销相对可控但如果把序列长度也切了中间通信量直接翻数倍性能会血崩。3.4 性能调优实战从50分提到90分的完整清单数据并行和混合并行都跑通之后就该谈性能了。我总结了一个调优顺序按优先级来不要乱。第一优先级是通信拓扑。MindSpore底层在GPU上用NCCL、在昇腾上用HCCL默认的通信拓扑是Ring AllReduce也就是所有设备组成一个环两两通信各传一部分数据。Ring在8卡以下表现稳定但到16卡、32卡多级Ring可能会导致延迟线性增加。这时候要检查NCCL的环境变量比如NCCL_IB_DISABLE是否误设为1导致走了慢速网络路径还有NCCL_DEBUGINFO能帮你看到实际走的通信拓扑和耗时。曾经遇到过一台服务器换了网卡驱动后NCCL回退到了TCP socket通信训练速度直接砍了80%查日志才发现是IB没识别到。第二优先级是通信和计算重叠。MindSpore的分布式中梯度AllReduce是同步的——每个step要等8卡都把梯度算完才开始通信。想隐藏这个延迟可以用梯度累积和微batch流水。MindSpore 1.10支持把一个小step内的梯度先做本地累积攒够多个微batch再统一做AllReduce。这样通信次数直接除以累积次数吞吐提升非常明显代价是模型收敛速度会略微变化需要稍微调整学习率。第三优先级是消除数据加载瓶颈。这块思路和MySQL之类的数据库调优本质一样先看瓶颈在IO还是计算再决定怎么优化。如果GPU/NPU利用率上不去先起个profiler看timeline确认训练step里哪个环节耗时最长。我之前遇到过一个项目数据里有大量小文件每个step要读几百个随机小文件IO延时把训练拖垮了。后来先做了数据预读和批量打包处理训练速度直接翻倍。调数据集并行worker数也属于这个环节。第四优先级才是混合精度和padding优化。FP16的算子加速明显但要注意照着一份loss scale配置来。还有模型内的动态shape问题比如GPT的padding mask导致序列长度变化MindSpore图模式对此支持得不算好推荐把序列长度固定成统一长度牺牲一点算力换稳定性。4. 踩坑实录与问题排查手册4.1 高频报错速查表报错现象根本原因解决办法HCCL/NCCL通信超时设备间网络不通或SecurityGroup拦截了IB端口检查{NCCL_DEBUGINFO}日志确认走IB/RoCE路径shape mismatch at allreduce数据集batch大小各卡不一致数据集加{drop_remainderTrue}训练中期loss为NaN梯度累积时loss scale太小或梯度爆炸开动态loss scale或调低学习率初始参数不一致导致通信出问题未设置随机种子{parameter_broadcast}未开迁移前统一种子开启parameter_broadcast某rank退出、其他卡hang住单卡OOM或者该rank数据异常定位问题rank日志先解决OOM再整体重启混合并行时显存不降反升中间激活被重复复制导致通信buffer过大开recompute并降低batch_size这张表虽然只有七行但覆盖了我在项目里见过的绝大多数分布式排查需求。遇到问题先对照这个表定位方向比盲目看日志高效得多。4.2 我亲身踩过的三个深坑第一个坑是数据集分片不对loss曲线看起来居然也正常。那一次跑的是8卡数据并行因为代码里误把num_shards8写成了num_shards1所有卡都在读同一份数据。按理说这会造成严重的过拟合但因为数据集本身很大、shuffle打乱了顺序前1000步的loss下降曲线跟正常训练几乎一模一样。直到我对比了8张卡的loss发现完全一致正常应该略有差异才意识到问题。这个教训说明光看loss曲线不够要看每卡的loss是否“正常地不完全一致”。第二个坑是通信超时。训练到第5000步左右整体速度突然掉了一半日志里不断出现NCCL timeout相关报错。查了一整天最后发现是宿主机防火墙规则更新把NCCL用到的TCP端口堵了但IB路径没走通导致它悄悄回退到了TCP速度自然就慢了。解决方式很简单确认IB/RoCE链路通畅并让NCCL优先使用IB如果实在无法走IB就明确把TCP端口放开避免来回回退。第三个坑是混合并行的显存爆炸问题。当时我们信心满满地把GPT模型切成4卡跑起来却直接OOM。看显存占用发现每卡的峰值显存比理论估算高了近40%原因是中间激活没有触发重计算而且其中某个线性层因为切分策略设置成了(8, 1)使得通信缓存buffer保留了一整份完整中间结果叠加上去了。后来给每个Transformer层开了recompute并且把通信密集的算子切分策略调整到(1, 1)维持完整计算、只在权重侧切显存一下子降了30%以上。4.3 调试工具箱日志、环境变量与VSCode远程排查分布式训练的调试难度比单卡高的多我的经验是“先看日志、再看拓扑、后改代码”。日志层面开NCCL_DEBUGINFO能输出每个通信算子的执行细节这在排查hang住和timeout问题时几乎是必备手段。同时建议每个rank进程把日志写到独立文件避免8个进程向同一个stdout输出导致日志互相穿插根本没法看。环境变量里值得注意的还有ASCEND_GLOBAL_LOG_LEVEL昇腾环境和GLOG_vMindSpore通用把日志级别调到INFO可以看到模型编译和资源分配信息。当年排查一个初始化失败问题就是靠GLOG日志发现某个后端没有正确加载才定位到是动态库缺失。远程调试的话我推荐直接用VSCode的Remote SSH插件连到训练服务器上看日志和代码比在终端里用tail -f舒服很多。MindSpore本身也提供了一个轻量级的命令行工具查看模型编译信息配合VSCode的静态检查能省很多事。具体到分布式场景真正值得推荐的调试姿势是先用1~2卡启动小规模训练确认单卡脚本逻辑没问题再逐步扩到8卡甚至更多。一次从1卡直接跳到32卡的教训是惨痛的——出问题时根本分不清是脚本问题还是扩卡引入的通信问题。5. 最后再聊几句实际的MindSpore 1.10的分布式能力确实够用了但它不是一个“开箱即用”的框架需要你在并行模式、切分策略、通信优化上都动脑。我个人做下来的体感是数据并行属于“必须掌握的基础技能”混合并行则更像一门手艺活需要针对模型形状、硬件拓扑反复试验。踩过这些坑之后我养成了两个习惯推荐给你。第一每次跑分布式训练前先花10分钟确认数据集切分方式和通信环境这两块出了问题往往要烧掉几天时间。第二任何并行方案都从最小规模验证开始2卡跑通了再推到8卡、16卡每一步都记录日志。多卡环境的不确定性比单卡大得多但只要你按部就班来踩坑的数量是可控的。这篇文章里的所有配置和案例都来自我实际跑过的项目不一定适用你遇到的所有场景但思路和排查路径是可以复用的。希望你能省下我当初踩坑的时间。
返回列表