PyTorch梯度累积实现引发原地操作错误,求原因分析
问题原因分析
你遇到的这个RuntimeError是因为梯度累积的实现逻辑完全错误,再加上误用retain_graph=True导致的:
- 你把多个batch的模型输出和标签拼接起来,等到累积步数才计算loss反向传播,但PyTorch的计算图是动态生成的,每次
net(inputs)都会构建新图,而optimizer.step()会更新模型参数,这会导致旧计算图里的参数变量版本被修改,反向传播时就会出现版本不匹配的错误。 retain_graph=True在这里完全多余,它会强制保留计算图,既浪费内存,又加剧了变量版本冲突的问题。- 梯度累积的核心是累积梯度,而非累积输出和标签,你搞反了实现方向。
正确的梯度累积实现方式
梯度累积的核心逻辑是:
- 每个batch前向传播计算loss,不清零梯度,直接反向传播让梯度累积在模型参数上。
- 当累积到指定步数时,执行
optimizer.step()更新参数,随后清零梯度。 - 每次累积周期结束后,无需保留计算图,正常释放即可。
修正后的代码
import torch import torchvision import torchvision.transforms as transforms transform = transforms.Compose( [transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))]) batch_size = 4 trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform) trainloader = torch.utils.data.DataLoader(trainset, batch_size=batch_size, shuffle=True, num_workers=2) testset = torchvision.datasets.CIFAR10(root='./data', train=False, download=True, transform=transform) testloader = torch.utils.data.DataLoader(testset, batch_size=batch_size, shuffle=False, num_workers=2) classes = ('plane', 'car', 'bird', 'cat', 'deer', 'dog', 'frog', 'horse', 'ship', 'truck') import torch.nn as nn import torch.nn.functional as F class Net(nn.Module): def __init__(self): super().__init__() self.conv1 = nn.Conv2d(3, 6, 5) self.pool = nn.MaxPool2d(2, 2) self.conv2 = nn.Conv2d(6, 16, 5) self.fc1 = nn.Linear(16 * 5 * 5, 120) self.fc2 = nn.Linear(120, 84) self.fc3 = nn.Linear(84, 10) def forward(self, x): x = self.pool(F.relu(self.conv1(x))) x = self.pool(F.relu(self.conv2(x))) x = torch.flatten(x, 1) # flatten all dimensions except batch x = F.relu(self.fc1(x)) x = F.relu(self.fc2(x)) x = self.fc3(x) return x net = Net() import torch.optim as optim criterion = nn.CrossEntropyLoss() optimizer = optim.SGD(net.parameters(), lr=0.001, momentum=0.9) # 定义梯度累积步数 accumulation_steps = 101 with torch.autograd.set_detect_anomaly(True): for epoch in range(2): # loop over the dataset multiple times running_loss = 0.0 for i, data in enumerate(trainloader, 0): inputs, labels = data outputs = net(inputs) loss = criterion(outputs, labels) # 缩放loss:累积N个batch的梯度,等价于batch_size*N的效果,需保持loss量级一致 loss = loss / accumulation_steps # 反向传播,累积梯度 loss.backward() # 达到累积步数,更新参数 if (i + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad() print(f'[{epoch + 1}, {i + 1:5d}] loss: {running_loss / accumulation_steps:.3f}') running_loss = 0.0 running_loss += loss.item() * accumulation_steps print('Finished Training')
额外说明
- 为什么要缩放loss?因为累积了
accumulation_steps个batch的梯度,相当于把batch_size放大了对应倍数,为了保持loss的量级和原batch训练一致,需要把每个batch的loss除以累积步数,这样最终的梯度更新效果和用大batch训练完全一致。 - 绝对不要拼接output和label,这种方式会生成超长计算图,既占用大量内存,又会因为参数更新导致旧计算图失效,完全是错误的实现思路。
内容的提问来源于stack exchange,提问作者JobHunter69
相关产品推荐
相关产品推荐

