深入理解Pytorch微调torchvision模型


Posted in Python onNovember 11, 2021

一、简介

在本小节,深入探讨如何对torchvision进行微调和特征提取。所有模型都已经预先在1000类的magenet数据集上训练完成。 本节将深入介绍如何使用几个现代的CNN架构,并将直观展示如何微调任意的PyTorch模型。
本节将执行两种类型的迁移学习:

  • 微调:从预训练模型开始,更新我们新任务的所有模型参数,实质上是重新训练整个模型。
  • 特征提取:从预训练模型开始,仅更新从中导出预测的最终图层权重。它被称为特征提取,因为我们使用预训练的CNN作为固定 的特征提取器,并且仅改变输出层。

通常这两种迁移学习方法都会遵循一下步骤:

  • 初始化预训练模型
  • 重组最后一层,使其具有与新数据集类别数相同的输出数
  • 为优化算法定义想要的训练期间更新的参数
  • 运行训练步骤

二、导入相关包

from __future__ import print_function
from __future__ import division
import torch
import torch.nn as nn
import torch.optim as optim
import numpy as np
import torchvision 
from torchvision import datasets,models,transforms
import matplotlib.pyplot as plt
import time
import os
import copy
print("Pytorch version:",torch.__version__)
print("torchvision version:",torchvision.__version__)

运行结果

深入理解Pytorch微调torchvision模型

三、数据输入

数据集——>我在这里

链接:https://pan.baidu.com/s/1G3yRfKTQf9sIq1iCSoymWQ
提取码:1234

#%%输入
data_dir="D:\Python\Pytorch\data\hymenoptera_data"
# 从[resnet,alexnet,vgg,squeezenet,desenet,inception]
model_name='squeezenet'
# 数据集中类别数量
num_classes=2
# 训练的批量大小
batch_size=8
# 训练epoch数
num_epochs=15
# 用于特征提取的标志。为FALSE,微调整个模型,为TRUE只更新图层参数
feature_extract=True

四、辅助函数

1、模型训练和验证

  • train_model函数处理给定模型的训练和验证。作为输入,它需要PyTorch模型、数据加载器字典、损失函数、优化器、用于训练和验 证epoch数,以及当模型是初始模型时的布尔标志。
  • is_inception标志用于容纳 Inception v3 模型,因为该体系结构使用辅助输出, 并且整体模型损失涉及辅助输出和最终输出,如此处所述。 这个函数训练指定数量的epoch,并且在每个epoch之后运行完整的验证步骤。它还跟踪最佳性能的模型(从验证准确率方面),并在训练 结束时返回性能最好的模型。在每个epoch之后,打印训练和验证正确率。
#%%模型训练和验证
device=torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
def train_model(model,dataloaders,criterion,optimizer,num_epochs=25,is_inception=False):
    since=time.time()
    val_acc_history=[]
    best_model_wts=copy.deepcopy(model.state_dict())
    best_acc=0.0
    for epoch in range(num_epochs):
        print('Epoch{}/{}'.format(epoch, num_epochs-1))
        print('-'*10)
        # 每个epoch都有一个训练和验证阶段
        for phase in['train','val']:
            if phase=='train':
                model.train()
            else:
                model.eval()
                
            running_loss=0.0
            running_corrects=0
            # 迭代数据
            for inputs,labels in dataloaders[phase]:
                inputs=inputs.to(device)
                labels=labels.to(device)
                # 梯度置零
                optimizer.zero_grad()
                # 向前传播
                with torch.set_grad_enabled(phase=='train'):
                    # 获取模型输出并计算损失,开始的特殊情况在训练中他有一个辅助输出
                    # 在训练模式下,通过将最终输出和辅助输出相加来计算损耗,在测试中值考虑最终输出
                    if is_inception and phase=='train':
                        outputs,aux_outputs=model(inputs)
                        loss1=criterion(outputs,labels)
                        loss2=criterion(aux_outputs,labels)
                        loss=loss1+0.4*loss2
                    else:
                        outputs=model(inputs)
                        loss=criterion(outputs,labels)
                        
                    _,preds=torch.max(outputs,1)
                    
                    if phase=='train':
                        loss.backward()
                        optimizer.step()
                        
                # 添加
                running_loss+=loss.item()*inputs.size(0)
                running_corrects+=torch.sum(preds==labels.data)
                
            epoch_loss=running_loss/len(dataloaders[phase].dataset)
            epoch_acc=running_corrects.double()/len(dataloaders[phase].dataset)
            
            print('{}loss : {:.4f} acc:{:.4f}'.format(phase, epoch_loss,epoch_acc))
            
            if phase=='train' and epoch_acc>best_acc:
                best_acc=epoch_acc
                best_model_wts=copy.deepcopy(model.state_dict())
            if phase=='val':
                val_acc_history.append(epoch_acc)
            
        print()

    time_elapsed=time.time()-since
    print('training complete in {:.0f}s'.format(time_elapsed//60, time_elapsed%60))
    print('best val acc:{:.4f}'.format(best_acc))
    
    model.load_state_dict(best_model_wts)
    return model,val_acc_history

2、设置模型参数的'.requires_grad属性'

当我们进行特征提取时,此辅助函数将模型中参数的 .requires_grad 属性设置为False。
默认情况下,当我们加载一个预训练模型时,所有参数都是 .requires_grad = True,如果我们从头开始训练或微调,这种设置就没问题。
但是,如果我们要运行特征提取并且只想为新初始化的层计算梯度,那么我们希望所有其他参数不需要梯度变化。

#%%设置模型参数的.require——grad属性
def set_parameter_requires_grad(model,feature_extracting):
    if feature_extracting:
        for param in model.parameters():
            param.require_grad=False

靓仔今天先去跑步了,再不跑来不及了,先更这么多,后续明天继续~(感谢有人没有催更!感谢监督!希望继续监督!)

以上就是深入理解Pytorch微调torchvision模型的详细内容,更多关于Pytorch torchvision模型的资料请关注三水点靠木其它相关文章!

Python 相关文章推荐
python比较两个列表大小的方法
Jul 11 Python
用python记录运行pid,并在需要时kill掉它们的实例
Jan 16 Python
python之PyMongo使用总结
May 26 Python
详解Python 实现元胞自动机中的生命游戏(Game of life)
Jan 27 Python
python 3.6.5 安装配置方法图文教程
Sep 18 Python
详解python列表生成式和列表生成式器区别
Mar 27 Python
Python GUI编程完整示例
Apr 04 Python
关于Python形参打包与解包小技巧分享
Aug 24 Python
python  ceiling divide 除法向上取整(或小数向上取整)的实例
Dec 27 Python
Python读取JSON数据操作实例解析
May 18 Python
python使用opencv resize图像不进行插值的操作
Jul 05 Python
python“静态”变量、实例变量与本地变量的声明示例
Nov 13 Python
Python 中 Shutil 模块详情
Nov 11 #Python
django 认证类配置实现
Nov 11 #Python
Python Pandas数据分析之iloc和loc的用法详解
据Python爬虫不靠谱预测可知今年双十一销售额将超过6000亿元
Python 详解通过Scrapy框架实现爬取百度新冠疫情数据流程
python中tkinter复选框使用操作
Nov 11 #Python
Python中的变量与常量
Nov 11 #Python
You might like
PHP面向对象分析设计的经验原则
2008/09/20 PHP
PHP下打开URL地址的几种方法小结
2010/05/16 PHP
JS中encodeURIComponent函数用php解码的代码
2012/03/01 PHP
php的XML文件解释类应用实例
2014/09/22 PHP
PHP中set_include_path()函数相关用法分析
2016/07/18 PHP
使用Post提交时须将空格转换成加号的解释
2013/01/14 Javascript
浅析JavaScript中的delete运算符
2013/11/30 Javascript
jQuery使用之处理页面元素用法实例
2015/01/19 Javascript
基于JavaScript制作霓虹灯文字 代码 特效
2015/09/01 Javascript
Backbone.js框架中简单的View视图编写学习笔记
2016/02/14 Javascript
jQuery 获取遍历获取table中每一个tr中的第一个td的方法
2016/10/05 Javascript
走进javascript——不起眼的基础,值和分号
2017/02/24 Javascript
关于JavaScript中高阶函数的魅力详解
2018/09/07 Javascript
vue-cli 首屏加载优化问题
2018/11/06 Javascript
详解nodejs 开发企业微信第三方应用入门教程
2019/03/12 NodeJs
详解小程序毫秒级倒计时(适用于拼团秒杀功能)
2019/05/05 Javascript
vue拖拽组件 vuedraggable API options实现盒子之间相互拖拽排序
2019/07/08 Javascript
vue项目中实现缓存的最佳方案详解
2019/07/11 Javascript
在Vue环境下利用worker运行interval计时器的步骤
2019/08/01 Javascript
js实现点击按钮随机生成背景颜色
2020/09/05 Javascript
使用python脚本自动创建pip.ini配置文件代码实例
2019/09/20 Python
Python编程快速上手——Excel到CSV的转换程序案例分析
2020/02/28 Python
Python with语句用法原理详解
2020/07/03 Python
世界上最大的餐具公司:Oneida
2016/12/17 全球购物
预订全球最佳旅行体验:Viator
2018/03/30 全球购物
诺心蛋糕官网:LE CAKE
2018/08/25 全球购物
SIMON MILLER官网:洛杉矶的生活方式品牌
2020/10/19 全球购物
养殖项目策划书范文
2014/01/13 职场文书
《油菜花开了》教学反思
2014/02/22 职场文书
年会搞笑主持词
2014/03/27 职场文书
大学生党员自我批评思想汇报
2014/10/10 职场文书
2014年幼儿园教研工作总结
2014/12/04 职场文书
2014年教师业务工作总结
2014/12/19 职场文书
Nginx tp3.2.3 404问题解决方案
2021/03/31 Servers
python-opencv 中值滤波{cv2.medianBlur(src, ksize)}的用法
2021/06/05 Python
MySQL数据库如何查看表占用空间大小
2022/06/10 MySQL