PyTorch 新手到老手都容易踩的坑:梯度、显存、多卡三座大山 PyTorch 新手到老手都容易踩的坑梯度、显存、多卡三座大山一、个性化深度引言上午还在跟同事说这轮训练稳了下午 OOM 了。不是 batch 太大是梯度累积到backward()时显存碎片化扛不住了。你说你用的是 PyTorch 自动求导但它的计算图默默地保存着你根本不需要的中间变量。PyTorch 的灵活性是把双刃剑。它可以让你随意构建动态计算图也可以让你随意地浪费显存、错误地累积梯度、混乱地分配多卡任务。这三座大山——梯度、显存、多卡——横亘在每一个从实验走向生产的 PyTorch 开发者面前。见证奇迹的时刻是你终于理解了torch.no_grad()和model.eval()不是可选的装饰而是显存管理的生死线。是你发现把loss.backward()放在循环里和放在循环外显存占用差了一个数量级。二、个性化原理剖析三座大山的底层逻辑是互相关联的。梯度计算的动态图如果不主动释放会保留到下一次backward()之前。这意味着在训练循环中做验证、做推理、做任何不需要梯度的操作时之前计算图占用的显存都不会被回收。这是看似突然 OOM的根本原因。三、个性化代码实践import torch import torch.nn as nn import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DataLoader, DistributedSampler import gc # # 第一座大山梯度 # def gradient_mistake_demo(): 设计原因展示最常见的四种梯度错误 model nn.Linear(10, 2) # 错误1: 梯度未清零——每次 backward 会累加到 .grad 上 # 设计原因PyTorch 默认累加梯度是为了方便 RNN 的 BPTT optimizer torch.optim.SGD(model.parameters(), lr0.01) for epoch in range(2): for batch in range(3): x torch.randn(32, 10) loss model(x).sum() loss.backward() # 忘记 optimizer.zero_grad()——梯度会累加3次 optimizer.step() # 正确做法 for batch in range(3): x torch.randn(32, 10) optimizer.zero_grad() # 设计原因必须放在 forward 之前 loss model(x).sum() loss.backward() optimizer.step() # 错误2: requires_grad 污染 # 设计原因任何 requires_gradTrue 的张量参与的运算都会创建计算图节点 a torch.randn(10, requires_gradTrue) b a * 2 # b.requires_grad True c b.detach() # 显式切断梯度c.requires_grad False # 设计原因使用 with torch.no_grad() 包裹所有不需要梯度的操作 with torch.no_grad(): logits model(torch.randn(1, 10)) pred logits.argmax(dim1) # 错误3: 梯度累积时忘记缩放 loss # 设计原因每步 backward 后 loss 应该除以累积步数 # 否则梯度量级会放大 accumulation_steps 倍 accumulation_steps 4 for i, batch in enumerate(range(8)): loss model(torch.randn(32, 10)).sum() # 设计原因scaled_loss 保持梯度量级一致 (loss / accumulation_steps).backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad() # 错误4: 在 backward 后持有计算图引用 # 设计原因var 持有计算图节点引用阻止显存回收 # 解决使用 var.detach() 或 var.item() 获取值后释放引用 # # 第二座大山显存 # class MemoryTracker: 设计原因封装显存监控方便定位泄漏点 staticmethod def report(): if torch.cuda.is_available(): allocated torch.cuda.memory_allocated() / 1024**3 reserved torch.cuda.memory_reserved() / 1024**3 max_allocated torch.cuda.max_memory_allocated() / 1024**3 # 设计原因reserved - allocated 碎片/缓存 print(fAllocated: {allocated:.2f}GB | Reserved: {reserved:.2f}GB | fPeak: {max_allocated:.2f}GB | Fragmented: {reserved-allocated:.2f}GB) staticmethod def reset_peak(): torch.cuda.reset_peak_memory_stats() def memory_pitfalls(): 设计原因常见显存陷阱及解决方案 # 陷阱1: 保留中间激活 # 设计原因默认 retain_graphFalse 会在 backward 后释放中间结果 # 如果 grad_output 不是标量需要传入 grad_tensors x torch.randn(100, 100, requires_gradTrue) y x.sum() y.backward() # 标量图自动释放 # 陷阱2: 显存碎片化 # 设计原因大量小张量的分配/释放导致碎片使用 empty_cache 整理 # 但不建议频繁调用——它本身也有开销 # torch.cuda.empty_cache() # 陷阱3: DataLoader 的 pin_memory 占用额外显存 # 设计原因pin_memoryTrue 在 CPU 锁页不占 GPU 显存 # 但 non_blockingTrue 的传输会短暂占用 # 陷阱4: checkpoint 保存时持有模型引用 # 设计原因torch.save 不会自动释放显存保存后马上 del state {model: model.state_dict()} torch.save(state, checkpoint.pt) del state # 显式释放 # 陷阱5: 列表累积未 detach 的张量 # 设计原因list.append(tensor) 持有引用阻止计算图释放 losses [] for _ in range(100): l torch.randn(1, requires_gradTrue) losses.append(l.item()) # .item() 返回 Python 标量不持有引用 # # 第三座大山多卡 # class MultiGPUManager: 设计原因多卡训练的配置是高度场景化的 这里提供一个基础配置模板注释标注了每个选择的理由。 staticmethod def setup_ddp(): 设计原因DDP 是当前多卡训练的标准方案 # 设计原因NCCL 后端在 GPU 间通信最快GLOO 用于 CPU dist.init_process_group(backendnccl) local_rank int(os.environ.get(LOCAL_RANK, 0)) torch.cuda.set_device(local_rank) return local_rank staticmethod def create_model_and_loader(model, dataset, local_rank, batch_size): 设计原因多个容易踩坑的细节集中处理。 # 设计原因模型先 to device 再包装 DDP避免设备错乱 model model.to(local_rank) # 设计原因find_unused_parametersFalse 提升性能 # 但如果模型有未参与 loss 的参数会报错 model DDP(model, device_ids[local_rank], find_unused_parametersFalse) # 设计原因DistributedSampler 保证每张卡看到不重叠的数据 # shuffleTrue 是必须的否则每个 epoch 每张卡看相同数据 sampler DistributedSampler(dataset, shuffleTrue) loader DataLoader( dataset, batch_sizebatch_size, samplersampler, num_workers4, pin_memoryTrue, # 设计原因drop_lastTrue 避免最后 batch 不整除导致的 # all-reduce 阻塞 drop_lastTrue ) return model, loader, sampler staticmethod def train_epoch(model, loader, optimizer, sampler, epoch): 设计原因DDP 训练的 epoch 模板 # 设计原因每个 epoch 必须调用 set_epoch # 否则每张卡每个 epoch 的 shuffle 结果是一样的 sampler.set_epoch(epoch) model.train() for batch_idx, (data, target) in enumerate(loader): data, target data.cuda(), target.cuda() optimizer.zero_grad() output model(data) loss nn.CrossEntropyLoss()(output, target) loss.backward() optimizer.step() # 设计原因只在 rank 0 打印避免刷屏 if dist.get_rank() 0 and batch_idx % 100 0: print(fEpoch {epoch} Batch {batch_idx} Loss {loss.item():.4f}) # # 综合诊断工具 # class PyTorchDiagnostics: 设计原因一次性诊断脚本快速定位三座大山的问题 staticmethod def diagnose_gradient(model: nn.Module): 设计原因检查梯度是否正常 issues [] for name, param in model.named_parameters(): if param.requires_grad and param.grad is not None: grad_norm param.grad.norm().item() if grad_norm 100: issues.append(f{name}: 梯度爆炸({grad_norm:.1f})) elif grad_norm 1e-7: issues.append(f{name}: 梯度消失({grad_norm:.2e})) if torch.isnan(param.grad).any(): issues.append(f{name}: 梯度包含 NaN) return issues staticmethod def diagnose_memory(): 设计原因显存快照诊断 if not torch.cuda.is_available(): return {error: CUDA not available} return { allocated_gb: torch.cuda.memory_allocated() / 1024**3, reserved_gb: torch.cuda.memory_reserved() / 1024**3, max_allocated_gb: torch.cuda.max_memory_allocated() / 1024**3, fragmentation: (torch.cuda.memory_reserved() - torch.cuda.memory_allocated()) / 1024**3, device_count: torch.cuda.device_count(), current_device: torch.cuda.current_device() }四、个性化边界权衡梯度裁剪 vs 不加裁剪裁剪防止梯度爆炸训练稳定。但裁剪阈值是敏感参数过大无效过小有害。不加裁剪保留原始梯度信息但在 RNN/Transformer 训练中可能爆炸。实际选择默认开启 clip_grad_norm_阈值从 1.0 开始调。观察到 loss 震荡时降低阈值。checkpoint 保存频率每 N 步保存丢失的进度有限但磁盘 I/O 频繁影响训练速度。每 epoch 保存I/O 开销小但一旦中断丢失一个完整 epoch。实际选择每 1000 步 每 epoch 保存。保留最近 3 个 checkpoint 做滚动清理。DDP vs FSDP 选型DDP通信高效但需要每张卡能装下完整模型。适合模型 10B 参数。FSDP单卡放不下时必须用但通信开销大配置复杂。实际选择7B 以下用 DDP7B-70B 用 FSDP CPU Offload70B 以上用模型并行 流水线并行。五、总结PyTorch 的三座大山是互相关联的梯度计算图的保留直接导致显存膨胀显存的碎片化在多卡场景下被放大多卡通信的开销又反向影响梯度同步的效率。解决之道在于理解每个操作的显存生命周期——backward()后的图何时释放、no_grad()的作用域覆盖了哪些操作、DDP 的 gradient reduction 在哪个时机触发。通过torch.cuda.memory_summary()跟踪显存分配通过梯度范数检查定位爆炸/消失通过nvidia-smi观察多卡间的负载均衡是诊断这三座大山的基础手段。框架的灵活性赋予开发者控制力但控制力需要精确操作来兑现。

本月热点