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

PyTorch中NTM控制器反向传播二次遍历图问题求解(无需retain_graph=True)

问题描述

训练PyTorch中带2D注意力融合控制器的NTM模型时,调用loss.backward()出现如下错误:

RuntimeError: Trying to backward through the graph a second time (or directly access saved tensors after they have already been freed). Saved intermediate values of the graph are freed when you call .backward() or autograd.grad(). Specify retain_graph=True if you need to backward through the graph a second time or if you need to access saved tensors after calling backward.

错误出现在NTM控制器MLP的反向传播阶段,使用retain_graph=True会导致内存暴涨,希望修复底层问题。控制器简化代码片段如下:

class NTMController(nn.Module):
"""
A fully advanced NTM controller that fuses 2D attention maps with the base representation.
Produces control signals (read/write keys, etc.) via an MLP.
"""
def __init__(self, d_model=128, mem_dim=128, hidden_dim=256, n_layers=3,
             fuse_in_channels=32, fuse_out_channels=32):
    super().__init__()
    self.d_model = d_model
    self.mem_dim = mem_dim

    # Layers for 2D attention fusion:
    self.conv_merge = nn.Conv2d(fuse_in_channels, fuse_out_channels, kernel_size=3, padding=1)
    self.resblock = BasicResBlock2D(fuse_out_channels, fuse_out_channels, stride=1)
    self.final_linear = nn.Linear(fuse_out_channels, fuse_out_channels)

    # MLP for generating control signals; input is base representation concatenated with fused 2D features.
    mlp_in_dim = d_model + fuse_out_channels
    layers = []
    in_dim = mlp_in_dim
    for _ in range(n_layers):
        layers.append(nn.Linear(in_dim, hidden_dim))
        layers.append(nn.ReLU())
        in_dim = hidden_dim
    # Final layer produces 4*mem_dim + 1 outputs.
    layers.append(nn.Linear(hidden_dim, 4 * mem_dim + 1))
    self.mlp = nn.Sequential(*layers)

def _fuse_2d_maps(self, attn_intra: Optional[torch.Tensor], attn_hier: Optional[torch.Tensor]) -> Optional[torch.Tensor]:
    if attn_intra is None and attn_hier is None:
        return None
    # If both maps are provided, interpolate to the same spatial size and concatenate along the channel dim.
    if attn_intra is not None and attn_hier is not None:
        B, C_i, Hi, Wi = attn_intra.shape
        B2, C_h, Hh, Wh = attn_hier.shape
        Hmax, Wmax = max(Hi, Hh), max(Wi, Wh)
        if (Hi, Wi) != (Hmax, Wmax):
            attn_intra = F.interpolate(attn_intra, size=(Hmax, Wmax), mode='bilinear', align_corners=False)
        if (Hh, Wh) != (Hmax, Wmax):
            attn_hier = F.interpolate(attn_hier, size=(Hmax, Wmax), mode='bilinear', align_corners=False)
        x = torch.cat([attn_intra, attn_hier], dim=1)
    else:
        x = attn_intra if attn_intra is not None else attn_hier

    x = F.relu(self.conv_merge(x))
    x = self.resblock(x)
    B, C, H, W = x.shape
    x_pool = F.adaptive_avg_pool2d(x, (1, 1)).view(B, C)
    return self.final_linear(x_pool)

def forward(self, reps: torch.Tensor,
            attn_hier: Optional[torch.Tensor] = None,
            attn_intra: Optional[torch.Tensor] = None) -> tuple:
    # Fuse attention maps into a feature vector.
    fused_2d = self._fuse_2d_maps(attn_intra, attn_hier)
    if fused_2d is None:
        fused_2d = torch.zeros(reps.size(0), self.final_linear.out_features, device=reps.device)
    # Concatenate the base representation with the fused 2D features.
    cat_inp = torch.cat([reps, fused_2d], dim=-1)
    # Use clone() here to avoid potential in-place modifications:
    out = self.mlp(cat_inp.clone())    # <--- is where the traceback points
    # Split the output into control signals.
    read_key, write_key, erase_raw, add_vec, scale_raw = torch.split(
        out, [self.mem_dim, self.mem_dim, self.mem_dim, self.mem_dim, 1], dim=-1
    )
    scale = torch.sigmoid(scale_raw).squeeze(-1)
    # For now, duplicate scale for read and write operations.
    return read_key, write_key, erase_raw, add_vec, scale, scale

错误具体出现在self.mlp(cat_inp.clone())处,且仅调用一次loss.backward(),需解答以下问题:

  1. 导致计算图被二次遍历的原因可能是什么?
  2. 结合克隆张量、分离张量或原地操作时,有哪些常见陷阱会引发该错误?
  3. 如何重构代码,使每个前向传播仅构建一个反向传播图,无需使用retain_graph=True?

1. 计算图二次遍历的可能原因

  • 输入张量复用:reps、attn_intra或attn_hier若在之前的前向传播中已参与过反向传播(比如多批次复用、循环重复使用),其计算图缓存已被释放,再次用于前向传播后反向会触发错误。
  • ResBlock原地操作:BasicResBlock2D若包含x += residual这类原地加法,会破坏原张量的计算图追踪,导致后续反向时试图访问已被修改的张量节点。
  • 输出张量重复关联:返回的scale被重复返回(return ... scale, scale),若后续代码对两个scale分别进行梯度操作,会导致同一计算图被两次遍历。
  • 克隆操作误用:cat_inp.clone()默认保留计算图,若cat_inp的计算图已部分释放,克隆后的张量反向时会试图访问已失效的节点。

2. 克隆、分离与原地操作的常见陷阱

  • 克隆保留计算图但未管理生命周期:使用clone()而不搭配detach()时,克隆张量仍依赖原张量的计算图。若原张量计算图已释放,克隆张量反向时会报错。
  • 分离张量不当接入:对张量使用detach()后又重新接入计算图,会导致计算图断裂,反向时要么丢失梯度,要么触发二次遍历错误。
  • 原地操作破坏计算图:x.add_()、x[:] = ...、x += y等原地修改操作会覆盖张量的data指针,Autograd无法追踪原始计算路径,反向时要么丢失梯度,要么触发已释放节点的访问错误。
  • 重复使用输出张量:同一输出张量多次返回或重复用于后续计算,会导致同一计算图节点关联到多个损失计算,触发二次反向遍历。

3. 代码重构方案

方案1:移除不必要的克隆操作

cat_inp是torch.cat生成的新张量,无原地修改风险,直接传入MLP即可:

out = self.mlp(cat_inp)  # 去掉.clone()

同时检查外部代码,确保reps、attn_intra、attn_hier仅单次参与前向传播,每次前向都使用新张量实例。

方案2:修复ResBlock的原地操作

若BasicResBlock2D包含原地操作,修改为非原地版本:

# 错误的原地操作示例
def forward(self, x):
    residual = x
    x = self.conv1(x)
    x = self.relu(x)
    x = self.conv2(x)
    x += residual  # 原地加法破坏计算图
    x = self.relu(x)
    return x

# 修复为非原地操作
def forward(self, x):
    residual = x
    x = self.conv1(x)
    x = self.relu(x)
    x = self.conv2(x)
    x = x + residual  # 生成新张量,保留计算图
    x = self.relu(x)
    return x

方案3:避免重复返回同一张量

将重复的scale返回改为生成新张量,避免同一计算图节点被多次关联:

return read_key, write_key, erase_raw, add_vec, scale.clone(), scale.clone()

或确保后续代码仅使用其中一个scale进行梯度相关操作。

方案4:显式切断旧计算图(可选)

若输入张量确实需要复用,每次前向传播前对其detach(),切断旧计算图关联:

def forward(self, reps: torch.Tensor, attn_hier: Optional[torch.Tensor] = None, attn_intra: Optional[torch.Tensor] = None) -> tuple:
    # 切断旧计算图,确保每次前向构建新图
    reps = reps.detach()
    if attn_hier is not None:
        attn_hier = attn_hier.detach()
    if attn_intra is not None:
        attn_intra = attn_intra.detach()
    # 后续逻辑不变
    fused_2d = self._fuse_2d_maps(attn_intra, attn_hier)
    ...

内容的提问来源于stack exchange,提问作者Foenkel

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 05:24:50