混合精度训练中梯度dtype疑问与torch.cuda.amp.autocast机制探究
关于torch.cuda.amp混合精度训练的疑问
为探究torch.cuda.amp.autocast的工作原理,我开展了如下实验:
import torch import torch.nn as nn import torch.optim as optim class CustomModel(nn.Module): def __init__(self, input_size, hidden_size, num_classes): super(CustomModel, self).__init__() self.fc1 = nn.Linear(input_size, hidden_size) self.fc2 = nn.Linear(hidden_size, hidden_size) self.relu = nn.ReLU() self.fc3 = nn.Linear(hidden_size, num_classes) def forward(self, x): out = self.fc1(x) out = self.relu(out) out = self.fc2(out) out = self.fc3(out) return out # 假设X_train_tensor和y_train_tensor已定义并移至cuda input_size = X_train_tensor.shape[1] hidden_size = 16000 num_classes = 2000 model = CustomModel(input_size, hidden_size, num_classes).to('cuda') scaler = torch.cuda.amp.GradScaler() criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=0.001) num_epochs = 1 for epoch in range(num_epochs): optimizer.zero_grad() with torch.cuda.amp.autocast(dtype=torch.float16, enabled=True): outputs = model(X_train_tensor) loss = criterion(outputs, y_train_tensor) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() print(outputs.dtype) print(model.fc1.weight.grad.dtype) print(model.fc2.weight.grad.dtype) print("Done!")
运行后得到输出:
torch.float16 torch.float32 torch.float32 Done!
我对此感到困惑,不清楚混合精度训练中梯度应采用何种dtype;此外,若梯度均为float32,是否仍有必要使用GradScaler?恳请解答。
解答
1. 混合精度训练中梯度的dtype设定
在PyTorch的混合精度训练中,梯度默认以float32存储,这是设计上的合理选择:
float16的动态范围远小于float32,梯度值通常较小,用float16存储容易出现下溢(数值变成0),导致训练无法收敛。autocast仅控制前向传播中运算的dtype(比如线性层、激活等用float16加速),而反向传播时,PyTorch会自动将梯度计算的中间结果转换为float32,最终权重的梯度也保持为float32,以此保证梯度的数值稳定性。
你实验中看到梯度是float32,完全符合混合精度训练的规范,这是正常且预期的行为。
2. 梯度为float32时仍需使用GradScaler的原因
即使梯度最终是float32,GradScaler依然必不可少,核心原因是前向传播的loss是float16类型:
- 当用
float16计算loss时,loss的数值范围可能很小,直接反向传播会导致梯度下溢(因为float16无法表示极小的梯度值)。 GradScaler的作用是先将loss放大若干倍(比如2^k),让反向传播得到的梯度也对应放大,避免下溢;之后在更新权重前,再将缩放后的梯度缩回去(通过scaler.step(optimizer)自动处理),保证权重更新的正确性。- 即使梯度最终存储为
float32,如果不做缩放,前向传播的loss在转换为float16时已经丢失了部分精度,导致梯度计算出现数值问题,GradScaler正是用来解决这个问题的。
内容的提问来源于stack exchange,提问作者熊fiona
相关产品推荐
相关产品推荐

