Pytorch原生AMP支持使用方法(1.6版本)
AMP:Automatic mixed precision,自动混合精度,可以在神经网络推理过程中,针对不同的层,采用不同的数据精度进行计算,从而实现节省显存和加快速度的目的。
在Pytorch 1.5版本及以前,通过NVIDIA出品的插件apex,可以实现amp功能。
从Pytorch 1.6版本以后,Pytorch将amp的功能吸收入官方库,位于torch.cuda.amp模块下。
本文为针对官方文档主要内容的简要翻译和自己的理解。
1. Introduction
torch.cuda.amp提供了对混合精度的支持。为实现自动混合精度训练,需要结合使用如下两个模块:
torch.cuda.amp.autocast:autocast主要用作上下文管理器或者装饰器,来确定使用混合精度的范围。torch.cuda.amp.GradScalar:GradScalar主要用来完成梯度缩放。
2. Typical Mixed Precision Training
一个典型的amp应用示例如下:
# 定义模型和优化器
model = Net().cuda()
optimizer = optim.SGD(model.parameters(), ...)
# 在训练最开始定义GradScalar的实例
scaler = GradScaler()
for epoch in epochs:
for input, target in data:
optimizer.zero_grad()
# 利用with语句,在autocast实例的上下文范围内,进行模型的前向推理和loss计算
with autocast():
output = model(input)
loss = loss_fn(output, target)
# 对loss进行缩放,针对缩放后的loss进行反向传播
# (此部分计算在autocast()作用范围以外)
scaler.scale(loss).backward()
# 将梯度值缩放回原尺度后,优化器进行一步优化
scaler.step(optimizer)
# 更新scalar的缩放信息
scaler.update()
3. Working with Unscaled Gradients
待更新
4. Working with Scaled Gradients
待更新
5. Working with Multiple Models, Losses, and Optimizers
如果模型的Loss计算部分输出多个loss,需要对每一个loss值执行scaler.scale。
如果网络具有多个优化器,对任一个优化器执行scaler.unscale_,并对每一个优化器执行scaler.step。
而scaler.update只在最后执行一次。
应用示例如下:
scaler = torch.cuda.amp.GradScaler()
for epoch in epochs:
for input, target in data:
optimizer0.zero_grad()
optimizer1.zero_grad()
with autocast():
output0 = model0(input)
output1 = model1(input)
loss0 = loss_fn(2 * output0 + 3 * output1, target)
loss1 = loss_fn(3 * output0 - 5 * output1, target)
scaler.scale(loss0).backward(retain_graph=True)
scaler.scale(loss1).backward()
# 选择其中一个优化器执行显式的unscaling
scaler.unscale_(optimizer0)
# 对每一个优化器执行scaler.step
scaler.step(optimizer0)
scaler.step(optimizer1)
# 完成所有梯度更新后,执行一次scaler.update
scaler.update()
6. Working with Multiple GPUs
针对多卡训练的情况,只影响autocast的使用方法,GradScaler的用法与之前一致。
6.1 DataParallel in a single process
在每一个不同的cuda设备上,torch.nn.DataParallel在不同的进程中执行前向推理,而autocast只在当前进程中生效,因此,如下方式的调用是不生效的:
model = MyModel()
dp_model = nn.DataParallel(model)
# 在主进程中设置autocast
with autocast():
# dp_model的内部进程并不会对autocast生效
output = dp_model(input)
# loss的计算在主进程中执行,autocast可以生效,但由于前面执行推理时已经失效,因此整体上是不正确的
loss = loss_fn(output)
有效的调用方式如下所示:
# 方法1:在模型构建中,定义forwar函数时,采用装饰器方式
MyModel(nn.Module):
...
@autocast()
def forward(self, input):
...
# 方法2:在模型构建中,定义forwar函数时,采用上下文管理器方式
MyModel(nn.Module):
...
def forward(self, input):
with autocast():
...
# DataParallel的使用方式不变
model = MyModel().cuda()
dp_model = nn.DataParallel(model)
# 在模型执行推理时,由于前面模型定义时的修改,在各cuda设备上的子进程中autocast生效
# 在执行loss计算是,在主进程中,autocast生效
with autocast():
output = dp_model(input)
loss = loss_fn(output)
6.2 DistributedDataParallel, one GPU per process
torch.nn.parallel.DistributedDataParallel在官方文档中推荐每个GPU执行一个实例的方法,以达到最好的性能表现。
在这种模式下,DistributedDataParallel内部并不会再启动子进程,因此对于autocast和GradScaler的使用都没有影响,与典型示例保持一致。
6.3 DistributedDataParallel, multiple GPUs per process
与DataParallel 的使用相同,在模型构建时,对forward函数的定义方式进行修改,保证autocast在进程内部生效。
Pytorch原生AMP支持使用方法(1.6版本)的更多相关文章
- 原生JS添加节点方法与jQuery添加节点方法的比较及总结
一.首先构建一个简单布局,来供下边讲解使用 1.HTML部分代码: <div id="div1">div1</div> <div id="d ...
- 原生JavaScript支持6种方式获取元素
一.原生JavaScript支持6种方式获取元素 document.getElementById('id'); document.getElementsByName('name'); document ...
- thinkPHP框架中执行原生SQL语句的方法
这篇文章主要介绍了thinkPHP框架中执行原生SQL语句的方法,结合实例形式分析了thinkPHP中执行原生SQL语句的相关操作技巧,并简单分析了query与execute方法的使用区别,需要的朋友 ...
- 现有语言不支持XXX方法
史上最强大的IDE也会有bug的时候哈,今天遇到这个问题特别郁闷,百度了下,果然也有人遇到过这个问题 解决方法: 1.调用的时候参数和接口声明的参数不一致(检查修改) 2.继承接口中残留一个废弃的方法 ...
- 原生JS事件绑定方法以及jQuery绑定事件方法bind、live、on、delegate的区别
一.原生JS事件绑定方法: 1.通过HTML属性进行事件处理函数的绑定如: <a href="#" onclick="f()"> 2.通过JavaS ...
- 原生JS中apply()方法的一个值得注意的用法
今天在学习vue.js的render时,遇到需要重复构造多个同类型对象的问题,在这里发现原生JS中apply()方法的一个特殊的用法: var ary = Array.apply(null, { &q ...
- PHPnow开启PHP扩展里openssl支持的方法
PHPnow 是 Win32 下绿色的 Apache + PHP + MySQL 环境套件包.简易安装.快速搭建支持虚拟主机的 PHP 环境.更多介绍<PHP服务套件 PHPnow1.5.6&g ...
- 扩展原生js的一些方法
扩展原生js的Array类 Array.prototype.add = function(item){ this.push(item); } Array.prototype.addRange = fu ...
- 原生Js 两种方法实现页面关键字高亮显示
原生Js 两种方法实现页面关键字高亮显示 上网看了看别人写的,不是兼容问题就是代码繁琐,自己琢磨了一下用两种方法都可以实现,各有利弊. 方法一 依靠正则表达式修改 1.获取obj的html2.统一替换 ...
随机推荐
- luogu CF125E MST Company wqs二分 构造
LINK:CF125E MST Company 难点在于构造 前面说到了求最小值 可以二分出斜率k然后进行\(Kruskal\) 然后可以得到最小值.\(mx\)为值域. 得到最小值之后还有一个构造问 ...
- 有关WebSocket必须了解的知识
一.前言 最近之前时间正好在学习java知识,所以自个想找个小项目练练手,由于之前的ssm系统已经跑了也有大半年了,虽然稀烂,但是功能还是勉强做到了,所以这次准备重构ssm系统,改名为postCode ...
- async和await的使用总结 ~ 竟然一直用错了c#中的async和await的使用。。
对于c#中的async和await的使用,没想到我一直竟然都有一个错误.. ..还是总结太少,这里记录下. 这里以做早餐为例 流程如下: 倒一杯咖啡. 加热平底锅,然后煎两个鸡蛋. 煎三片培根. 烤两 ...
- 【NOIP2016】组合数问题 题解(组合数学+递推)
题目链接 题目大意:给定$n,m,k$,求满足$k|C_i^j$的$C_i^j$的个数.$(0\leq i\leq n,1\leq j\leq \min(i,m))$. --------------- ...
- Caffe CuDNN版本与环境不同导致make错误
1.将./include/caffe/util/cudnn.hpp 换成最新版的caffe里的cudnn的实现,即相应的cudnn.hpp. 2.将./include/caffe/layers里的,所 ...
- C++Primer学习日记
计划:4.27-4.30 完成IO库.顺序容器两章 4/28 ------------------------------------------------- 为什么要使用using namespa ...
- Java web Cookie详解(持久化+原理详解+共享问题+设置中文+发送多个Cookie)
Java web Cookie详解 啥是cookie? 查询有道词典得: web和饼干有啥关系? 这个谜底等等来为大家揭晓 会话技术 web中的会话技术类似于生活中两个人聊天,不过web中的会话指的是 ...
- 谈谈代码评审(code review)
什么是代码评审(code review)? 根据维基百科的定义,代码评审是一种通过若干人员检阅源代码方式来进行的软件质量保证活动.根据软件工程的经典理论,代码评审应该是收益很高的活动,因其产生在Cod ...
- Newbe.Claptrap 框架如何实现在多种框架之上运行?
Newbe.Claptrap 框架如何实现在多种框架之上运行?最近整理了一下项目的术语表.今天就谈谈什么是 Claptrap Box. 特别感谢 kotone 为本文提供的校对建议! Newbe.Cl ...
- 尾递归(java)
一般递归: 一个过程或函数在其定义或说明中有直接或间接调用自身的一种方法,它通常把一个大型复杂的问题层层转化为一个与原问题相似的规模较小的问题来求解,递归策略只需少量的程序就可描述出解题过程所需要的多 ...