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

PyTorch梯度计算报错:原地操作修改梯度变量问题排查

PyTorch原地操作引发梯度计算错误的解决方案

问题描述

运行PyTorch代码时,在loss.backward()执行阶段报错:

one of the variables needed for gradient computation has been modified by an inplace operation

已尝试在loss.backward()中添加retain_graph=True,但问题未解决。

相关代码

main.py

dataset = Data(params)
detector = Detector(params)
optimizer = torch.optim.Adam(detector.parameters(), lr=params['learning_rate'])

Loss = []  
criterion = torch.nn.MSELoss()
for epoch in range(params['maxEpoch']):
    y, h_a, h_b, plus, hTy, hTh = dataset.generate()
    x_ = detector(hTy, hTh)

    loss = 0.0
    optimizer.zero_grad()
    for _ in range(1, params['DetNet_layer']):
        loss += criterion(x_[:,:,_], torch.from_numpy(plus).to(torch.double)) * math.log(_)

    loss.backward(retain_graph=True)
    optimizer.step()

    Loss.append(loss.item())

Detector类forward方法

def forward(self, HTy, HTH):
    HTy_torch = torch.from_numpy(HTy).unsqueeze(1)
    HTH_torch = torch.from_numpy(HTH).unsqueeze(1)

    x_torch = torch.from_numpy(np.zeros((self.batch_size, 1, self.L)))
    v_torch = torch.from_numpy(np.zeros((self.batch_size, 1, self.L)))

    for i in range(1, self.L):
        x_tmp, v_tmp = self.layers[i](HTy_torch, HTH_torch, x_torch[:, :, i-1], v_torch[:, :, i-1])
        x_torch[:, :, i] = x_tmp
        v_torch[:, :, i] = v_tmp

    return x_torch

问题原因

确实是切片赋值操作x_torch[:, :, i] = x_tmp和v_torch[:, :, i] = v_tmp导致的原地操作。这种直接修改张量切片的方式会覆盖原始张量的部分数据,破坏了PyTorch梯度计算依赖的计算图结构——梯度追踪需要记录张量的完整修改路径,原地操作会打断这条路径,导致反向传播时找不到正确的梯度来源。

retain_graph=True在这里无效,因为问题根源不是计算图被提前释放,而是计算图本身被原地操作破坏了。

修改方案

核心思路:用列表收集中间结果,最后拼接成张量

避免预先创建全零张量再原地修改,改用列表保存每一步的x和v输出,最后通过torch.cat拼接成最终张量,这样所有中间张量都能被正确追踪梯度。

修改后的Detector forward方法:

def forward(self, HTy, HTH):
    # 转换为PyTorch张量并统一dtype(和后续计算匹配)
    HTy_torch = torch.from_numpy(HTy).unsqueeze(1).double()
    HTH_torch = torch.from_numpy(HTH).unsqueeze(1).double()

    # 初始化第一步的x和v,直接用PyTorch创建张量,保证梯度追踪
    x_init = torch.zeros((self.batch_size, 1, 1), dtype=torch.double)
    v_init = torch.zeros((self.batch_size, 1, 1), dtype=torch.double)
    x_list = [x_init]
    v_list = [v_init]

    for i in range(1, self.L):
        x_tmp, v_tmp = self.layers[i](HTy_torch, HTH_torch, x_list[-1], v_list[-1])
        # 确保输出张量的维度和列表中元素一致(保持[batch,1,1]以便拼接)
        x_tmp = x_tmp.unsqueeze(2) if x_tmp.dim() == 2 else x_tmp
        v_tmp = v_tmp.unsqueeze(2) if v_tmp.dim() == 2 else v_tmp
        x_list.append(x_tmp)
        v_list.append(v_tmp)

    # 在维度2上拼接所有中间结果,得到最终的x_torch
    x_torch = torch.cat(x_list, dim=2)
    return x_torch

额外优化点

  1. 移除retain_graph=True:每个epoch都会重新构建计算图,不需要保留上一轮的计算图,去掉该参数能节省内存。
  2. 统一张量dtype:确保所有参与计算的张量dtype一致(比如都用double),避免隐式类型转换带来的问题。
  3. 提前转换plus张量:在main.py中可以提前把plus转为PyTorch张量,避免每次循环重复转换:
    # 在epoch循环内generate之后执行
    plus_tensor = torch.from_numpy(plus).to(torch.double)
    # 循环内使用plus_tensor
    loss += criterion(x_[:,:,_], plus_tensor) * math.log(_)
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 19:50:26