资讯动态

大模型并行与分布式训练实践:从并行范式到多机调优全攻略

发布时间:2026/9/5 1:55:14 来源:尧图企业网站定制
说实话第一次系统看完CS336 Lecture 8的讲义时我最大的感受不是“原来有这么多并行方法”而是“之前吭哧吭哧调了那么久的多卡训练居然全是在凭感觉试错”。我当时把单机训练脚本改到多机整整卡了一宿NCCL初始化超时、GPU利用率上不去、loss震荡到怀疑人生最后才发现问题出在最基础的通信配置上。所以这次把Stanford CS336第八讲“并行与分布式实践”连同我自己的实测记录和踩坑过程一起整理出来希望能让准备迈进分布式训练大门的朋友们少走一点弯路。这篇内容不只是把讲义里的概念翻译一遍我会把并行范式之间的区别、集合通信的底层逻辑、从单机到多机的部署细节以及分布式场景里常遇到的各类工程问题都串起来讲清楚。它适合正准备做多卡训练、需要搞懂分布式训练原理或者已经在跑分布式任务但总觉得哪里不对劲的开发者和算法工程师参考。1. 为什么要系统学并行先从单卡的天花板说起1.1 模型变大之后显存和算力都不够用了很多人第一次接触并行训练是因为一张卡跑不动了。这个“跑不动”其实有两层含义一是显存放不下二是算力跟不上。以现在主流的大语言模型为例一个参数量在7B左右的模型光参数用FP16存就需要大约14GB。但训练过程不是只存参数这么简单还要存梯度、优化器状态以及前向传播过程中的中间激活值。如果用Adam优化器每个参数在训练中会额外占用fp32的动量项和方差项这个开销是8字节再加上梯度本身。算下来一个7B模型光优化器状态就已经轻松超过单张A100 80GB的显存上限了。所以就算你不关心训练速度单纯想把这个模型塞进显存多卡就已经不是可选项而是必选项。算力也是一样。单卡训练一个7B模型哪怕只跑1个token的batch计算量都高得吓人。H100的理论算力大约是989 TFLOPSFP16稀疏场景更高但单卡训练大模型的理论吞吐离实用吞吐还是差得很远。因为训练是串行的一个step算完才能算下一个step你没法靠一张卡把几周的墙钟时间压缩成几天唯一的办法就是让多张卡一起算。1.2 “并行”不是简单把活分出去我见过很多刚开始接触分布式的同学有一个非常天真的想法既然一张卡放不下那我把它拆成两半放两张卡速度应该也能翻倍吧现实远没这么简单。当模型被拆到多张卡上之后卡与卡之间需要在每个step内频繁同步梯度和中间结果这种通信会产生巨大的开销。如果并行策略设计不好经常会看到8张卡跑起来单卡算力利用率还不到50%甚至出现“越加卡越慢”的情况。根本原因不是卡不行而是通信变成了瓶颈——数据在卡之间传输的时间远远大于真正计算的时间。所以学习并行与分布式实践的第一步是建立一个观念分布式系统里的博弈从来都在“计算”和“通信”之间。你做的每一个并行决策本质上都在做一道取舍题用多少通信换计算用多少重复存储换少通信。1.3 这门课教的核心思路CS336这门课整体的风格偏系统会把LLM从数据准备、模型结构、训练优化到部署讲得很扎实。到了第八讲主题非常明确在真实的训练集群上如何把一个大模型高效地跑起来。它不会只给你一个高层概念而是细化到每层网络传输什么、每个并行范式里通信量和显存占用如何计算、激活值怎么在设备之间流转这些层度。这才是做分布式训练最有价值的部分。因为只有把这些细节理解透了你在跑崩了的时候才知道查哪里在调参的时候才知道改哪里。我整体的建议是先搞懂几个并行的基本范式理解它们各自解决什么问题、代价是什么、适合什么场景然后再上手框架提供的分布式原语自己动手去配置一次多卡训练最后才是根据实测结果做调优。2. 主流并行范式拆解数据并行、张量并行与流水线并行2.1 数据并行最简单也最常用数据并行Data Parallelism的思路一句话就可以说清楚每张卡都放一份完整的模型副本然后把训练数据切成很多份每张卡分到不同的数据去做前向和反向计算最后把所有卡算出来的梯度做一次全局同步再拿同步后的梯度去更新每一张卡上的模型。举个例子假设我有8张卡、一个batch size等于64的数据集在数据并行下每张卡会单独处理8条样本即micro-batch size8。每个step结束后8张卡各自得到一份梯度这些梯度要做一次AllReduce操作得到所有梯度的平均值然后每张卡用这个平均梯度更新自己的模型副本。为什么这个方案能跑起来而且很稳定因为模型各卡完全相同每张卡看到的数据又不一样同步梯度等价于用更大的batch做了训练。这也是数据并行和“在线学习”的本质区别。但它的短板也很明显每张卡都要存完整的模型、梯度和优化器状态当模型大到超过单卡显存容量时数据并行直接失效。适合数据并行的场景是模型单卡勉强能塞下但你想通过增大batch来提升训练吞吐。这是最常见也最容易起步的方案。像PyTorch DDPDistributedDataParallel底层就是这么做的只是它在同步梯度的通信效率上做了不少优化。2.2 张量并行把模型横向切开当模型大到单卡放不下数据并行就不够用了这时候需要的是把模型本身切开。张量并行Tensor Parallelism的做法是针对某一层的权重矩阵做切分让不同的GPU分别计算同一层输出的不同部分再通过通信把结果拼起来。拿Transformer里的Self-Attention来举例。Q、K、V三个矩阵是由输入乘以三个权重矩阵得到的。张量并行里可以把权重矩阵按列切成两份分别放在两张GPU上每张卡各自完成部分头部的计算最后再做一次AllReduce把结果合并。MLP层也有类似的做法比如把第一个线性层的权重按列切、第二个线性层按行切这样计算过程就均匀分布在多张卡上。张量并行的优势是能真正把单层参数拆开让超大模型塞进多卡。缺点是每一层的前向和反向都需要一次甚至多次通信通信频率极高。所以它基本都是用在单机多卡场景靠NVLink这类高速卡间互联来降低通信延迟。如果跨机做张量并行网络往返延迟会直接拖垮训练速度。2.3 流水线并行把层纵向切段流水线并行Pipeline Parallelism的角度又不一样。它是把模型按层切成多个阶段比如一个24层的Transformer切成4段每段分配到不同的GPU上。数据像流水线一样先经过第1段所在GPU算完传给第2段再传给第3段最后在第4段完成前向。但是这个方案会暴露一个经典问题bubble气泡。想象一下如果只有一条数据在流水线上流动那么任何时刻只有一段GPU在计算其余都是空闲的。为了填满流水线我们需要把batch切细成多个micro-batch让不同micro-batch错峰在不同段上计算这样多个GPU才能同时忙起来。即便如此流水线在启动和排空阶段仍然会有一定比例的空闲这就是气泡开销。流水线并行最大的价值是解决超大模型跨机部署问题。因为机器之间的通信带宽远不如卡间NVLink所以把通信频率控制在“每传播一个micro-batch才传一次”的流水线方案比张量并行更适合跨机场景。2.4 三种范式怎么选说实话真正生产级的训练很少只用一种并行而是几种并行组合使用。组合的思路通常是这样机器内部多卡用张量并行切分超大层机器之间用流水线并行把模型的层切成多段分散到不同机器在整套模型拓扑之上再做数据并行来扩大总体batch和处理更多数据。这种组合其实符合一个朴素的工程原则通信代价昂贵的层面少通信通信快的层面多通信。单机内的NVLink带宽很高适合频繁通信的张量并行跨机的网络延迟高带宽低适合通信频率较低的流水线并行和数据并行。我把三种常见范式整理了一下方便照着选型并行范式切分对象通信开销适用规模主要代价数据并行训练数据按batch切每step一次梯度同步通信量随模型增大而增大单卡能放下的中大型模型每卡完整副本显存冗余多张量并行层内权重矩阵按行/列切每层多次AllReduce非常频繁单卡放不下且单机内多卡场景通信依赖卡间高速互联不适合跨机流水线并行模型按层切成多个阶段每个micro-batch传递一次激活和梯度较稀疏超大模型跨机部署存在流水线气泡调micro-batch平衡麻烦还有一类技术叫序列并行Sequence Parallelism主要用来省显存因为Transformer处理长序列时中间激活值特别占空间。它的做法是把序列维度也切分到不同设备上在LayerNorm和Dropout这类“按token独立操作”的部分并行计算。下面第3节会进一步提到这套思想和ZeRO系列有重叠但也有本质差异。其实最值得推荐的实操路径是在单机上先跑通DDP再去理解张量并行的切分逻辑最后才上流水线。因为DDP实现最简单、调试相对容易是建立分布式直觉最快的方式。3. 集合通信与分布式框架的底层逻辑3.1 别绕过NCCL它是分布式训练的地基现在做GPU分布式训练几乎绕不开NCCLNVIDIA Collective Communications Library。无论是PyTorch DDP还是DeepSpeed、Megatron底层多卡通信基本都是通过NCCL来执行的。NCCL提供了一系列集合通信原语比如AllReduce、Broadcast、AllGather、ReduceScatter等。最核心的一个概念就是AllReduce——把多张卡上各自的张量规约成一个结果比如求和或求平均然后把这个结果广播给所有卡。数据并行里同步梯度用的就是这个操作。AllReduce最简单的实现是“主节点聚合再分发”但这样主节点会成为瓶颈。NCCL实际采用的多是Ring AllReduce把参与通信的GPU首尾相连成一个环每个GPU只和相邻GPU通信通过两个阶段reduce-scatter all-gather完成全量规约。拿传梯度的场景打个比方就好比班上8个同学各自写了一部分答案要把8份答案合并成完整版本给每个人都发一份。如果所有版本都先交给班长汇总再由班长分发班长会被累死但如果是环形传递每个同学只和前后同桌各传一次经过两轮就能让所有同学手里都拿到完整合并结果而且每个同学的工作量非常均衡。这也是为什么Ring AllReduce在卡数较多时扩展性更好。实操中需要记住的是NCCL的通信是异步的它以独立的通信算子形式提交到GPU流上如果不做同步你观察到的通信时间可能并不真实。训练脚本里的torch.cuda.synchronize不只是为了算时间更是为了确保你测到的耗时是真实耗时。3.2 PyTorch DDP为什么是默认起点PyTorch的DistributedDataParallel是所有分布式训练的起点。很多人刚接触时可能会用DataParallelDP但我建议直接放弃DP转DDP。DP是单进程多线程模型受Python GIL限制且主卡通信压力巨大几乎没法扩展到多机。DDP是真正的多进程模型每个GPU对应一个进程通信路径更合理扩展性远好于DP。DDP的底层实现里有一个很关键的细节它并不是在每个step里对所有梯度做一次全量AllReduce而是利用梯度计算顺序在每个参数的梯度算好后就立刻启动通信让通信和反向传播计算重叠起来。这就是为什么你在训练日志里看到单个step的时间往往比“反向传播时间AllReduce时间”要小很多的原因——通信被隐藏到计算里了。用DDP启动训练时主要入口是torch.distributed.init_process_group和torch.nn.parallel.DistributedDataParallel。这里有个容易踩的坑进程初始化依赖环境变量RANK、WORLD_SIZE、MASTER_ADDR、MASTER_PORT。RANK是当前进程在全局的编号WORLD_SIZE是总进程数。多机场景下每台机器上的local_rank指的是该进程在本地GPU中的编号比如8卡机器上的0-7。3.3 ZeRO与显存冗余的进一步优化DDP虽然解决了数据并行的通信效率问题但前面说过它每卡都要存一份完整参数、梯度和优化器状态。对大模型来说这是巨大的显存浪费。ZeROZero Redundancy Optimizer系列是微软DeepSpeed提出的一套显存优化方案。它的核心思想不是减少总显存占用而是把原来所有卡上都重复存的那份优化器状态、梯度和参数切分到不同的卡上每卡只存自己负责的那个分片。用的时候再通过通信把需要的部分取回来。以Adam优化器为例ZeRO的三个阶段是这样递进的ZeRO-1把优化器状态momentum和variance切分到各卡每卡只更新自己那部分的参数状态。ZeRO-2把梯度也切分每卡在反向传播时只保留属于自己参数分片的那部分梯度其余分片通过ReduceScatter丢弃。ZeRO-3把参数本身也切分前向和反向过程中用到哪个参数就临时通过AllGather把那段参数广播到所有需要它的卡上。ZeRO-2在生产中使用最广泛因为它在显存节省和通信开销之间比较平衡。ZeRO-3虽然能跑更大的模型但每个step都要做参数AllGather通信量剧增如果没有高速网络收益会被通信成本吞掉。所以实际工程里经常看到“ZeRO-3 checkpointing CPU offload”组合是一种不得已而为之但确实有效的方案。下面给一个最简单的DeepSpeed配置示例让大家对实际落地的样子有直观感受{ train_batch_size: 256, gradient_accumulation_steps: 2, zero_optimization: { stage: 2, offload_optimizer: { device: cpu, pin_memory: true } , allgather_partitions: true, reduce_scatter: true }, fp16: { enabled: true, auto_cast: true, loss_scale: 0 } }这个配置常见于训练参数量在10B级别的模型。开ZeRO-2后单卡显存占用通常会比纯DDP下降40%-50%如果再把优化器状态offload到CPU还能进一步省出容量。3.4 分布式训练中的超参调整逻辑分布式训练不是简单地把代码改成分布式就能跑batch size变了学习率和训练步数也要跟着变。核心原则通常遵循线性缩放规则当全局batch size扩大N倍时学习率也近似扩大N倍但需要配合warmup。原因是更大batch意味着梯度估计更准每个step的更新方向更接近真实梯度方向所以可以迈更大的步子。但如果batch扩大得太厉害比如从256加到4096学习率单纯线性放大很容易发散这时反而需要给学习率设置一个上限或者用更复杂的warmup和衰减策略。梯度累积gradient accumulation是扩大全局batch而又不增加单卡显存占用的常用手段。它把一个大batch拆成多个micro-batch逐步累加梯度攒够后做一次参数更新。比如全局batch size要256单卡一次只能塞下2条样本8卡就是16条/step那么需要累积16个step才够更新一次参数。注意累积时损失的scale要正确最好在累积完成后统一除以累积步数不要每个micro-batch都除一遍否则数值容易漂移。4. 从单机脚本到多机多卡一次完整的实操记录4.1 节点规划和环境准备多机训练和单机训练在启动方式上就不同。单机只要设好CUDA_VISIBLE_DEVICES就能控制多卡多机则必须有明确的节点角色划分通常一台机器被指定为master主节点其余机器是worker从节点。主节点负责维护训练进程的协调信息包括NCCL初始化时的全局通信组建立、状态同步等。实际操作中master的IP和端口要能被所有节点访问。如果节点之间有防火墙一定要在实验前确认端口在放行列表里。网络方面多机训练最理想的方案是用RDMARemote Direct Memory Access网络或者InfiniBand它的带宽远高于普通以太网CPU开销也小很多。没有RDMA的话NCCL会自动回退到TCP Socket通信这时候要注意设置NCCL_SOCKET_IFNAME指定通信网卡避免NCCL误用存储网络或者管理网络导致初始化和通信都异常缓慢。4.2 一份可直接参考的启动脚本这里拿出一份我曾经用于4机8卡共32卡训练任务的启动脚本简化版关键位置都会注释说明# 每台机器上执行的命令以4机8卡为例 export MASTER_ADDR192.168.1.10 # 主节点IP export MASTER_PORT29500 # 通信端口 export NCCL_IB_DISABLE0 # 启用IB/RoCE export NCCL_SOCKET_IFNAMEeth0 # 指定通信网卡别用docker0之类 torchrun \ --nnodes4 \ --nproc_per_node8 \ --rdzv_endpoint192.168.1.10:29500 \ --rdzv_backendc10d \ --max_restarts0 \ train.py \ --model-config config/7b.json \ --batch-size 2 \ --grad-accum-steps 8在torchrun启动方式下每台机器执行的命令基本是一样的大家会通过--rdzv_endpoint指定的主机做分布式发现。这种方式比老的--master_addr手动传参更省心因为不需要分别在每台机器上改不同的启动参数所有机器跑同一条命令即可。4.3 启动代码里的关键部分训练脚本内部要和启动器约定好如何获取进程身份信息。在torchrun模式下通过环境变量就能拿到import os import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP def setup(): dist.init_process_group( backendnccl, init_methodenv:// ) torch.cuda.set_device(int(os.environ[LOCAL_RANK])) def main(): setup() local_rank int(os.environ[LOCAL_RANK]) global_rank int(os.environ[RANK]) world_size int(os.environ[WORLD_SIZE]) model build_model() model model.cuda(local_rank) model DDP(model, device_ids[local_rank], output_devicelocal_rank) dataset get_dataset() sampler torch.utils.data.distributed.DistributedSampler( dataset, num_replicasworld_size, rankglobal_rank, shuffleTrue ) loader DataLoader(dataset, batch_size2, samplersampler) for epoch in range(10): sampler.set_epoch(epoch) for step, batch in enumerate(loader): ...这里最容易被新手忽略的是DistributedSampler。如果不用它每个进程都会从原始数据集的第0条开始取数导致32个进程读了完全一样的32份数据等于实际batch size只在单卡范围内分布式扩展的收益荡然无存。sampler.set_epoch(epoch)的作用是保证每个epoch里数据切分方式不同避免模型在重复的“伪全局batch”上训练。4.4 训练监控和瓶颈定位跑起来之后不要只看loss还要看实时状态指标。我常用的监控维度以及判断标准如下GPU利用率如果低于80%说明GPU在等数据或等其他卡的通信结果需要排查数据加载或网络通信。网络收发吞吐用nvidia-smi dmon或NCCL的debug日志看通信量如果收发一直顶着上限通信已经成为瓶颈。各卡loss差异如果同一step不同卡的loss差别很大多半是数据sampler没有正确shuffle或者batch归一化相关参数跨卡同步出了问题。训练日志step时长的抖动如果step time间隔忽高忽低多半是某台机器存在资源争抢比如有人在这台机器上额外跑了任务。实测下来训练速度上不去的排查顺序通常是先看GPU利用率利用率低则看数据加载器有没有瓶颈数据加载正常则看通信时间占比通信占比高则考虑是否切分不合理、是否应该把某些操作从同步改成异步。5. 跳出模型训练并行与分布式在工程侧的延展学CS336之前我以为“并行分布式”是很窄的领域后来做业务系统才发现同一套思维在工程上遍地都是。分布式锁、分布式任务调度、分布式事务、监控系统部署本质上都是“多节点协作时如何保证正确性和效率”的问题。这里挑三个我在工程里经常用到、且和并行训练技术栈高度相关的方向展开。5.1 分布式任务调度多个worker别重复干活训练任务自己有一套调度逻辑业务系统里的定时任务同样需要调度。比如每天凌晨要跑一堆报表任务、定期清理数据、或者定时触发模型重训如果系统是多节点部署的直接在每个节点上配cron就会导致同一个任务被执行多遍。现成的方案里XXL-Job算是国内用得很多的分布式任务调度平台。它的核心模型是“调度中心”加“执行器”。调度中心负责任务的编排和触发执行器部署在业务节点上负责真正的逻辑执行。同一个任务调度中心只会选一台执行器去触发不会广播给所有节点。XXL-Job里有很多值得琢磨的设计比如故障转移如果执行器A执行任务时挂了调度中心会把任务重新路由到执行器B执行保证任务最终完成。再比如分片广播一个大任务可以配置成按节点分片执行每台机器只处理分给自己的那部分数据有点像数据并行里的DistributedSampler。做这种分布式任务编排时真正需要注意的坑是“任务幂等性”。因为不管是故障转移还是超时重试任务都有概率被重复执行。如果你在任务里不加幂等控制补一张数据表的时候可能补出重复数据发消息的时候可能发出两条。务实做法是在任务表里增加一个执行批次ID任务开始前先查批次是否已成功成功过就直接返回。5.2 Redis分布式锁并发控制的经典解法多个节点并发处理同一笔资源时靠本地锁不够必须用分布式锁。最常见的实现是Redis分布式锁。一个看似能用的错误姿势是先setnx成功后再expire。如果setnx之后、expire之前进程崩溃锁就永远不释放其他节点永远拿不到锁。正确做法是一条命令完成加锁和过期设置SET lock_key unique_value NX PX 30000这个命令的意思只有lock_key不存在时才设置成功NX同时带上30秒过期时间PXvalue是本次加锁的唯一标识。释放锁时也不能简单地DEL而是要先比对value是否还是自己的唯一标识再用Lua脚本保证比对和删除是原子操作if redis.call(get, KEYS[1]) ARGV[1] then return redis.call(del, KEYS[1]) else return 0 end如果不做唯一标识校验可能会出现一个经典的并发事故线程A的锁因为业务处理时间过长过期了线程B加锁成功并开始处理此时线程A处理完毕执行DEL把线程B的锁删掉了然后线程C又加锁成功导致B和C同时进入临界区。唯一标识就是用来防止“删别人的锁”的。5.3 分布式事务多节点一致性怎么保证训练场景里多个GPU同步梯度业务场景里则是多个服务共同完成一笔操作。比如一个下单流程里订单服务要写订单表库存服务要扣库存两个操作分布在不同的数据库里。如果扣库存成功但订单入库失败怎么保证两边数据一致常见方案有2PC两阶段提交、TCCTry-Confirm-Cancel、Saga和本地消息表。它们的核心逻辑都不复杂难点在于工程落地的边界条件。方案核心思路优点缺点适用场景2PC/XA先prepare再commit引入事务协调者强一致实现标准协调者单点阻塞时间长单体跨库且对一致性要求极高TCCTry阶段预留资源Confirm确认Cancel回滚业务侵入可接受性能较好每个操作都要写三个方法开发量大需要较强的资源管控和回滚语义Saga每个本地事务成功则触发下一个失败则执行反向补偿不锁资源适合长事务最终一致中间状态对外可见跨服务业务流程较长允许短暂不一致本地消息表业务操作和消息写入同一个本地事务由消息队列异步通知简单可靠实现成本低需要消费者幂等消息有延迟异步解耦但需要最终一致的场景做分布式事务选择时我个人的建议是别动不动就上分布式事务框架先想清楚业务是否能接受最终一致性。能接受就用Saga或本地消息表接受不了再考虑TCC甚至2PC。分布式事务是一个高成本方案它的每一次回滚都可能触发一整条链路的数据修正复杂度远超出单机事务想象。一个高效的做法是把分布式系统拆成“主流程事务 对账补偿”两步走。先保证主流程的业务在自己的本地事务里完成再通过消息或定时任务去驱动其他节点的数据同步。即使中间失败对账脚本也能捞回来。这是很多长链路业务在实践的稳健路线。6. 实战中会遇到的常见问题与排查方法这个部分是我最想分享的。分布式训练和单机训练最大的不同是出问题时错误信息往往不直观有时候甚至不报错只是训练突然变慢了几倍你完全没察觉。下面按症状分类给出我踩过坑的经验总结。6.1 NCCL初始化超时或直接卡住不动这个错误几乎是每个多机训练新手都会遇到的具体表现是训练启动后进程一直停在“NCCL version xxx starting”或“init_process_group”这里然后过一段时间报超时。排查顺序一般是先用telnet或nc验证节点之间的端口是否互通nc -zv MASTER_IP 29500不通则看防火墙和安全组。检查NCCL指定的网卡是否选对了如果机器上有多个网口用NCCL_SOCKET_IFNAME强制指定。检查所有节点的共享存储是否一致训练脚本路径、Python环境是否一致。如果节点数量很多可以先从小规模两机启动验证网络基础再逐步扩到全量。在NCCL相关超时排查里大量问题其实出在没有使用共享文件系统或者每台机器读取的代码版本不一致导致训练配置都不同。务必保证所有节点的代码和数据路径完全一致。6.2 多卡训练比单卡还慢GPU利用率上不去这种情况最让人崩溃。明明加了卡速度却上不去甚至更慢。需要先确认瓶颈是在数据读取还是通信。先用pytorch的profiler去抓单个step的耗时构成。如果发现DataLoader加载数据的时间远大于计算时间说明数据管道卡脖子了。解决办法包括开num_workers、用pin_memoryTrue、换更快的存储例如把数据缓存到本地NVMe而不是全走网络文件系统。如果通信占比很高就要考虑分布式策略是否需要调整。比如对单机8卡场景做了跨机张量并行通信可能直接拖垮性能。跨机的通信带宽通常只有卡间NVLink的几分之一如果频繁做AllReduce耗时必然飙升。想定位通信瓶颈一个实用的命令是设置NCCL_DEBUGINFO运行一小段训练看NCCL日志里通信的内核耗时占比。6.3 loss发散或震荡确认真实batch size变化分布式训练默认batch size是单卡batch乘上卡数如果全局batch size比以前单卡训练时大了N倍而学习率没有做对应的调整loss很容易发散。还有一个隐蔽的坑在BatchNorm这类需要跨卡统计的层上。对于卷积网络或带BatchNorm的模型分布式训练时默认每个进程独立计算均值和方差这会导致global batch size很大但BatchNorm统计口径仍然是小batch影响模型收敛。常见解法是启用SyncBatchNorm让所有卡共享统计量。对Transformer类模型因为用的是LayerNorm不存在这个问题但很多迁移到LLM训练的团队早期可能并不会注意到这点区别。6.4 显存不够训练跑到中途OOM每次说“显存不够”其实都要先搞清楚是什么占了显存。用torch.cuda.memory_summary()可以非常直观地看到各分配器的内存占用情况。优化手段按成本从低到高排列开梯度检查点activation checkpointing用少量重计算换取显存通常能把激活显存降低60%-70%。开启混合精度训练参数和激活的主存储仍然可以是fp32但前向和反向的主要计算放到fp16/bf16显存能省接近一半。调低batch size配合梯度累积保持全局batch不变。用ZeRO-2或ZeRO-3分散优化器状态和参数。最终手段才是模型并行切分或CPU offload因为它们的工程复杂度最高。6.5 快速自查清单症状优先排查方向快速处理建议启动卡住超时节点连通性、NCCL网卡选择、master地址是否正确用nc测试端口设NCCL_SOCKET_IFNAME训练很慢GPU利用率低数据加载、通信占比、存储IO用profiler抓step构成检查num_workers和pin_memoryloss发散学习率、batch size扩展、数据sampler检查全局学习率是否随batch增大而调大确认sampler设置正确显存OOM激活值、优化器状态、梯度开gradient checkpointing换ZeRO降batch加梯度累积各卡loss差异大数据sampler、模型初始化确认使用DistributedSampler且所有进程初始化为同一随机种子通信stuck或hang住NCCL版本、共享文件系统统一容器镜像和依赖版本用NCCL_DEBUGINFO追踪关键日志每次我重新回看CS336这个lecture都会有个强烈的感受并行分布式本身并不是什么高不可攀的黑魔法它就是一套“如何让多台机器高效协作解决问题”的工程方法论。真正把它跑到产业级的时候一半的功夫其实都在那些不起眼的细节上比如网卡选对没有、sampler有没有用对、超时日志是不是被吞了。把基础概念搞扎实再愿意动手实测几次这些“看着吓人”的问题基本都是纸老虎。希望这篇来自实战一线的实践笔记能让你在搭起第一套分布式训练时比我当时少熬几个通宵。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价