为PyTorch模型添加标量参数触发RuntimeError问题求助
核心代码逻辑如下:
class WholeModel: def __init__(...): self.new_parameter = Parameter(torch.scalar_tensor(0.1, requires_grad=True)) self.model = self.make_model() def make_model(self): d = distribution() # returns a Distribution which is a Module d = transform_distribution(d, self.new_parameter) d.register_parameter(name='new', param=self.new_parameter) return d
运行时触发错误: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.
尝试过的无效操作:
- 将
self.new_parameter = Parameter(torch.scalar_tensor(0.1))改为self.new_parameter = torch.scalar_tensor(0.1)并移除register_parameter,代码可正常运行但参数无法被学习 - 用
tensor替代scalar_tensor,无论是否设置requires_grad都会触发错误
移除重复的参数注册
你同时将new_parameter注册为WholeModel的属性和子模块d的参数,这会导致参数在计算图中被多次关联,反向传播时触发重复回溯问题。只需在WholeModel中保留self.new_parameter = Parameter(torch.tensor(0.1, requires_grad=True))即可,无需再调用d.register_parameter——因为self.model = d已经让d成为WholeModel的子模块,参数会自动纳入模型的可学习参数列表。排查
transform_distribution的实现
错误根源大概率在该函数内部:- 检查是否存在对
new_parameter的原地修改操作,这类操作会破坏计算图完整性 - 若函数内多次使用该参数构建计算图,可尝试在非梯度依赖的环节用
detach()临时断开图关联,避免计算图冗余
- 检查是否存在对
验证计算图的重复使用情况
如果在第一次反向传播时添加retain_graph=True后错误消失,说明存在多次回溯同一计算图的情况。需要排查代码中是否有其他提前触发梯度计算的逻辑,比如子模块内部的隐式反向传播操作,找到重复触发点并修正。简化参数初始化
用Parameter(torch.tensor(0.1, requires_grad=True))替代scalar_tensor初始化,两者功能一致,可避免底层实现差异带来的潜在问题。
内容的提问来源于stack exchange,提问作者Lucidnonsense

