PyTorch中如何在优化器下一步访问自定义梯度范数参数
问题需求
希望基于前一步梯度的范数作为阈值对SGD的梯度进行裁剪,需要访问前一状态的梯度范数。现有代码已计算出current_norm,需将该值传递到下一步,作为torch.nn.utils.clip_grad_norm_的max_norm参数使用,如何实现?
原代码
model = Classifier(784, 125, 65, 10) criterion = torch.nn.CrossEntropyLoss() optimizer = torch.optim.SGD(model.parameters(), lr = 0.1) for epoch in range(epochs): correct, total, epoch_loss = 0, 0, 0.0 for images, labels in trainloader: images, labels = images.to(DEVICE), labels.to(DEVICE) optimizer.zero_grad() outputs = net(images) loss = criterion(outputs, labels) loss.backward() torch.nn.utils.clip_grad_norm_(net.parameters(), max_norm=1.0) for param in net.parameters(): if param.grad is not None: param.grad += torch.randn_like(param.grad) * noise_scale optimizer.step() current_norm = torch.max(torch.tensor([torch.norm(p.grad, 2) for p in net.parameters()])) # Metrics epoch_loss += loss total += labels.size(0) correct += (torch.max(outputs.data, 1)[1] == labels).sum().item() epoch_loss /= len(trainloader.dataset) epoch_acc = correct / total
解决方案
核心思路是初始化一个变量保存上一步的梯度范数,在每次迭代中先使用这个变量作为裁剪阈值,再更新该变量为当前计算出的梯度范数,供下一次迭代使用。
修改后的代码如下:
model = Classifier(784, 125, 65, 10) criterion = torch.nn.CrossEntropyLoss() optimizer = torch.optim.SGD(model.parameters(), lr = 0.1) # 初始化前一步梯度范数,首次迭代用默认值1.0,与模型保持同设备 prev_norm = torch.tensor(1.0, device=DEVICE) for epoch in range(epochs): correct, total, epoch_loss = 0, 0, 0.0 for images, labels in trainloader: images, labels = images.to(DEVICE), labels.to(DEVICE) optimizer.zero_grad() outputs = net(images) loss = criterion(outputs, labels) loss.backward() # 使用前一步的梯度范数作为本次裁剪的阈值 torch.nn.utils.clip_grad_norm_(net.parameters(), max_norm=prev_norm) for param in net.parameters(): if param.grad is not None: param.grad += torch.randn_like(param.grad) * noise_scale optimizer.step() # 计算当前梯度范数,更新prev_norm供下一次迭代使用 current_norm = torch.max(torch.tensor([torch.norm(p.grad, 2) for p in net.parameters()], device=DEVICE)) prev_norm = current_norm # Metrics epoch_loss += loss total += labels.size(0) correct += (torch.max(outputs.data, 1)[1] == labels).sum().item() epoch_loss /= len(trainloader.dataset) epoch_acc = correct / total
关键说明
- 初始化
prev_norm时指定与模型相同的设备(device=DEVICE),避免张量设备不匹配的运行错误 - 迭代流程严格遵循:用前一步范数裁剪→更新梯度→执行优化步→计算当前范数→更新保存的范数值,确保下一次迭代使用的是上一轮的梯度范数
内容的提问来源于stack exchange,提问作者S.Dasgupta
相关产品推荐
相关产品推荐

