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(),需解答以下问题:
- 导致计算图被二次遍历的原因可能是什么?
- 结合克隆张量、分离张量或原地操作时,有哪些常见陷阱会引发该错误?
- 如何重构代码,使每个前向传播仅构建一个反向传播图,无需使用
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

