【PyTorch】with torch.no_grad() 详解

发布时间:2026/7/23 18:23:58
【PyTorch】with torch.no_grad() 详解 在阅读以下内容前请务必先大致了解计算图机制特别是叶子节点pytorch——计算图与动态图机制with torch.no_grad()在 PyTorch官网 中的定义为Context-manager that disables gradient calculation.意思是with torch.no_grad()是一个用于禁用梯度的上下文管理器。禁用梯度计算对于推理是很有用的当我们确定不会调用Tensor.backward()时它将减少计算的内存消耗。需要注意它并不会直接修改张量自身的requires_grad属性本质上在其内部执行的所有运算不会构建计算图节点也不会跟踪依赖。所以即便输入张量requires_gradTrue运算生成的新张量无法继承梯度依赖默认requires_gradFalse。除此之外with torch.no_grad()还通常与原地操作in-place operation组合在一起。原地操作有明确定义当计算图追踪处于开启状态时对于 requires_gradTrue 的叶子张量leaf tensor不能使用 inplace operation因为原地操作会直接覆盖张量内存中的原始数值。若仍在计算图构建阶段叶子张量原始值是反向传播求导必需的数据一旦原值被覆盖梯度与张量数值之间的依赖关系断裂求导无法正常进行因此框架提前抛出异常。withtorch.no_grad():forparaminparams:param-lr*param.grad/batch_size param.grad.zero_()# 清空当前梯度我们针对以上操作进行探究以更好理解该情况下with torch.no_grad()的作用。不使用with torch.no_grad()直接进行原地操作forparaminparams:param-lr*param.grad/batch_size param.grad.zero_()# 清空当前梯度运行上面的代码会报错错误信息为RuntimeError: a leaf Variable that requires grad is being used in an in-place operation.意思是在原地操作中使用了需要梯度的叶子节点。如果打印验证可以发现无论是否包裹with torch.no_grad()param.requires_grad始终为True。这与前文的定义是吻合的在计算图追踪关闭时即使是叶子节点并具requires_grad属性也可以进行就地操作。与原地操作无关的补充两点客观现象lr * param.grad / batch_size运算生成临时张量该临时张量的requires_grad Falseparam.grad作为参数附属的梯度张量本身requires_grad恒为False该性质与是否启用with torch.no_grad()无关。不使用with torch.no_grad()进行赋值操作forparaminparams:paramparam-lr*param.grad/batch_sizeprint(param.is_leaf)# Falseparam.grad.zero_()# 清空当前梯度运行上面的代码会报错错误信息为AttributeError: NoneType object has no attribute zero_。我们都知道赋值操作会新创建一块内存以存放数据所以根据计算图理论此时的param是中间节点不再是叶子节点不具有grad属性了。使用with torch.no_grad()进行赋值操作withtorch.no_grad():forparaminparams:print(param.requires_grad)# Trueparamparam-lr*param.grad/batch_sizeprint(param.is_leaf)# Trueprint(param.requires_grad)# Falseparam.grad.zero_()# 清空当前梯度运行上面的代码同样会报错错误信息为AttributeError: NoneType object has no attribute zero_。在torch.no_grad()上下文内不再构建计算图。执行param param - lr * param.grad / batch_size依旧生成一块全新张量该张量在无梯度追踪环境下创建不存在运算依赖因此is_leafTrue上下文内运算不会继承梯度标记故requires_gradFalserequires_gradFalse的张量不会分配梯度缓冲区.grad恒为None因此调用.zero_()报错。