模型跨设备迁移后,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

