如何在保留单精度原始权重的同时用半精度节省内存?
问题:单精度权重模型中用半精度计算梯度以节省内存
我希望在训练采用单精度权重的模型时,通过半精度计算来节省内存。尝试使用autocast后,模型确实能以半精度进行预测,但生成的梯度仍为单精度,既影响性能又无法达到内存节省的目的。请问是否有办法让PyTorch以半精度计算梯度,并使用这些梯度来更新原始的单精度权重?
原测试代码
import torch class KekNet (torch.nn.Module): def __init__(self): super(KekNet, self).__init__() self.layer1 = torch.nn.Conv2d(3, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), dtype=torch.float32) def forward(self, x, features=False): return self.layer1(x) device = torch.device("cuda") # HALF-DATA AUTOCAST net = KekNet().to(device) loss_l2 = torch.nn.MSELoss(reduction='none') g_params = [{'params': net.parameters(), 'weight_decay': 0}] optimizerG = torch.optim.RMSprop(g_params, lr=3e-5, alpha=0.99, eps=1e-07, weight_decay=0) schedulerG = torch.optim.lr_scheduler.CosineAnnealingLR(optimizerG, T_max=300) X = torch.randn((40,3,555,555), dtype=torch.float16, device =device) with torch.autocast(device_type='cuda', dtype=torch.float16): Y_h=net(X) Y = torch.randn_like(Y_h) loss = loss_l2(Y_h, Y).mean() loss.backward() print(f"-autocast\r\ndata precision: {X.dtype}\r\npred precision: {Y_h.dtype}\r\ngrad precision: {net.layer1.weight.grad.dtype}\r\n") optimizerG.step() schedulerG.step()
原运行结果
data precision: torch.float16 pred precision: torch.float16 grad precision: torch.float32
解决方案
可以通过**梯度钩子(Gradient Hook)**实现半精度梯度存储,同时保持权重为单精度。具体逻辑如下:
- 用
torch.autocast维持前向传播的半精度计算。 - 给模型的每个参数注册梯度钩子,在反向传播生成梯度后立即将其转换为
float16,减少内存占用。 - 优化器执行更新时,会自动将半精度梯度转换为与权重一致的
float32,不影响权重更新的精度。
修改后的代码
import torch class KekNet(torch.nn.Module): def __init__(self): super(KekNet, self).__init__() self.layer1 = torch.nn.Conv2d(3, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), dtype=torch.float32) def forward(self, x, features=False): return self.layer1(x) device = torch.device("cuda") net = KekNet().to(device) # 注册梯度钩子,将梯度转为float16 for param in net.parameters(): param.register_hook(lambda grad: grad.half()) loss_l2 = torch.nn.MSELoss(reduction='none') g_params = [{'params': net.parameters(), 'weight_decay': 0}] optimizerG = torch.optim.RMSprop(g_params, lr=3e-5, alpha=0.99, eps=1e-07, weight_decay=0) schedulerG = torch.optim.lr_scheduler.CosineAnnealingLR(optimizerG, T_max=300) X = torch.randn((40,3,555,555), dtype=torch.float16, device=device) with torch.autocast(device_type='cuda', dtype=torch.float16): Y_h = net(X) Y = torch.randn_like(Y_h) loss = loss_l2(Y_h, Y).mean() loss.backward() print(f"-autocast\r\ndata precision: {X.dtype}\r\npred precision: {Y_h.dtype}\r\ngrad precision: {net.layer1.weight.grad.dtype}\r\n") optimizerG.step() schedulerG.step()
修改后运行结果
data precision: torch.float16 pred precision: torch.float16 grad precision: torch.float16
注意事项
- 若训练中出现半精度梯度的数值溢出问题,可以配合
torch.cuda.amp.GradScaler使用,先缩放梯度再转换为float16。 - 部分对数值精度敏感的任务,需要验证半精度梯度是否会影响最终训练效果。
内容的提问来源于stack exchange,提问作者Anonymous
相关产品推荐
相关产品推荐

