当前位置: 首页 > news >正文

【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()的作用。

  1. 不使用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属性也可以进行就地操作。
与原地操作无关的,补充两点客观现象:

  1. lr * param.grad / batch_size运算生成临时张量,该临时张量的requires_grad = False
  2. param.grad作为参数附属的梯度张量,本身requires_grad恒为False,该性质与是否启用with torch.no_grad()无关。
  1. 不使用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属性了。

  1. 使用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依旧生成一块全新张量:

  1. 该张量在无梯度追踪环境下创建,不存在运算依赖,因此is_leaf=True
    上下文内运算不会继承梯度标记,故requires_grad=False
  2. requires_grad=False的张量不会分配梯度缓冲区,.grad恒为None,因此调用.zero_()报错。
http://www.cnnetsun.cn/news/3606718.html

相关文章:

  • 如何基于小型工控机与openwrt打造家用软路由
  • 麒麟信安登录央视, 深度展现为中国信息安全铸“魂”之路
  • 【模拟IC学习笔记】 PSS和Pnoise仿真
  • PairLIE论文阅读笔记
  • AI生成音乐版权雷区全扫描,律师+音频工程师双视角解析:93%创作者不知的4类侵权高发场景
  • Qt 定时器放在线程中执行,支持随时开始和停止定时器。
  • AI模型优化与多模态学习实战指南
  • 动态住宅代理使用指南:粘性会话 vs. 每次请求轮换 IP,如何选择?
  • 关于在VM虚拟机下,安装OpenWrt软路由,所遇错误及解决方法。
  • 如何在PVE(Proxmox)中安装OpenWrt软路由?
  • AI如何革新本科生论文写作:从选题到答辩的全流程优化
  • 随身wifi刷openwrt变软路由--103s没网的解决
  • OpenWrt 软路由介绍
  • CTFHub-WEB-文件上传
  • YOLOv5工业检测优化:CSP-EDLAN模型实战解析
  • Github项目分享——免费的编程中文书籍索引
  • 【单片机毕业设计推荐】 基于 STM32 或 51 单片机的智能感应自动门控制系统设计与实现,基于 STM32 或 51 单片机的带人数统计功能智能门禁装置设计(012403)
  • ARM DAP与JTAG指令深度解析:从调试原理到故障排查实战
  • AI技术如何革新论文查重系统
  • Tiva™ TM4C129LNCZAD PWM故障保护与中断控制全解析
  • 从技术层级角度看多链路聚合通信技术
  • TM4C I2C从机中断与FIFO配置实战:从寄存器原理到高效数据引擎
  • 从零构建C++轻量级系统监控方案:原理、实现与生产实践
  • AI招聘解决方案:重构招聘价值链的16项核心技术
  • 索引失效避坑: 明明是等值查询,为何EXPLAIN显示走了全表扫描?
  • 嵌入式外设识别与EPI接口:从GPIO身份验证到高速并行通信
  • 深入Tiva™ TM4C129时钟与电源管理:从寄存器配置到低功耗实战
  • DSP上AES算法性能优化实战:从C代码到线性汇编的深度调优
  • npm : 无法加载文件 C:\Program Files\nodejs\npm.ps1,因为在此系统上禁止运行脚本。
  • 图像滤波原理与OpenCV实践指南