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(), lr=0.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_grad=True 的张量参与的运算都会创建计算图节点 a = torch.randn(10, requires_grad=True) 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(dim=1) # 错误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(f"Allocated: {allocated:.2f}GB | Reserved: {reserved:.2f}GB | " f"Peak: {max_allocated:.2f}GB | Fragmented: {reserved-allocated:.2f}GB") @staticmethod def reset_peak(): torch.cuda.reset_peak_memory_stats() def memory_pitfalls(): """设计原因:常见显存陷阱及解决方案""" # 陷阱1: 保留中间激活 # 设计原因:默认 retain_graph=False 会在 backward 后释放中间结果 # 如果 grad_output 不是标量,需要传入 grad_tensors x = torch.randn(100, 100, requires_grad=True) y = x.sum() y.backward() # 标量,图自动释放 # 陷阱2: 显存碎片化 # 设计原因:大量小张量的分配/释放导致碎片,使用 empty_cache 整理 # 但不建议频繁调用——它本身也有开销 # torch.cuda.empty_cache() # 陷阱3: DataLoader 的 pin_memory 占用额外显存 # 设计原因:pin_memory=True 在 CPU 锁页,不占 GPU 显存 # 但 non_blocking=True 的传输会短暂占用 # 陷阱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_grad=True) losses.append(l.item()) # .item() 返回 Python 标量,不持有引用 # =========================== # 第三座大山:多卡 # =========================== class MultiGPUManager: """ 设计原因:多卡训练的配置是高度场景化的, 这里提供一个基础配置模板,注释标注了每个选择的理由。 """ @staticmethod def setup_ddp(): """设计原因:DDP 是当前多卡训练的标准方案""" # 设计原因:NCCL 后端在 GPU 间通信最快,GLOO 用于 CPU dist.init_process_group(backend='nccl') 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_parameters=False 提升性能, # 但如果模型有未参与 loss 的参数会报错 model = DDP(model, device_ids=[local_rank], find_unused_parameters=False) # 设计原因:DistributedSampler 保证每张卡看到不重叠的数据 # shuffle=True 是必须的,否则每个 epoch 每张卡看相同数据 sampler = DistributedSampler(dataset, shuffle=True) loader = DataLoader( dataset, batch_size=batch_size, sampler=sampler, num_workers=4, pin_memory=True, # 设计原因:drop_last=True 避免最后 batch 不整除导致的 # all-reduce 阻塞 drop_last=True ) 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(f'Epoch {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 以下用 DDP,7B-70B 用 FSDP + CPU Offload,70B 以上用模型并行 + 流水线并行。
五、总结
PyTorch 的三座大山是互相关联的:梯度计算图的保留直接导致显存膨胀,显存的碎片化在多卡场景下被放大,多卡通信的开销又反向影响梯度同步的效率。解决之道在于理解每个操作的显存生命周期——backward()后的图何时释放、no_grad()的作用域覆盖了哪些操作、DDP 的 gradient reduction 在哪个时机触发。通过torch.cuda.memory_summary()跟踪显存分配,通过梯度范数检查定位爆炸/消失,通过nvidia-smi观察多卡间的负载均衡,是诊断这三座大山的基础手段。框架的灵活性赋予开发者控制力,但控制力需要精确操作来兑现。
