pytorch 使用半精度模型部署的操作


Posted in Python onMay 24, 2021

背景

pytorch作为深度学习的计算框架正得到越来越多的应用.

我们除了在模型训练阶段应用外,最近也把pytorch应用在了部署上.

在部署时,为了减少计算量,可以考虑使用16位浮点模型,而训练时涉及到梯度计算,需要使用32位浮点,这种精度的不一致经过测试,模型性能下降有限,可以接受.

但是推断时计算量可以降低一半,同等计算资源下,并发度可提升近一倍

具体方法

在pytorch中,一般模型定义都继承torch.nn.Moudle,torch.nn.Module基类的half()方法会把所有参数转为16位浮点,所以在模型加载后,调用一下该方法即可达到模型切换的目的.接下来只需要在推断时把input的tensor切换为16位浮点即可

另外还有一个小的trick,在推理过程中模型输出的tensor自然会成为16位浮点,如果需要新创建tensor,最好调用已有tensor的new_zeros,new_full等方法而不是torch.zeros和torch.full,前者可以自动继承已有tensor的类型,这样就不需要到处增加代码判断是使用16位还是32位了,只需要针对input tensor切换.

补充:pytorch 使用amp.autocast半精度加速训练

准备工作

pytorch 1.6+

如何使用autocast?

根据官方提供的方法,

答案就是autocast + GradScaler。

1,autocast

正如前文所说,需要使用torch.cuda.amp模块中的autocast 类。使用也是非常简单的:

如何在PyTorch中使用自动混合精度?

答案:autocast + GradScaler。

1.autocast

正如前文所说,需要使用torch.cuda.amp模块中的autocast 类。使用也是非常简单的

from torch.cuda.amp import autocast as autocast

# 创建model,默认是torch.FloatTensor
model = Net().cuda()
optimizer = optim.SGD(model.parameters(), ...)

for input, target in data:
    optimizer.zero_grad()

    # 前向过程(model + loss)开启 autocast
    with autocast():
        output = model(input)
        loss = loss_fn(output, target)

    # 反向传播在autocast上下文之外
    loss.backward()
    optimizer.step()

2.GradScaler

GradScaler就是梯度scaler模块,需要在训练最开始之前实例化一个GradScaler对象。

因此PyTorch中经典的AMP使用方式如下:

from torch.cuda.amp import autocast as autocast

# 创建model,默认是torch.FloatTensor
model = Net().cuda()
optimizer = optim.SGD(model.parameters(), ...)
# 在训练最开始之前实例化一个GradScaler对象
scaler = GradScaler()

for epoch in epochs:
    for input, target in data:
        optimizer.zero_grad()

        # 前向过程(model + loss)开启 autocast
        with autocast():
            output = model(input)
            loss = loss_fn(output, target)

        scaler.scale(loss).backward()
        scaler.step(optimizer)
        scaler.update()

3.nn.DataParallel

单卡训练的话上面的代码已经够了,亲测在2080ti上能减少至少1/3的显存,至于速度。。。

要是想多卡跑的话仅仅这样还不够,会发现在forward里面的每个结果都还是float32的,怎么办?

class Model(nn.Module):
    def __init__(self):
        super(Model, self).__init__()

    def forward(self, input_data_c1):
     with autocast():
      # code
     return

只要把forward里面的代码用autocast代码块方式运行就好啦!

自动进行autocast的操作

如下操作中tensor会被自动转化为半精度浮点型的torch.HalfTensor:

1、matmul

2、addbmm

3、addmm

4、addmv

5、addr

6、baddbmm

7、bmm

8、chain_matmul

9、conv1d

10、conv2d

11、conv3d

12、conv_transpose1d

13、conv_transpose2d

14、conv_transpose3d

15、linear

16、matmul

17、mm

18、mv

19、prelu

那么只有这些操作才能半精度吗?不是。其他操作比如rnn也可以进行半精度运行,但是需要自己手动,暂时没有提供自动的转换。

Python 相关文章推荐
Python切片用法实例教程
Sep 08 Python
Python实现115网盘自动下载的方法
Sep 30 Python
在Python中用split()方法分割字符串的使用介绍
May 20 Python
Django URL传递参数的方法总结
Aug 28 Python
Python星号*与**用法分析
Feb 02 Python
在matplotlib的图中设置中文标签的方法
Dec 13 Python
Python实现对特定列表进行从小到大排序操作示例
Feb 11 Python
Python实现将字符串的首字母变为大写,其余都变为小写的方法
Jun 11 Python
详解python实现交叉验证法与留出法
Jul 11 Python
Django MEDIA的配置及用法详解
Jul 25 Python
wxPython:python首选的GUI库实例分享
Oct 05 Python
解决keras backend 越跑越慢问题
Jun 18 Python
解决Pytorch半精度浮点型网络训练的问题
May 24 #Python
Python办公自动化之Excel(中)
May 24 #Python
PyTorch梯度裁剪避免训练loss nan的操作
May 24 #Python
python3读取文件指定行的三种方法
May 24 #Python
pytorch中Schedule与warmup_steps的用法说明
May 24 #Python
Python Pycharm虚拟下百度飞浆PaddleX安装报错问题及处理方法(亲测100%有效)
May 24 #Python
pytorch交叉熵损失函数的weight参数的使用
May 24 #Python
You might like
德生S2000电路分析
2021/03/02 无线电
JpGraph php柱状图使用介绍
2011/08/23 PHP
php数组中删除元素的实现代码
2012/06/22 PHP
php获取表单中多个同名input元素的值
2014/03/20 PHP
PHP自定义函数实现数组比较功能示例
2017/10/19 PHP
Laravel模型事件的实现原理详解
2018/03/14 PHP
ThinkPHP5与单元测试PHPUnit使用详解
2020/02/23 PHP
二级域名转向类
2006/11/09 Javascript
JavaScript对象反射用法实例
2015/04/17 Javascript
使用JQuery选择HTML遍历函数的方法
2016/09/17 Javascript
深入理解JS继承和原型链的问题
2016/12/17 Javascript
vue-router动态设置页面title的实例讲解
2018/08/30 Javascript
JavaScript常用工具方法封装
2019/02/12 Javascript
详解jQuery设置内容和属性
2019/04/11 jQuery
Vue中keep-alive组件的深入理解
2020/08/23 Javascript
windows下wxPython开发环境安装与配置方法
2014/06/28 Python
在Python的Flask框架下收发电子邮件的教程
2015/04/21 Python
python返回昨天日期的方法
2015/05/13 Python
Python用61行代码实现图片像素化的示例代码
2018/12/10 Python
Python Matplotlib库安装与基本作图示例
2019/01/09 Python
Python使用字典的嵌套功能详解
2019/02/27 Python
python将数据插入数据库的代码分享
2020/08/16 Python
Ubuntu权限不足无法创建文件夹解决方案
2020/11/14 Python
Python3 + Appium + 安卓模拟器实现APP自动化测试并生成测试报告
2021/01/27 Python
文员岗位职责范本
2014/03/08 职场文书
奉献演讲稿范文
2014/05/21 职场文书
阅兵口号
2014/06/19 职场文书
教室布置标语
2014/06/26 职场文书
2014年安置帮教工作总结
2014/12/11 职场文书
十佳少年事迹材料
2014/12/25 职场文书
2016年感恩节寄语
2015/12/07 职场文书
《蟋蟀的住宅》教学反思
2016/02/17 职场文书
为什么 Nginx 比 Apache 更牛逼
2021/03/31 Servers
Mysql8.0递归查询的简单用法示例
2021/08/04 MySQL
Java 异步任务计算FutureTask
2022/04/28 Java/Android
CSS 鼠标选中文字后改变背景色的实现代码
2023/05/21 HTML / CSS