You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在保留单精度原始权重的同时用半精度节省内存?

问题:单精度权重模型中用半精度计算梯度以节省内存

我希望在训练采用单精度权重的模型时,通过半精度计算来节省内存。尝试使用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)**实现半精度梯度存储,同时保持权重为单精度。具体逻辑如下:

  1. 用torch.autocast维持前向传播的半精度计算。
  2. 给模型的每个参数注册梯度钩子,在反向传播生成梯度后立即将其转换为float16,减少内存占用。
  3. 优化器执行更新时,会自动将半精度梯度转换为与权重一致的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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.29 23:26:10