☰
torch.compile+梯度累积:PyTorch训练提速与显存优化实战
2026/10/2 2:49:19 网站建设 项目流程

在日常训练模型的时候,我经常听到两种声音:一种说“GPU太贵,显存不够,batch size只能调到4”,另一种说“训练太慢,一个epoch要跑一天”。这两个痛点其实可以同时缓解,靠的就是 torch.compile 和梯度累积。torch.compile 是 PyTorch 2.x 带来的编译加速方案,能在不改变模型逻辑的前提下把前向和反向的计算图优化到接近硬件极限;梯度累积则是把多个小 batch 的梯度攒起来,攒够一定数量再更新一次权重,变相扩大 batch size,解决显存瓶颈。两者结合起来,既能提升单次迭代的速度,又能减少权重更新次数、稳定训练曲线,非常适合在消费级显卡上训练 YOLO、EasyOCR 这类检测/识别模型,或者微调 RoBERTa、ResNet 这类预训练模型。

这篇文章我会先从原理层面讲清楚 torch.compile 到底编译了什么、梯度累积为什么有效,再给出可以直接抄的实操代码和参数设置,最后把我踩过的坑和排查思路整理成表格。适合正在用 PyTorch 训练自己的模型、觉得显存不够或者训练太慢的开发者,看完就能自己动手改。

1. 整体设计思路:为什么要把编译加速和梯度累积放在一起

1.1 两个技术各自解决的问题

先拆开看。torch.compile 解决的是“单步计算效率低”的问题。PyTorch 默认是 eager mode,也就是每执行一行算子,就启动一次 GPU kernel,然后等待结果返回,再执行下一行。这种动态图模式好处是灵活,坏处是大量时间浪费在 kernel 启动和 Python 解释上,尤其当模型里有大量小算子(比如残差连接、归一化、激活函数)时,开销非常可观。torch.compile 会把整个模型或某个模块的 forward 过程捕获成一张计算图,通过 TorchInductor 生成优化的 kernel,甚至能把多个小算子融合成一个 kernel,同时减少 Python 层调度和显存读写。

梯度累积解决的是“一次更新需要多大 batch”的问题。在 SGD 系列优化器中,权重更新量 = 学习率 × 梯度均值(或均值相关的统计量)。标准的做法是每看到一个 batch 就更新一次,但 batch size 受限时,梯度噪声很大,训练不稳定,尤其是目标检测、语义分割这类任务。梯度累积的做法是:连续跑 N 个 batch,每个 batch 都做前向和反向,但不更新权重,把梯度累加下来,等到第 N 个 batch 结束,再用累加后的梯度除以 N(或做等价处理)来更新一次权重。这样等效的 batch size = 单卡 batch size × 累积步数,而显存占用只有单卡 batch size 的量。

1.2 组合使用的收益模型

把两者放在一起,收益不是简单的相加,而是乘。因为梯度累积意味着每个权重更新周期内要跑 N 次前向反向,这 N 次重复计算完全可以用 torch.compile 加速。换句话说,同样的 wall-clock 时间内,编译加速让单次前向反向更快,梯度累积让权重更新更少、更平滑,两者各管一个维度。

我习惯用一张简化的“训练时间 = 前向反向次数 × 单次耗时 + 更新次数 × 更新耗时”来评估。torch.compile 降低“单次耗时”,梯度累积不直接降低总前向反向次数(因为等效 batch 变大后,总样本数不变,迭代次数会变少,但每次迭代需要更多前向反向),但它减少了“更新次数”,还减少了优化器状态更新的开销。关键收益是,当你被迫使用小 batch 时,梯度累积能让你的等效 batch 大到足以触发训练稳定性,而 torch.compile 能弥补累积带来的额外循环开销。

1.3 适用场景和踩坑前提

不是所有模型都适合无脑套这两个技术。torch.compile 目前对动态 shape、控制流较强的模型(比如 NLP 里带有复杂 padding 的 decoder)支持不够好,首次编译可能有几分钟的“预热”开销。梯度累积如果使用不当,会出现梯度累积数值溢出、BN 层统计量不准确、学习率等效缩放错误等问题。所以,我的建议是:先用一个小数据集或单个 batch 跑通,验证 loss 曲线正常,再正式开启长训练。这篇文章后面给的例子,都是我实际用在 YOLO 系列和 OCR 模型训练上的方案。

2. torch.compile 加速原理解析与实战配置

2.1 编译模式:inductor、reduce-overhead 和 dynamic

torch.compile 有三种常用模式,直接决定加速比和编译时间。只传一个 model 对象而不写 mode,默认是 “default” 模式,对应的后端是 TorchInductor,它会做算子融合和内存规划,但保留了一些动态性,所以编译时间中等,加速比通常 10% 到 30%。reduce-overhead 模式会把 CUDA graph 也利用上,减少 kernel 启动的 overhead,在多次小算子场景下提升更明显,但显存占用会略高,首次编译也更慢。max-autotune 模式会针对每个算子做 autotuning,找最佳配置,加速比最高,但编译可能需要十几分钟,不适合日常调试。

实际使用中,我一般在本地调试时用 default 或直接不开编译,在服务器上正式训练时用 reduce-overhead。用代码表示就是:

model = torch.compile(model, mode="reduce-overhead")

需要注意,torch.compile 返回的是一个新的 wrapped model,原来的 model 参数不能被直接拿去保存或做 torch.jit 序列化。如果你要保存权重,应该保存原始 model 的 state_dict,或者保存 compiled_model 内部的原始模型权重。这个细节我后面会在常见问题里说。

2.2 动态 shape 与 padding 陷阱

训练检测、OCR 模型时,输入图片的长宽经常不固定。YOLO 在训练时一般会做 letterbox,把图片统一缩放到固定尺寸,比如 640×640,shape 是静态的,相对安全。但如果你用的是可变 batch 或者动态 padding 的 NLP 模型,torch.compile 可能会因为 shape 变化而重新编译,反而比原始速度更慢。

TorchInductor 支持 dynamic shape,但需要显式标记,或者在编译后对不同的 shape 逐个触发编译。实际做法是:

model = torch.compile(model, dynamic=True)

或者更精细地,使用 torch._dynamo.mark_dynamic 标记输入张量的某个维度为动态。我在微调 RoBERTa 中文模型时,batch 内样本长度差异大,使用 dynamic=True 后,编译只做了一次,后续不同长度序列也能复用,不然每遇到一个新长度就重新编译一次,训练速度会拖慢好几倍。

2.3 图模式下的内存图与显存变化

torch.compile 在 reduce-overhead 模式下会使用 CUDA graph 捕获整个计算图,这会让显存占用比 eager 模式高一些,因为需要额外缓存 graph 的内存池。如果你之前已经因为显存不够才用梯度累积,那么再叠加 reduce-overhead 可能会显存溢出。我的做法是:先不开 CUDA graph,只做算子融合,跑一个 batch 看看显存峰值,再决定要不要加 reduce-overhead。

对于显存敏感的场景,我更推荐用 default 模式,然后用梯度累积来弥补速度。比如一个原本 batch size 8 就爆显存的任务,你可以设置 batch size 4,累积步数为 2,同时开 torch.compile default 模式。这样显存占用和原来 batch size 4 差不多,但等效 batch size 还是 8,训练稳定性不变,单步速度还提升了 15% 左右,整体收益非常明显。

3. 梯度累积的正确实现与学习率调整

3.1 朴素实现和经典误区

最简单的梯度累积写法是这样:

scaler = torch.cuda.amp.GradScaler() # 混合精度 for i, (images, targets) in enumerate(loader): with torch.autocast(device_type="cuda", dtype=torch.float16): loss = model(images, targets)["loss"] scaler.scale(loss).backward() if (i + 1) % accum_steps == 0: scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_none=True)

很多新手会在每次 loss.backward() 后直接调用 optimizer.step(),这是错的;还有人在累积结束后没有把梯度除以累积步数,导致梯度绝对值偏大 N 倍,学习率等效变大,容易发散。正确做法是:在累积的最后一步,用 scaler.scale(loss / accum_steps).backward(),或者在 optimizer.step() 前手动把梯度除以 accum_steps。注意,如果你用了 GradScaler,对 loss 除以 accum_steps 后再调用 backward,scaler 内部还是会根据 loss scale 调整梯度,所以要把除法放在 scale 之前还是之后,需要想清楚。

推荐的做法是:

loss = loss / accum_steps scaler.scale(loss).backward()

这样每一项梯度都是原始梯度的 1/accum_steps,累加 N 次后正好等于平均梯度。有些开源实现喜欢用 loss.backward(retain_graph=True),那个是给对抗生成网络用的,训练常规模型时不要加 retain_graph,否则不仅慢,还会造成梯度累积。

3.2 学习率、warmup 与等效 batch size

当你把等效 batch size 从 B 增大到 B×N,学习率应该怎么变?常见经验是线性缩放规则:如果 batch size 变成原来的 N 倍,学习率可以乘以 N 的平方根或者直接乘以 N,但实际任务要谨慎。我自己的经验是:在目标检测和 OCR 这类任务上,等效 batch size 翻倍后,学习率只调高 30% 到 50% 就足够,不要直接翻倍,否则训练初期 loss 很容易震荡。

原因是,梯度累积并没有真正同时看到 N 个 batch 的样本,权重更新间隔内的梯度是多次独立采样求平均,方差确实变小了,但 BN 层的统计量仍然基于单一 batch 计算。如果模型里有 BatchNorm,累积跨多步时,BN 的 running_mean 和 running_var 在每个小 batch 都会更新,这跟大 batch 训练时用完整 batch 统计 BN 是有区别的。很多检测框架(比如 YOLO)在训练时默认不用 BN,或者用了 BatchNorm 但实际统计的是每个小 batch 的分布,所以问题不明显。但如果你微调 ResNet 这类高频使用 BN 的分类模型,就要留意这个差异。

对于 warmup,梯度累积会让参数更新频率变低,warmup 的总步数应当对应“更新次数”而不是“迭代次数”。比如原来每个 epoch 更新 1000 次,warmup 10 个 epoch;现在累积步数 N=4,每个 epoch 只更新 250 次,那 warmup 应该仍然是 10 个 epoch 对应的 2500 次更新。很多代码用 iteration 计数,要小心换算。

3.3 梯度裁剪、EMA 和噪声

如果你的训练流程里用了梯度裁剪(grad clip),在梯度累积时应该在累积完成后统一裁剪,而不是每个小 batch 都裁剪。如果每个小 batch 都裁剪,会破坏梯度的比例关系,等于加了随机噪声。同样,EMA(指数移动平均)更新一般也应该在权重更新步进行,而不是每个小 batch 都更新。我踩过的一个坑是:把 grad clip 放在累积循环内部,结果模型收敛很慢,去掉之后才恢复正常。这个代码片段看一下:

if (i + 1) % accum_steps == 0: # 在这里统一裁剪 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=10.0) scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_none=True)

另外,如果用了分布式训练(DDP),梯度累积时要注意 DDP 内部的梯度同步问题。在每次 backward 时,DDP 会触发跨卡通信和梯度 allreduce。如果你只在最后一步才想同步梯度,可以调用 model.no_sync() 来抑制中间步的同步,否则梯度累积会变成“每个小 batch 都同步一次”,通信开销巨大。正确写法用 context 包住中间步,只在最后一步正常反向。这里我不展开 DDP,但这个坑很常见,值得先记住。

4. 实操:在你的模型上实际应用(YOLO/EasyOCR/预训练微调)

4.1 一个通用的训练循环模板

我直接给一个可以复用的模板,集成 torch.compile、混合精度和梯度累积。这个模板主要针对单卡训练,用在 YOLO、EasyOCR 和分类模型上都可以,只需要替换 model、loss 计算部分。

import torch model = get_model() # 你自己的模型 model = model.cuda() model = torch.compile(model, mode="default") # 先不用 reduce-overhead optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-2) scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer, max_lr=1e-4, total_steps=10000) scaler = torch.cuda.amp.GradScaler() accum_steps = 4 grad_clip = 10.0 model.train() optimizer.zero_grad(set_to_none=True) for i, (images, targets) in enumerate(loader): images = images.cuda(non_blocking=True) targets = [t.cuda(non_blocking=True) for t in targets] with torch.autocast(device_type="cuda", dtype=torch.float16): loss_dict = model(images, targets) loss = loss_dict["loss"] loss = loss / accum_steps scaler.scale(loss).backward() if (i + 1) % accum_steps == 0: scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip) scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_none=True) scheduler.step()

注意几个点:我在每个小 batch 都除以 accum_steps,这样最后累积梯度就是平均梯度。GradScaler 在自动混合精度下,如果梯度出现 inf 或 nan,它会跳过优化器 step,所以累积步内如果某一步 loss 溢出,最终 step 会被跳过,不用手动处理。

4.2 在 YOLOv8 上具体怎么改

YOLOv8 的官方训练代码逻辑在 ultralytics 框架里,但你也可以直接用这个模板替换它的核心训练循环。如果你的环境里可以直接用训练 CLI,它本身已经集成了 AMP 和 torch.compile?实际上,Ultralytics 在 8.x 版本中确实支持torch.compile选项,比如yolo train model=yolov8s.pt compile=True。但默认没有开梯度累积。如果你想结合使用,最好写一个基于该框架的简单脚本。我的做法是:

先加载预训练模型,比如 yolov8s.pt,然后替换 data loader 的 batch size 为一个很小的值,再用我上面的模板重写 for 循环,把 model.forward 输出直接作为 loss。YOLOv8 的 loss 计算在模型内部完成,返回一个 Loss 对象,你需要从中取出总损失。

如果你不想重写整个框架,只想看看梯度累积的效果,可以先把 batch size 设成原来的一半,accum_steps=2,学习率不变。这是最简单的介入方式,也能明显看到显存占用下降,训练曲线略有变化。对于模型训练参数含义,我补充一下:YOLO 的batch参数是你每个 step 喂进去的图片张数,accumulate或accum参数就是累积步数。在 Ultralytics 中,默认会根据 batch size 自动计算一个累积步数,但它内部实现不一定和我上面的一致,具体要看源码。

4.3 EasyOCR 自定义模型训练时如何配置

EasyOCR 本质上是在训练一个文本检测器 + 文本识别器,常用的是 CRAFT 检测模型和 CRNN 识别模型。如果训练自己的 EasyOCR 模型,训练集通常是一批裁剪好的单行文本图片,长宽变化很大,需要做定长 padding 或者动态 batch 构建。对识别模型做 torch.compile 时,强烈建议采用动态 shape 或者固定 sequence length,否则每次遇到新长度都会重新编译。我在训练自己的 OCR 模型时,把图片高度固定为 64,宽度动态变化,但通过torch.randn(1, 3, 64, width)的形式模拟 batch,发现 torch.compile 在 dynamic 模式下可以工作,但速度提升没有定长情况明显。所以,如果你追求训练速度,尽量把 batch 内的图片 resize 成接近的宽高比,或者按宽度 bucket 分桶,让同一 batch 宽度一致,这样编译收益最大。

梯度累积在 OCR 任务中很有用,因为单张图片分辨率高,显存很快就被吃满。我常用的配置是:batch size 8,accum_steps 2,等效 batch 16,学习率从 5e-4 降到 3e-4 使用。虽然等效 batch 变大,但因为样本宽度不均,梯度噪声还是很大,所以学习率不要调太高。我建议在训练初期用一个小验证集快速测试 3 到 5 个 epoch,观察 loss 震荡幅度,再决定要不要动学习率。

4.4 微调 RoBERTa / ResNet 预训练模型的要点

对于中文 RoBERTa 这类 transformer 模型,torch.compile 的收益主要来自层内多个注意力头的并行融合。我的实测是,在单卡 A100 上微调 RoBERTa-base,开启 compile 后单步训练时间大约降低 20% 到 30%,但第一次编译需要 3 到 5 分钟。如果你做超参搜索,频繁变更模型结构或输入长度,编译缓存会失效,反而得不偿失。建议:先跑一个 epoch 把编译缓存保存下来,之后继续训练就很快。可以通过设置环境变量TORCHINDUCTOR_CACHE_DIR指定缓存目录。

ResNet 这类 CNN 模型,卷积算子被 torch.compile 融合的空间没有 transformer 大,因为 cuDNN 本身已经优化很好了,但残差连接、BN、ReLU 之间的内存读写优化还是有效果。更重要的是,微调 ResNet 时如果用了大量数据增强,输入 shape 固定,compile 很稳定。我习惯在微调时同时开启算术强度更高的 reduce-overhead 模式,因为 CNN 的 kernel 普遍比较大,CUDA graph 能减少 launch 延迟。

4.5 混合精度与 torch.compile 的配合

PyTorch 2.x 的 torch.compile 可以和 torch.cuda.amp 混合使用,但要注意几个 subtleties。第一,建议先开启 autocast,再调用 model。因为 autocast 是上下文环境,compile 会捕获算子类型,如果编译后你在外部切换 dtype,可能触发重新编译。所以最好把 autocast 和 model 调用都放在同一个作用域里。第二,GradScaler 和 compile 兼容,但如果你发现反向传播后梯度为 None,关闭 compile 后再试,可能是 torch.compile 在特定模型下对未使用的参数做了图裁剪,导致某些参数没有梯度。这个问题在带 frozen backbone 的微调任务中偶有发生,后面我会讲排查方法。

第三,如果你用的是 bfloat16,则不需要 GradScaler,直接with torch.autocast(device_type="cuda", dtype=torch.bfloat16)即可。BF16 在部分 GPU(如 A100、3090)上表现不错,但 20 系及以前的卡支持差,还是用 fp16。

5. 常见问题与排查技巧实录

5.1 torch.compile 首次编译很慢或卡住

这是新手问得最多的问题。torch.compile 首次编译确实慢,尤其是 reduce-overhead 和 max-autotune 模式,一个小模型也可能花 10 分钟。正常现象,不是死机。判断方法是观察 GPU 利用率很低但 CPU 利用率高,说明正在 auto tune。如果实在等不了,可以先用 default 模式跑,以后代码不变的情况下第二次启动会走缓存,编译时间大幅缩短。

如果开启 compile 后程序崩掉,多半是因为模型里有跟设备绑定、或者包含不可追踪的操作。排查方法是设置TORCH_COMPILE_DEBUG=1,会输出详细日志,或者在编译前把模型的一部分 disable 编译。也可以按模块粒度编译,只编译模型中的几个 main layer,把不兼容的层(如自定义 CUDA kernel)排除:

model.encoder = torch.compile(model.encoder) model.head = model.head # 不编译

5.2 梯度累积后 loss 曲线出现周期尖刺

经典现象:每 accum_steps 次迭代,loss 突然升高一次。原因大概率是学习率调度器只在更新步执行,而 loss 打印是在每个小 batch 都记录,当第 N 个小 batch 的累计梯度被应用后,权重发生较大变化,loss 自然突跳。这其实是正常的,但如果尖刺不是在你预期位置出现,就要检查 scheduler.step() 是否在累积步内被误调用了,导致部分小 batch 更新了优化器学习率,而权重没更新。

我在代码里特意把 scheduler.step() 放在 if 累积结束的块里,就是为了保证学习率只随权重更新步变化。还有,如果用 OneCycleLR,它的 step 是以更新步为单位的,不能混。

5.3 梯度累积后梯度 norm 异常变大/变小

如果你在累积内部用了 grad clip,会出现 norm 被 clip 到阈值后,后续小 batch 的梯度又叠加上去,导致最终 norm 比预期大。这就是为什么我强调要在累积结束后统一 clip。另一种情况是,你用了loss = loss / accum_steps,但代码不小心重复除以了多次。我见过有人先除以 accum_steps,又在 backward 里再除以一次,最后梯度变成原来的 1/accum_steps^2,模型训练很慢。

排查方法:在累积结束后、optimizer.step() 之前,打印模型某一层的梯度范数:

for name, p in model.named_parameters(): if p.grad is not None: print(name, p.grad.norm().item()) break

和单 batch 不除以 accum_steps 时对比一下,量级应该大约是原来的 1/accum_steps。如果差异明显,检查代码。

5.4 BN 层统计量在梯度累积中的偏差

如果你的模型使用 BatchNorm,且训练时 batch size 很小(比如 2 或 4),BN 本身的统计量就会很差,这时梯度累积并不能修复 BN。因为 BN 的 running_mean 是在每个小 batch 前向时更新的,它没有跨小 batch 汇总。可能出现训练 loss 正常、验证集精度上不去的问题。解决方案有几种:换用 GroupNorm 或 LayerNorm;或者增大单卡 batch size(如果你还有显存);或者使用torch.nn.SyncBatchNorm配合多卡。对于单卡,最稳妥还是尽量让单 batch 不低于 8,这样 BN 才可信。

5.5 torch.compile 后保存的模型权重无法加载

因为 torch.compile 返回的模型对象是一个 wrapper,直接torch.save(compiled_model.state_dict(), 'x.pt')保存的是优化后计算图的 buffer,里面可能包含额外键。正确做法是先暂时关闭编译,保存原始模型的状态:

torch.save({'model': model._orig_mod.state_dict(), 'optimizer': optimizer.state_dict()}, 'ckpt.pt')

这里_orig_mod是 torch.compile 内部保存原始模型属性的名字。如果你的代码里面用了深拷贝或者序列化,也可能因为 wrapper 导致失败。我在代码里设了一个分支,在保存 checkpoint 时使用model._orig_mod if hasattr(model, '_orig_mod') else model。

5.6 compile 后验证阶段显存不释放

有些用户在训练循环外单独跑验证集,发现显存占用不断上涨。这可能是验证阶段对模型输入做了不同 shape,触发新的编译,而旧编译缓存没有释放。建议把验证阶段也包进torch.no_grad()和torch.inference_mode(),并且最好在验证前先torch.cuda.empty_cache()。如果你用动态 shape,验证时 shape 变化导致重新编译,可以考虑固定验证集 batch size 和图片尺寸,或者把模型切回 eager 模式验证:

model.eval() with torch.inference_mode(): # 验证逻辑

这样至少不会把训练阶段的编译缓存污染了。

6. 从单卡扩展到多卡与预训练模型加载策略

6.1 DDP + 梯度累积的正确顺序

当 torch.compile 和 DDP 结合时,官方推荐先 compile 再包 DDP,也就是model = torch.compile(model); model = DDP(model, ...),或者反过来也行,但要注意顺序会影响编译缓存和通信。我更习惯先 DDP 再 compile,因为 DDP 包装后 model.module 是原始模型,compile 后_orig_mod嵌套关系更简单。但实际中,先 compile 再 DDP 能减少 DDP 的 hook 复制开销,理论上更快。

DDP 下梯度累积时,必须用model.no_sync()抑制中间步的梯度 allreduce:

context = model.no_sync() if (i + 1) % accum_steps != 0 else nullcontext() with context: loss = model(images, targets)["loss"] / accum_steps scaler.scale(loss).backward()

这样每次反向时,除了累积步的最后一步,其他步不会触发跨卡梯度同步。这个细节确实容易漏,一旦漏掉,多卡训练速度反而比单卡还慢,因为每步都在通信。

6.2 预训练模型的加载与 compile 缓存

以 yolo 预训练模型下载和 resnet 预训练模型为例,加载预训练权重时要注意:如果先用 torch.compile 再加载权重,要确认 state_dict 的 key 对齐。我建议在 compile 之前加载预训练权重,然后再 compile。这样 compile 只是包装了模型,权重不会被修改。

另外,torch.compile 的缓存与 Python 版本、PyTorch 版本、模型结构 hash 强相关。如果你同时跑多个项目,最好为每个项目设置不同的 TORCHINDUCTOR_CACHE_DIR,避免缓存冲突。比如:

export TORCHINDUCTOR_CACHE_DIR=/data/tmp_inductor_cache/yolo_v8

这样第二次启动同一个训练任务时,编译时间可以降到几秒。

6.3 超参数速查表

我根据自己的经验整理了一张表,适合单卡 12GB 显存、训练检测/识别模型时参考:

参数默认值开 torch.compile 后开梯度累积后组合使用建议
batch size161688
accum_steps1122
等效 batch size16161616
学习率1e-41e-41e-41.3e-4
warmup 步数1000 更新步1000按更新步换算250 更新步
编译模式-default-default
混合精度fp16fp16fp16fp16

注意,这个表以 12GB 显存为例,实际使用时你需要自己测一下峰值显存。我的经验是,开启 torch.compile 的 reduce-overhead 模式会增加约 5% 到 10% 显存,所以如果显存很紧张,先不用这个模式。

7. 我踩过的几个坑和最终推荐配置

7.1 坑:把 compile 放在 AMP 上下文里

我有一次把 torch.compile 放在了 autocast 的作用域内调用,结果编译后的模型行为怪怪的,loss 偶尔变成 nan。原因是 compile 捕获到的算子类型可能是 fp16,但后续 forward 进入了 fp32 上下文,导致 recompile 或类型不匹配。正确做法是:compile 在普通上下文中调用,forward 时再用 autocast 包住。代码结构:

model = torch.compile(model) # 不在 autocast 里 for step: with torch.autocast(...): loss = model(input)

7.2 坑:动态 shape 引发反复编译

我最初在 OCR 模型上没动 dynamic 设置,结果训练到第 10 个 epoch 后显存溢出,因为每出现新宽度都编译一份新 graph,缓存占满显存。后来用torch._dynamo.mark_dynamic(input, 3)标记宽度维度后,编译次数大幅下降。对于输入 [B, C, H, W],标记第 3 维动态。如果你不喜欢硬编码,也可以用dynamic=True并配合TORCH_LOGS="recompiles"观察重编译次数。理想情况下,整个训练过程重编译次数应该是个位数,否则就说明 shape 处理有问题。

7.3 坑:训练完保存的 model 是 compiled_module

这里我再强调一遍。使用 torch.compile 后,model变量类型变成了OptimizedModule,它的 state_dict 和原始模型不是一回事。我吃过亏,保存后加载,报 key mismatch。最后我用的是model._orig_mod.state_dict()。如果你实在不确定,可以打印model.state_dict().keys()看是否包含_orig_mod.前缀,如果有,就用_orig_mod重取。

7.4 我目前最常用的训练启动配置

以训练一个自定义 YOLOv8 模型为例,我最终采用的是:

batch_size 8 accumulate 4 amp true compile mode=default

学习率初始 0.002,warmup 3 个 epoch,之后 cos 衰减。这个配置在 8GB 显存的消费卡上可以稳定运行,等效 batch 32,训练曲线明显比 batch 8 直接训练平滑。速度方面,开了 compile 后单 iter 时间约 0.32 秒,不开约 0.41 秒,提升约 22%。显存峰值比不开编译时高约 300MB,但在 8GB 卡上还能接受。

如果你的卡是 24GB,我会推荐 batch 16、accumulate 2、compile mode=reduce-overhead。此时显存占用约 15GB,等效 batch 32,单步耗时比 eager 模式快接近 30%。

7.5 最后的小技巧:先用小模型验证流程再上全量

还有一个很实用的习惯:在正式启动长训练前,把模型换成一个很小的版本(比如 torchvision 的 resnet18 或者自己写一个 3 层 CNN),输入也换成一个很小的随机张量,跑 50 个 step 看看 loss 是否能正常下降。这样能迅速验证你的训练循环、compile 配置、梯度累积和保存逻辑是否正确,而不用等 YOLO 或者 OCR 大模型跑完一遍才知道问题。我自己所有实验都会做这一步骤,能省下很多时间。

把 torch.compile 和梯度累积放在一起,其实就是用编译加速把省下来的时间投入到更多的有效迭代中,同时用梯度累积把 batch size 的物理限制变成逻辑上的可调参数。这套组合拳几乎适用于所有 PyTorch 训练任务,只要留意动态 shape、BN 统计、学习率缩放和保存格式这些细节,就能稳定地看到显存压力减小、单步速度变快、loss 曲线更平滑。你也可以在自己的数据集上测试这两个特性,先用小模型跑通,再把它们集成到你现有的训练脚本里。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询