【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_grad=True,运算生成的新张量无法继承梯度依赖,默认requires_grad=False。
除此之外,with torch.no_grad()还通常与原地操作(in-place operation)组合在一起。原地操作有明确定义:
当计算图追踪处于开启状态时,对于 requires_grad=True 的叶子张量(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 = False;param.grad作为参数附属的梯度张量,本身requires_grad恒为False,该性质与是否启用with torch.no_grad()无关。
- 不使用
with torch.no_grad()进行赋值操作
forparaminparams:param=param-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)# Trueparam=param-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_leaf=True;
上下文内运算不会继承梯度标记,故requires_grad=False;requires_grad=False的张量不会分配梯度缓冲区,.grad恒为None,因此调用.zero_()报错。
