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

模型跨设备迁移后,backward指定inputs参数无法正常工作的原因

问题代码

import torch
import torch.nn as nn

class Model(nn.Module):
    def __init__(self):
        super().__init__()
        self.weight_mul = nn.Parameter(torch.randn(D,))
        self.weight = nn.Parameter(torch.randn(D,))

    def forward(self, x):
        x = x * self.temp_weight
        return x

D = 5

x = torch.randn(D,).cuda()
model = Model()
model.cuda()
model.temp_weight = model.weight * model.weight_mul
model.cpu(); model.cuda()
output = model(x)
output.sum().backward(inputs=[model.weight, model.weight_mul])

print(model.weight.grad)
print(model.weight_mul.grad)

运行后输出的.grad均为None,但移除backward()的inputs参数,或者移除model.cpu(); model.cuda()设备迁移步骤,就能让反向传播正常工作,请问原因是什么?


原因解析

1. 设备迁移导致计算图与当前模型参数断裂

当执行model.cpu(); model.cuda()时,PyTorch会为模型的每个Parameter创建新的张量实例(设备迁移操作会生成新张量,原张量被丢弃)。而model.temp_weight是在第一次model.cuda()后,用当时的参数张量计算得到的,它的计算图依赖的是旧的参数张量,而非设备迁移后的新参数。

此时你在backward()中指定inputs为当前模型的新参数,这些参数根本不在output的计算图路径上,自然无法计算梯度,最终.grad为None。

2. 移除inputs参数能工作的原因

如果不指定inputs,backward()会默认遍历整个计算图,计算所有requires_grad=True的叶子节点的梯度——也就是旧参数张量的梯度。不过这些旧参数已经不再是模型的属性(模型现在持有新的参数张量),但此时能得到非None的梯度,所以看起来“正常工作”。

3. 移除设备迁移步骤能工作的原因

如果跳过model.cpu(); model.cuda(),model.temp_weight的计算图依赖的就是当前模型持有的参数张量,此时在backward()中指定inputs为这些参数,它们在计算图的正向路径上,反向传播时就能正常计算并累积梯度。

额外说明:temp_weight是普通张量而非Parameter,设备迁移时PyTorch不会自动同步它的设备或更新它的计算图依赖,这也是计算图断裂的关键因素。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 00:31:20