ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

torch.compile+梯度累积:16G显存也能撑起大batch训练

torch.compile+梯度累积:16G显存也能撑起大batch训练 开头先讲一个我自己的经历手里只有一张 16G 显存的卡却要同时跑图像分类和中文文本微调batch size 从 64 一路降到 8、降到 4终于不 OOM 了但 loss 曲线抖得像心电图。后来我靠两个不起眼的组合拳把日子过舒服了一个是torch.compile把模型从逐行解释执行变成整图编译优化另一个是梯度累积把多个小 batch 的梯度攒起来拼出一个等效大 batch。这两招加起来改动不到十行代码却同时解决了显存不够和收敛太慢两个问题。无论你是在训练自己的 OCR 模型还是微调中文 RoBERTa只要训练循环里还有backward()下面这些经验应该都能直接用上。1. 一张 16G 的卡如何撑起 batch size 32 的训练1.1 训练慢的真相算力空闲与 batch 太小大多数人觉得训练慢是显卡不够好但实际卡脖子的问题往往有两个。第一个是算力没喂饱。PyTorch 默认的 eager 模式是一步步解释执行的每做一次张量运算都要经过 Python 解释器、PyTorch 调度、CUDA kernel launch 这一整套流程。一个简单的卷积加 ReLU 加 BN就会产生三到四次 kernel launch每次 launch 都有固定开销。模型不大时这些开销占比极高GPU 大部分时间在等指令真正算数的时间反而很短。这就像一个大厨每次做一道菜都要重新读一遍菜谱、重新洗一遍锅而不是把十道菜的流程一次性背熟后连续出餐。第二个是 batch size 被显存压得太小。显存除了装模型参数还要装每一层的中间激活值。batch size 越大反向传播需要保留的中间激活越多。我那张 16G 的卡跑 ResNet50 想开 batch 64直接 OOM被迫降到 16梯度噪声变大收敛变慢一个 epoch 里还经常出现剧烈震荡。这本质上是显存受限 → batch 过小 → 梯度估计不准 → 需要更多迭代才能收敛的连锁反应。1.2 两个技巧怎么分工为什么不冲突torch.compile和梯度累积解决的问题完全不同所以它们天然不冲突。torch.compile管的是微观效率让单个 batch 的前向和反向跑得更快梯度累积管的是宏观策略让多个小 batch 在权重更新上等效于一个超大 batch。打个比方你想搬 100 块砖上楼。梯度累积是让你一次别搬太多每趟 10 块分 10 趟避免把腰闪了——对应显存限制torch.compile则是把弯腰、搬砖、上楼、放下这一套动作优化成更流畅的节奏让每一趟都更省力。一个是改变搬运策略一个是优化动作效率两者叠加当然可以同时使用。还有个关键认知梯度累积本身通常不会变快它只是让你在不换卡的情况下用上大 batch 的收敛优势真正把墙钟时间压下来的是torch.compile。我见过不少人把梯度累积吹成加速技巧严格说它是个显存规避技巧加速是顺带的收敛红利。把这两件事分清楚后面调参才不会迷糊。2. torch.compile从逐行解释到整图编译提速到底靠什么2.1 Dynamo 捕获图、Inductor 生成代码一次 forward 的编译过程torch.compile是 PyTorch 2.0 引入的一套编译栈核心是 Dynamo 和 Inductor 两部分。第一次调用模型时Dynamo 会分析 Python 字节码把前向计算捕获成一张 FX Graph也就是把计算步骤从 Python 层抽象成一张静态图。然后 Inductor 后端会在这张图上做优化并用 Triton 生成专门的 GPU kernel。之后每次 forward 都直接执行优化好的计算流程不再经过 Python 解释器。这里有一个容易被忽略的点model torch.compile(model)返回的是一个编译后的新对象但它和原模型共享参数存储。所以你在训练循环里正常调用model(input)、正常backward()即可参数更新会同步到编译版本上。你不用手工切换什么。编译过程发生在第一次前向时所以首次调用会明显卡一下后面就顺了。这个机制对训练脚本几乎是透明的代价是你得理解图捕获的边界如果模型里有大量动态控制流、Python 列表推导、或者某些自定义算子Dynamo 可能无法完整捕获会退回到 Python 模式执行这种现象叫 graph break。图一旦断了性能提升就大打折扣。所以torch.compile不是无脑开关后面我会专门讲哪些模型适合。2.2 三种优化带来的收益分别在哪编译后主要换来三类收益。第一是算子融合。比如一个卷积后面跟着批归一化和 ReLU原本要启动三个 kernel、中间结果写回显存、再读出来融合后变成一个 kernel中间结果直接留在寄存器或片上缓存里。显存带宽是 GPU 上最稀缺的资源之一省掉这些中间读写收益非常直观。第二是减少 Python 调度开销。静态图模式下C 执行引擎可以直接按图调度不再每一轮都跑到 Python 解释器里去查属性、发指令。对层数多、小算子多的模型比如 Transformer 里一堆 LayerNorm、注意力矩阵乘法、残差 Add这个节省相当可观。第三是 CUDA Graph 带来的 kernel launch 削峰。默认模式下每个 op 都要 launch 一次一次 launch 大概 3 到 10 微秒。模型有几百个算子累加起来就是毫秒级损耗。torch.compile(modereduce-overhead)会把整张图打包成一个 CUDA Graph一次 launch 完成所有计算尤其适合推理或训练中计算图固定的大模型。2.3 mode 参数选型default、reduce-overhead 还是 max-autotunetorch.compile最常用的三个 mode我直接给选型建议模式适合场景首步编译代价提速潜力default大多数训练/微调模型结构复杂或输入有变化较低几秒到十几秒中等偏高reduce-overhead大模型、Transformer、输入 shape 固定较高可达几十秒高max-autotune模型定型后追求极限性能很高可能数分钟最高我自己的习惯是训练第一天先用default把流程跑通确认 loss 正常稳定后切到reduce-overhead再跑正式实验。max-autotune我一般只在推理阶段用训练阶段启动代价太大而且对一个小数据集来说那点额外收益不值当。如果你的模型里有大量动态 shape 或者控制流reduce-overhead容易触发 graph break反而比default更不稳。2.4 适用边界哪些模型收益明显哪些会白折腾从实测看收益最明显的是 CNN 和 Transformer 这两大类。ResNet、EfficientNet 这类卷积模型卷积和归一化的融合空间大BERT、GPT、RoBERTa 这类 Transformer小算子多、调度开销占比高编译效果也立竿见影。我微调中文 RoBERTa 时开reduce-overhead单 step 耗时有可感知的下降。收益不明显甚至变差的也有常见三类一是模型本身就很小比如只有一个全连接层的逻辑回归kernel launch 不是瓶颈编译开销反而拖后腿二是控制流极多的模型比如 RNN 里对序列逐步循环、或者带复杂 beam search 的解码部分图经常断三是自定义算子、第三方 CUDA 扩展较多的代码Dynamo 可能捕获不了白白增加一层风险。所以我的原则是先小步试、看日志、再全量跑。开 compile 后如果发现训练循环里频繁出现Graph break之类的警告说明你的模型结构对编译不友好要么改结构要么干脆别用。3. 梯度累积显存不够时的伪大 batch方案3.1 原理拆解为什么除以累积步数就能等效大 batch梯度累积的逻辑很简单原本想一次喂 32 个样本显存不允许那就一次喂 8 个连续喂 4 次把 4 次算出的梯度加起来再更新一次权重。这样权重每更新一次见过的样本数还是 32数学上约等于跑了一个 batch size 32。但要写对关键是每个 micro-batch 的 loss 要除以累积步数。假设每个 micro-batch 的 loss 已经是该 batch 内样本的平均 loss那么第 i 个 micro-batch 的梯度是平均梯度累积 4 步后再除以 4得到的才是 32 个样本的全局平均梯度。如果不做除法梯度范数会被放大 4 倍等效于学习率放大 4 倍训练很容易直接发散。我第一次写梯度累积就吃过这个亏loss 冲上天以后才反应过来是分母漏了。一个标准的累积循环长这样accum_steps 4 optimizer.zero_grad(set_to_noneTrue) for step, (inputs, labels) in enumerate(dataloader): loss model(inputs, labels) / accum_steps loss.backward() if (step 1) % accum_steps 0: optimizer.step() optimizer.zero_grad(set_to_noneTrue)注意zero_grad(set_to_noneTrue)这个细节。默认zero_grad()是把梯度置零梯度张量还在显存里set_to_noneTrue是直接把梯度张量释放置 None能省下一点显存。累积场景下我强烈建议用这个参数。3.2 学习率到底调不调这是最容易踩的坑梯度累积之后学习率要不要放大是被问得最多的问题也是最容易搜到互相矛盾答案的地方。我说说自己的结论。如果代码里严格做了loss loss / accum_steps那么多个 micro-batch 的平均梯度和直接用一个等效大 batch 算出的平均梯度在数学上是一致的。所以基线学习率不需要动——你原本用 batch 32 训练觉得 3e-4 合适改成 8×4 梯度累积后继续用 3e-4 是合理的起点。但这里有个隐蔽的第二层问题你可能本来就是因为显存不够才用小 batch现在等效 batch 从 16 变成了 64而大 batch 训练本身有一个学习率缩放的经验法则——batch 翻倍学习率可以按 1 到 2 倍往上调。问题在于你并不确定原来的 batch 16 配 3e-4 是不是最优组合。直接从 batch 16 跳到等效 batch 64把学习率也翻倍收敛到底变好还是变差很难拍脑袋判断。我的实操建议是第一轮循环保持原学习率不变跑几百个 step 看看 loss 曲线。如果曲线比原来更平滑、下降更快说明等效大 batch 确实带来了收益可以继续如果曲线收敛变慢再尝试把学习率按sqrt(等效batch / 原batch)的系数上调也就是 4 倍 batch 就乘 2而不是乘 4。这样稳妥得多至少不会第一个 epoch 就把训练搞爆。3.3 BatchNorm 与梯度累积的隐性冲突梯度累积在数学上等效大 batch但有一个东西不等效——BatchNorm。BN 在前向时用的是当前 micro-batch 内的均值和方差做归一化同时更新 running stats。真实的大 batch 是用 32 个样本的统计量归一化梯度累积却是每个小 batch 分别归一化相当于网络在 4 个不同的数据分布之间来回切换。micro-batch 越小这种统计噪声越大。此外running stats 的更新频率也从每 32 个样本更新一次变成了每 8 个样本更新一次等效于 momentum 被放大了 4 倍训练早期指数滑动平均更容易被噪声带着跑偏。解决办法按优先级排序第一如果模型能用 LayerNorm/GroupNorm/RMSNorm直接换掉一劳永逸这也是为什么很多 Transformer、TTS 模型在梯度累积下毫无压力第二如果必须用 BN尽量保证 micro-batch size 不小于 8让统计量别太离谱第三实在不行可以每几个 micro-batch 手动同步一次 BN 统计量但实现复杂、收益不稳定我一般不会走到这一步。4. 两招叠加训练循环的完整改造与 AMP 配合4.1 完整可运行的训练骨架代码把torch.compile和梯度累积放在一起时代码反而比单独用其中任何一个更简单。核心思路是只编译模型不要编译训练循环。import torch import torch.nn as nn from torch.cuda.amp import autocast, GradScaler micro_batch 8 # 单次前向喂的样本数 accum_steps 4 # 梯度累积步数 effective_batch micro_batch * accum_steps # 等效 batch 32 model MyModel().cuda() model torch.compile(model, modereduce-overhead) # 训练期推荐 optimizer torch.optim.AdamW(model.parameters(), lr2e-5) scaler GradScaler() train_loader DataLoader(dataset, batch_sizemicro_batch, shuffleTrue, num_workers4, pin_memoryTrue, persistent_workersTrue) for epoch in range(num_epochs): optimizer.zero_grad(set_to_noneTrue) for step, (inputs, labels) in enumerate(train_loader): inputs inputs.cuda(non_blockingTrue) labels labels.cuda(non_blockingTrue) with autocast(): loss model(inputs, labels) # micro-batch 平均 loss loss loss / accum_steps # 平均到等效 batch scaler.scale(loss).backward() if (step 1) % accum_steps 0: scaler.unscale_(optimizer) nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_noneTrue)这段代码里有三个位置最容易写错下面单独展开。4.2 GradScaler、梯度裁剪与累积边界的先后顺序先说混合精度和梯度累积的配合。AMP 下GradScaler会在 loss 上乘一个缩放因子避免半精度下的小梯度被冲刷成零。累积 N 个 micro-batch 时每个 loss 都要除以 N所以每个 micro-batch 的反向梯度都会被scaler重新缩放累积后的梯度是多个缩放值的和。scaler.step(optimizer)必须在累积边界处执行而不是每个 micro-batch 都执行否则权重更新频率会变成原来的 N 倍整个累积策略就失效了。再说梯度裁剪的位置。如果直接写clip_grad_norm_再scaler.step梯度是在缩放状态下被裁剪的裁剪阈值会被缩放因子污染导致裁剪效果不稳定。正确顺序是先scaler.unscale_(optimizer)把梯度还原到真实数值再裁剪再scaler.step。代码里我已经按这个顺序写了。最后一个容易漏掉的点epoch 结束时如果累积步数没有凑整剩下的梯度不要留到下一个 epoch。要么单独对残余梯度做一次scaler.step要么直接optimizer.zero_grad()把残渣清掉。我习惯直接清掉因为强行用不足 N 步的梯度更新等效 batch 变了学习率心态也跟着变收益说不清。4.3 编译热身与 DataLoader 细节torch.compile的编译发生在第一次前向而且训练模式下还需要编译反向图所以第一次调用可能慢到让你怀疑卡死了。我建议在正式训练前做一个编译热身def compile_warmup(model, sample_shape, labels_shape): model.train() dummy_x torch.randn(*sample_shape, devicecuda) dummy_y torch.randint(0, num_classes, labels_shape, devicecuda) with autocast(): loss model(dummy_x, dummy_y) / accum_steps scaler.scale(loss).backward() optimizer.zero_grad(set_to_noneTrue)跑完这个函数前向和反向的计算图都被编译缓存了正式训练时第一个 batch 不会再出现几十秒的卡顿。注意 dummy 数据的 shape 和类型要跟真实输入一致否则热身的图缓存用不上等于白做。DataLoader 方面的细节是pin_memoryTrue加non_blockingTrue让 CPU 到 GPU 的拷贝不阻塞计算。persistent_workersTrue可以避免每个 epoch 重新 fork worker 的开销。还有一个常见配置是torch.backends.cudnn.benchmark True但它只对固定 shape 有效如果你的输入长度动态变化benchmark 模式反而会反复搜索最优卷积算法拖慢速度这点要记住。4.4 显存怎么变化为什么不会爆之前说过真实大 batch 的显存瓶颈在中间激活值它随 batch 线性增长。而梯度累积的每个 micro-batch 前向结束后autograd 图在backward()完成后就释放了不会跨 micro-batch 累积。所以总的显存占用约等于模型参数 梯度 优化器状态 单个 micro-batch 的激活峰值跟你累积多少步关系不大。这就是为什么可以拿 8 的 micro-batch 拼出等效 64 的 batch而显存纹丝不动。需要提醒的是torch.compile(modereduce-overhead)因为用 CUDA Graph 缓存了一些执行资源显存占用会比 eager 模式略高一点点但通常只有几十到几百 MB相对你省下的激活显存完全值得。5. 实测效果与踩坑清单从性能不升反降到数值发散5.1 我在图像和文本任务上实测的数字拿我最近的两个任务说事。第一个是 ResNet50 图像分类16G 的卡原计划 batch 64降到 batch 16 才能跑。单看这一个改动训练收敛明显变差加上torch.compile(modedefault)后单 step 耗时有接近 25% 的下降再把 4 步累积起来等效回 batch 64loss 曲线跟原始 batch 64 基本重合墙钟时间反而更短。第二个是中文 RoBERTa 微调batch 8 搭配reduce-overhead单 step 提速接近 15%这个场景本身是 LayerNorm 结构梯度累积对归一化没有影响属于最省心的组合。任务基线torch.compile梯度累积实际感受ResNet50 分类batch 16 硬跑单 step 约快 25%等效 batch 64收敛更稳总时间反而缩短RoBERTa 中文微调batch 8单 step 约快 15%等效 batch 16显存稳定不再 OOM需要说明的是模型结构、GPU 型号、数据 shape 都会影响提速幅度数字仅供参考。如果你的收益明显高于或低于这个范围优先检查是不是踩了下边某个坑。5.2 坑一把训练循环整个塞进 compile我见过有人为了让速度更快把整个 train_step 包进一个函数再torch.compile结果训练损失不降反升还频繁报错。原因很简单optimizer.step()、scaler.step()、zero_grad()这些操作包含大量 Python 状态更新和 CUDA 同步它们不是前向/反向计算图的一部分硬塞进编译范围只会让图捕获失败或者触发多次重编译性能更差。正确做法永远是把torch.compile加在模型实例上而不是某个训练函数上。模型的前向和反向是纯计算适合编译优化器状态更新、学习率调度、梯度清零这些管理动作应该留在 Python 层跟着你的业务逻辑走。这个边界划清楚编译才能稳定发力。5.3 坑二动态 shape 导致反复重新编译我一开始在 NLP 任务上开reduce-overhead训练特别慢每跑几个 batch 就卡一下。查了半天发现是序列长度不固定每个新长度都会触发一次 guard 检查失败然后重新编译生成新 kernel。torch.compile本身对动态 shape 有保护机制但代价就是每来一个新形状就重来一次训练过程被编译切得支离破碎。解决思路主要有两种。一是把输入统一 padding 到固定长度简单粗暴但有效显存浪费一点换训练稳定。二是用torch._dynamo.mark_dynamic显式标记某些维度是动态的让编译器生成更通用的实现减少重编译次数。第二种方案省显存但 API 随版本有差异用时先查你当前 PyTorch 版本的文档。我个人的偏好是数据允许就 padding省心数据实在 padding 不动再用 mark_dynamic。5.4 坑三AMP 下的 Inf/Nan 与梯度裁剪顺序训练过程中 loss 突然变成 nan是很多人被劝退的最后一根稻草。在 AMP 梯度累积的组合里最常见的原因是梯度裁剪位置不对。如果你在scaler.step()之后裁剪或者干脆没裁剪累积若干步之后半精度梯度的下溢问题会被放大某个 batch 里一个小小的 inf 就可能污染整个累积周期。正确做法就是我第 4 节写的scaler.unscale_→clip_grad_norm_→scaler.step。另外GradScaler默认会动态调整缩放因子但它只根据上一步有没有出现 inf/inf来 update而梯度累积会把问题延迟到边界处才暴露所以如果训练跑到一半开始 nan先看累积边界附近的日志再检查裁剪顺序十有八九是这里出了问题。5.5 如何科学验证提速效果不要靠感觉快了来下结论。我每次调完都会做一个最小对照实验固定随机种子固定 epoch 数分别跑四组——什么都不加、只加 compile、只加梯度累积、两个都加。统计稳定运行后的 samples/s 或 tokens/s注意把第一次 warmup 的耗时排除掉因为编译首步会严重拉低平均值。显存监控用torch.cuda.max_memory_allocated()记录峰值而不是看 nvidia-smi 里瞬时跳动的数字。做完这组对照你才知道收益到底来自哪里也方便下次换模型时快速复制这套配置。带梯度累积时还有一个隐藏变量数据加载速度。累积步数越多模型更新频率越低如果 DataLoader 跟不上GPU 反而会出现等待。所以做对照实验时开torch.profiler看一眼 GPU 利用率如果低于 80%别急着怪 compile先查数据管线。我自己就遇到过 compile 后 GPU 利用率反而下降的情况最后发现是num_workers太少数据加载成了瓶颈。最后说点个人体会。这两个技巧之所以适合放在一起讲是因为它们都具备改动小、风险低、可组合的特点。torch.compile不是银弹遇到 graph break 该拆就拆梯度累积也不是直接把学习率翻倍就能起飞要分清等效大 batch 的数学和大 batch 调参的经验是两回事。但只要你按上面的顺序一步步试先跑通再优化最后再压榨绝大多数模型都能在有限显存下获得实打实的收益。
返回列表