如何复用前序计算结果且不破坏autograd自动求导机制?
报错核心原因
第一次调用loss.backward()时,PyTorch默认会释放当前计算图的中间缓存。你后续时间步用到的first_output是i=0步计算得出的张量,还绑定了i=0步的计算图,第二步反向传播时会尝试再次遍历已经被释放的i=0步计算图,因此触发报错。
你之前的尝试未生效,是因为detach()调用的位置不对,没有在给first_output赋值的节点就完成张量的计算图剥离;而直接修改非叶子节点的requires_grad属性本身就不符合PyTorch的机制,因此会报错。
正确解决方案
只需要在i=0步给first_output赋值时,对output做detach操作即可。detach会返回一个和原张量共享数据、但完全脱离原计算图的新张量,不会被autograd追踪梯度,完全满足你将其作为固定输入的需求。
修改后的代码如下:
for i in range(time): optimizer.zero_grad() if i == 0: first_output = model.output_init() output = model(input, first_output) loss = F.loss(output) if i == 0: # 从i=0的计算图中剥离output,赋值给first_output first_output = output.detach() loss.backward() optimizer.step()
如果担心张量内存共享带来的意外修改,可以额外加clone操作,写法为first_output = output.detach().clone(),仅会增加少量显存占用,不影响功能逻辑。
注意:不要为了规避报错直接给
loss.backward()加retain_graph=True参数,这个方案会持续保留所有历史步的计算图,显存占用会随时间步线性增长,而且会导致后续反向传播错误地把梯度累积到i=0步的计算图中,不符合你把first_output作为固定输入的需求。
内容的提问来源于stack exchange,提问作者shazoop
相关产品推荐
相关产品推荐

