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

PyTorch梯度累积实现引发原地操作错误,求原因分析

问题原因分析

你遇到的这个RuntimeError是因为梯度累积的实现逻辑完全错误,再加上误用retain_graph=True导致的:

  • 你把多个batch的模型输出和标签拼接起来,等到累积步数才计算loss反向传播,但PyTorch的计算图是动态生成的,每次net(inputs)都会构建新图,而optimizer.step()会更新模型参数,这会导致旧计算图里的参数变量版本被修改,反向传播时就会出现版本不匹配的错误。
  • retain_graph=True在这里完全多余,它会强制保留计算图,既浪费内存,又加剧了变量版本冲突的问题。
  • 梯度累积的核心是累积梯度,而非累积输出和标签,你搞反了实现方向。
正确的梯度累积实现方式

梯度累积的核心逻辑是:

  1. 每个batch前向传播计算loss,不清零梯度,直接反向传播让梯度累积在模型参数上。
  2. 当累积到指定步数时,执行optimizer.step()更新参数,随后清零梯度。
  3. 每次累积周期结束后,无需保留计算图,正常释放即可。
修正后的代码
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 02:05:34