升级SSH服务器后出现TBackward版本不匹配RuntimeError如何解决?
问题本质说明
你遇到的报错并非TBackward本身的版本问题,而是服务器更新后PyTorch版本升高,对梯度计算相关的inplace操作检查更加严格导致的。旧版本PyTorch未触发这类检查所以代码可以正常运行,TBackward是PyTorch内置的transpose算子反向传播模块,没有单独的安装包,所谓降级本质是降低用户环境下的PyTorch版本,不过更推荐优先修改代码解决问题,不会影响你其他依赖的版本兼容性。
优先方案:修改代码修复报错(无需降级)
你的代码存在两个核心问题,修复后即可在新版本PyTorch下正常运行:
- 损失计算逻辑错误:计算生成器G的损失时,你用的是加了
detach的fake样本计算的判别器输出,G根本无法拿到梯度,同时G的损失依赖的判别器D的前向结果,会因为你先更新D的参数被inplace修改,触发梯度检查报错。 - 训练更新顺序错误:先更新D的参数再计算G的反向传播,会直接破坏G的损失对应的计算图。
具体修改步骤如下:
- 调整损失计算函数,分开计算D和G需要的判别器输出:
def forward_n_get_loss(real_img, cond_img, G, D, c=100): fake_img = G(cond_img) real_pair = torch.cat((real_img, cond_img), 1) fake_pair = torch.cat((fake_img, cond_img), 1) # 计算D损失:fake样本加detach,不更新G的梯度 prob_real = D(real_pair) prob_fake_D = D(fake_pair.detach()) loss_D_real = nn.BCELoss()(prob_real, torch.ones_like(prob_real)) loss_D_fake = nn.BCELoss()(prob_fake_D, torch.zeros_like(prob_fake_D)) loss_D = (loss_D_real + loss_D_fake) * 0.5 # 计算G损失:重新过判别器,不加detach,保留G的梯度 prob_fake_G = D(fake_pair) loss_G_fake = nn.BCELoss()(prob_fake_G, torch.ones_like(prob_fake_G)) loss_G_FLIP_l1 = FLIP_l1_loss(fake_img, real_img) loss_G = loss_G_fake + c*loss_G_FLIP_l1 return loss_D, loss_G, loss_D_real, loss_D_fake, loss_G_fake, loss_G_FLIP_l1
- 调整训练循环的更新顺序,先算完两个损失的反向传播再更新参数:
# 计算D的梯度 if update_D: optimizer_D.zero_grad() losses[0].backward(retain_graph=True) # 计算G的梯度 if update_G: optimizer_G.zero_grad() losses[1].backward() # 无需再保留计算图 # 统一更新参数 if update_D: optimizer_D.step() if update_G: optimizer_G.step()
备选方案:降级PyTorch版本
如果你确定需要降级到旧版本PyTorch兼容原有代码,可以直接在用户权限下安装对应旧版本,不会影响系统全局环境:
- 首先确认你服务器的CUDA版本,执行
nvcc -V即可查看 - 执行对应pip命令安装旧版本,比如适配CUDA11.0的PyTorch1.7.1版本命令如下:
pip install --user torch==1.7.1+cu110 torchvision==0.8.2+cu110 torchaudio==0.7.2 -f https://download.pytorch.org/whl/torch_stable.html
如果你的CUDA版本是10.2,把命令里的cu110替换为cu102即可。
3. 安装完成后执行以下命令验证版本:
import torch print(torch.__version__)
输出对应安装的旧版本号即为安装成功。
内容的提问来源于stack exchange,提问作者JopBeuger
相关产品推荐
相关产品推荐

