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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 13:32:46