如何在torch.no_grad中排除指定张量以保留梯度计算?
在torch.no_grad模式下为特定张量保留梯度计算的方法
要实现全局处于torch.no_grad模式以节省资源,仅为少数特定张量开启梯度计算,可以通过以下两种核心方式实现:
1. 局部使用torch.set_grad_enabled(True)覆盖全局设置
torch.no_grad()会全局关闭梯度计算,但你可以在需要保留梯度的代码块外层,用torch.set_grad_enabled(True)上下文管理器局部重新开启梯度计算,其优先级高于外层的no_grad。
修改后的代码示例
import torch x = torch.tensor([1.], requires_grad = True) @torch.no_grad() def func(x): # 对需要梯度的张量操作,局部开启梯度计算 with torch.set_grad_enabled(True): a = torch.tensor([1.], requires_grad=True) # 显式设置需要梯度 b = a * 2 b.backward() print(a.grad) # 输出 tensor([2.]),成功计算梯度 # 其余操作仍处于no_grad模式,不会记录梯度 return x ** 2 z = func(x) print(z.requires_grad) # 输出 False,符合全局no_grad设置
2. 手动控制张量的requires_grad与retain_grad
如果是已存在的张量,或者需要更精细的控制:
- 对需要梯度的张量,显式设置
requires_grad=True(即使在no_grad里,该张量的操作也会被追踪梯度,前提是局部开启了梯度计算)。 - 若需要对非叶子张量(比如中间变量)保留梯度,需调用
张量.retain_grad(),否则反向传播后中间张量的梯度会被自动释放。
示例(针对非叶子张量)
import torch x = torch.tensor([1.], requires_grad = True) @torch.no_grad() def func(x): with torch.set_grad_enabled(True): a = torch.tensor([1.], requires_grad=True) b = a * 2 b.retain_grad() # 保留非叶子张量b的梯度 c = b * 3 c.backward() print(a.grad) # 输出 tensor([6.]) print(b.grad) # 输出 tensor([3.]) return x ** 2 z = func(x)
关键注意事项
torch.set_grad_enabled(True)可以嵌套在no_grad、torch.inference_mode()等上下文环境中,局部覆盖全局的梯度开关设置。- 必须确保需要梯度的张量显式设置
requires_grad=True,否则即使开启了梯度计算,也不会追踪该张量的梯度。 - 全局
no_grad模式下,未被局部开启梯度计算的操作不会生成计算图,能大幅节省内存占用和计算时间。
内容的提问来源于stack exchange,提问作者JobHunter69
相关产品推荐
相关产品推荐

