如何在PyTorch中避免原地操作同时降低GPU内存占用?
问题描述
在使用Transformer模型时遇到GPU内存不足的问题,原因是为了避免原地操作破坏autograd导致无法反向传播,将b += something这类原地操作改为了b = b + something的非原地操作,但非原地操作会复制张量数据,大幅增加内存占用。
尝试手动删除变量释放内存:
b_new = b + something del b torch.cuda.empty_cache()
但调用torch.cuda.memory_allocated()后发现内存并未释放,希望找到既能正常反向传播、又能拥有原地操作内存表现的解决方案。
示例代码如下:
from torch import nn import torch device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") X = torch.rand([500, 500, 512], device=device) class SimpleNetwork(nn.Module): def __init__(self): super().__init__() self.linears = nn.ModuleList([nn.Linear(512, 512) for i in range(10)]) self.mem_before = [] self.mem_after = [] def forward(self, x): for linear in self.linears: self.mem_before.append(torch.cuda.memory_allocated() /1024**2) # 对比非原地与原地操作 # x += linear(x) x = x + linear(x) self.mem_after.append(torch.cuda.memory_allocated() /1024**2) return x model = SimpleNetwork() model.to(device) for epoch in range(10): output = model(X) # 计算损失并反向传播 # 使用原地操作时无法执行反向传播
内存使用对比:
- 原地操作:内存占用稳定,无明显增长
- 非原地操作:内存随每一步计算持续上升
解决方案
1. 使用torch.add()指定out参数实现"伪原地"操作
PyTorch的torch.add()支持通过out参数指定输出张量,直接将结果写入原张量的内存空间,避免创建新张量,同时保留计算图用于反向传播:
# 替代 x = x + linear(x) torch.add(x, linear(x), out=x)
这种方式既复用了原张量的内存,又不会破坏autograd的计算追踪,能正常执行反向传播。
2. 配合垃圾回收与缓存清理
手动删除变量后,Python的垃圾回收机制可能不会立即释放内存,需显式触发垃圾回收再清理CUDA缓存:
b_new = b + something del b import gc gc.collect() # 触发Python垃圾回收 torch.cuda.empty_cache() # 清理CUDA未使用的缓存
注意:该方法仅能释放未被引用的张量内存,对仍在计算图中的张量无效。
3. 使用梯度检查点(Gradient Checkpointing)
通过torch.utils.checkpoint模块,正向传播时不存储所有中间激活值,反向传播时重新计算部分中间结果,大幅降低内存占用:
from torch.utils.checkpoint import checkpoint class SimpleNetwork(nn.Module): def __init__(self): super().__init__() self.linears = nn.ModuleList([nn.Linear(512, 512) for i in range(10)]) self.mem_before = [] self.mem_after = [] def forward_step(self, x, linear): self.mem_before.append(torch.cuda.memory_allocated() /1024**2) x = x + linear(x) self.mem_after.append(torch.cuda.memory_allocated() /1024**2) return x def forward(self, x): for linear in self.linears: x = checkpoint(self.forward_step, x, linear) return x
该方法会增加少量计算时间,但能显著减少内存占用,适合Transformer这类深层模型。
4. 混合精度训练
使用torch.cuda.amp将模型和张量转为半精度(FP16),在保证精度损失可控的前提下,将内存占用降低约一半:
from torch.cuda.amp import autocast, GradScaler scaler = GradScaler() model = SimpleNetwork().to(device) optimizer = torch.optim.Adam(model.parameters()) for epoch in range(10): optimizer.zero_grad() with autocast(): # 自动混合精度上下文 output = model(X) loss = ... # 定义你的损失函数 scaler.scale(loss).backward() # 缩放损失避免下溢 scaler.step(optimizer) scaler.update()
5. 梯度累积
如果batch size过大,可将多个小batch的梯度累积后再执行一次反向传播,等效于使用大batch size但内存占用更低:
accumulation_steps = 4 # 累积4个小batch的梯度 optimizer = torch.optim.Adam(model.parameters()) for epoch in range(10): optimizer.zero_grad() for step in range(accumulation_steps): output = model(X) loss = ... # 计算损失 loss = loss / accumulation_steps # 均分损失 loss.backward() # 累积梯度 optimizer.step() # 更新参数
内容的提问来源于stack exchange,提问作者Ciarán

