使用UNet3+触发RuntimeError:梯度变量遭原地操作修改
问题定位与解决方案
错误根源
你遇到的RuntimeError是因为原地操作(inplace operation)破坏了PyTorch Autograd的梯度追踪机制。模型中多处使用了inplace=True的ReLU和Dropout,这些操作会直接修改输入张量的内存,而梯度计算需要保留操作前的张量版本,导致依赖链中的张量版本不匹配(错误提示中的MulBackward0对应_dot_product里的乘法操作,其输出依赖的张量被原地修改)。
另外测试代码中存在变量名笔误,会导致额外的运行错误。
具体修改步骤
1. 移除所有原地操作参数
将模型中所有inplace=True的设置删除:
Encoder类:
class Encoder(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.relu = nn.ReLU() # 移除inplace=True self.conv_1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False) self.norm_1 = nn.BatchNorm2d(out_channels) self.conv_2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False) self.norm_2 = nn.BatchNorm2d(out_channels)Decoder类:
class Decoder(nn.Module): def __init__(self, in_channels, out_channels, down=None, up=None): super(Decoder, self).__init__() layers = [] if down: layers.append(nn.MaxPool2d(kernel_size=down, stride=down)) elif up: layers.append(nn.Upsample(scale_factor=up, mode='bilinear')) layers.extend([ nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False), nn.BatchNorm2d(out_channels), nn.ReLU() # 移除inplace=True ]) self.decoder = nn.Sequential(*layers)cls分支:
self.cls = nn.Sequential( nn.Dropout(p=0.5), # 移除inplace=True nn.Conv2d(make_divisible(1024 * channel_ratio), 2, kernel_size=1, stride=1, padding=0), nn.AdaptiveMaxPool2d(1), nn.Sigmoid() )
2. 修正测试代码的变量名笔误
测试代码中pre_data = model(in_data)和loss = loss_fc(out_data, pred_data)变量名不一致,修改为:
pred_data = model(in_data) optimizer.zero_grad() with torch.autograd.set_detect_anomaly(True): loss = loss_fc(out_data, pred_data) loss.backward()
3. 可选:优化张量操作的变量命名(避免潜在问题)
在_dot_product方法中,避免复用变量名,明确命名中间张量:
@staticmethod def _dot_product(seg, cls): b, c, h, w = seg.size() seg_flat = seg.view(b, c, h * w) # 改用新变量名 x = torch.einsum("ijk,ij->ijk", [seg_flat, cls]) x = x.view(b, c, h, w) return x
原理说明
PyTorch的Autograd通过追踪张量的操作历史来计算梯度,原地操作会直接覆盖原始张量的内存,导致梯度计算时无法找到操作前的张量版本(错误提示中的“当前版本为1;预期版本为0”就是这个原因)。移除inplace=True后,所有操作都会返回新的张量,保留完整的梯度依赖链,从而解决错误。
内容的提问来源于stack exchange,提问作者Big-Yellow-J
相关产品推荐
相关产品推荐

