pytorch 在网络中添加可训练参数,修改预训练权重文件的方法


Posted in Python onAugust 17, 2019

实践中,针对不同的任务需求,我们经常会在现成的网络结构上做一定的修改来实现特定的目的。

假如我们现在有一个简单的两层感知机网络:

# -*- coding: utf-8 -*-
import torch
from torch.autograd import Variable
import torch.optim as optim
 
x = Variable(torch.FloatTensor([1, 2, 3])).cuda()
y = Variable(torch.FloatTensor([4, 5])).cuda()
 
class MLP(torch.nn.Module):
  def __init__(self):
    super(MLP, self).__init__()
    self.linear1 = torch.nn.Linear(3, 5)
    self.relu = torch.nn.ReLU()
    self.linear2 = torch.nn.Linear(5, 2)
 
  def forward(self, x):
    x = self.linear1(x)
    x = self.relu(x)
    x = self.linear2(x)
 
    return x
 
model = MLP().cuda()
 
loss_fn = torch.nn.MSELoss(size_average=False)
optimizer = optim.SGD(model.parameters(), lr=0.001, momentum=0.9)
 
for t in range(500):
  y_pred = model(x)
  loss = loss_fn(y_pred, y)
  print(t, loss.data[0])
  model.zero_grad()
  loss.backward()
  optimizer.step()
 
print(model(x))

现在想在前向传播时,在relu之后给x乘以一个可训练的系数,只需要在__init__函数中添加一个nn.Parameter类型变量,并在forward函数中乘以该变量即可:

class MLP(torch.nn.Module):
  def __init__(self):
    super(MLP, self).__init__()
    self.linear1 = torch.nn.Linear(3, 5)
    self.relu = torch.nn.ReLU()
    self.linear2 = torch.nn.Linear(5, 2)
    # the para to be added and updated in train phase, note that NO cuda() at last
    self.coefficient = torch.nn.Parameter(torch.Tensor([1.55]))
 
  def forward(self, x):
    x = self.linear1(x)
    x = self.relu(x)
    x = self.coefficient * x
    x = self.linear2(x)
 
    return x

注意,Parameter变量和Variable变量的操作大致相同,但是不能手动调用.cuda()方法将其加载在GPU上,事实上它会自动在GPU上加载,可以通过model.state_dict()或者model.named_parameters()函数查看现在的全部可训练参数(包括通过继承得到的父类中的参数):

print(model.state_dict().keys())
for i, j in model.named_parameters():
  print(i)
  print(j)

输出如下:

odict_keys(['linear1.weight', 'linear1.bias', 'linear2.weight', 'linear2.bias'])
linear1.weight
Parameter containing:
-0.3582 -0.0283 0.2607
 0.5190 -0.2221 0.0665
-0.2586 -0.3311 0.1927
-0.2765 0.5590 -0.2598
 0.4679 -0.2923 -0.3379
[torch.cuda.FloatTensor of size 5x3 (GPU 0)]
 
linear1.bias
Parameter containing:
-0.2549
-0.5246
-0.1109
 0.5237
-0.1362
[torch.cuda.FloatTensor of size 5 (GPU 0)]
 
linear2.weight
Parameter containing:
-0.0286 -0.3045 0.1928 -0.2323 0.2966
 0.2601 0.1441 -0.2159 0.2484 0.0544
[torch.cuda.FloatTensor of size 2x5 (GPU 0)]
 
linear2.bias
Parameter containing:
-0.4038
 0.3129
[torch.cuda.FloatTensor of size 2 (GPU 0)]

这个参数会在反向传播时与原有变量同时参与更新,这就达到了添加可训练参数的目的。

如果我们有原先网络的预训练权重,现在添加了一个新的参数,原有的权重文件自然就不能加载了,我们需要修改原权重文件,在其中添加我们的新变量的初始值。

调用model.state_dict查看我们添加的参数在参数字典中的完整名称,然后打开原先的权重文件:

a = torch.load("OldWeights.pth") a是一个collecitons.OrderedDict类型变量,也就是一个有序字典,直接将新参数名称和初始值作为键值对插入,然后保存即可。

a = torch.load("OldWeights.pth")
 
a["layer1.0.coefficient"] = torch.FloatTensor([1.2])
a["layer1.1.coefficient"] = torch.FloatTensor([1.5])
 
torch.save(a, "Weights.pth")

现在权重就可以加载在修改后的模型上了。

以上这篇pytorch 在网络中添加可训练参数,修改预训练权重文件的方法就是小编分享给大家的全部内容了,希望能给大家一个参考,也希望大家多多支持三水点靠木。

Python 相关文章推荐
Pyramid Mako模板引入helper对象的步骤方法
Nov 27 Python
Python中利用函数装饰器实现备忘功能
Mar 30 Python
利用python3随机生成中文字符的实现方法
Nov 24 Python
python的re正则表达式实例代码
Jan 24 Python
Odoo中如何生成唯一不重复的序列号详解
Feb 10 Python
详解python--模拟轮盘抽奖游戏
Apr 12 Python
关于不懂Chromedriver如何配置环境变量问题解决方法
Jun 12 Python
Python 多个图同时在不同窗口显示的实现方法
Jul 07 Python
python基于pdfminer库提取pdf文字代码实例
Aug 15 Python
django 做 migrate 时 表已存在的处理方法
Aug 31 Python
解决运行django程序出错问题 'str'object has no attribute'_meta'
Jul 15 Python
Django多个app urls配置代码实例
Nov 26 Python
python PyQt5/Pyside2 按钮右击菜单实例代码
Aug 17 #Python
Pytorch 实现自定义参数层的例子
Aug 17 #Python
Python中PyQt5/PySide2的按钮控件使用实例
Aug 17 #Python
画pytorch模型图,以及参数计算的方法
Aug 17 #Python
pytorch 共享参数的示例
Aug 17 #Python
Pytorch卷积层手动初始化权值的实例
Aug 17 #Python
pytorch自定义初始化权重的方法
Aug 17 #Python
You might like
YII2框架中使用yii.js实现的post请求
2017/04/09 PHP
Laravel 不同生产环境服务器的判断实践
2019/10/15 PHP
JavaScript入门之基本函数详解
2011/10/21 Javascript
javascript的创建多行字符串的7种方法
2014/04/29 Javascript
使用node.js 制作网站前台后台
2014/11/13 Javascript
jQuery实现页面滚动时动态加载内容的方法
2015/03/20 Javascript
Bootstrap每天必学之按钮(一)
2015/11/24 Javascript
AngularJS 实现JavaScript 动画效果详解
2016/09/08 Javascript
jQuery插件autocomplete使用详解
2017/02/04 Javascript
JS实现多张图片预览同步上传功能
2017/06/23 Javascript
jquery.validate.js 多个相同name的处理方式
2017/07/10 jQuery
利用yarn代替npm管理前端项目模块依赖的方法详解
2017/09/04 Javascript
Vue多种方法实现表头和首列固定的示例代码
2018/02/02 Javascript
js实现图片3D轮播效果
2019/09/21 Javascript
小程序自动化测试的示例代码
2020/08/11 Javascript
[32:30]夜魇凡尔赛茶话会 第一期01:谁是卧底
2021/03/11 DOTA
Python调用C语言开发的共享库方法实例
2015/03/18 Python
使用python装饰器计算函数运行时间的实例
2018/04/21 Python
Python装饰器限制函数运行时间超时则退出执行
2019/04/09 Python
Python对接 xray 和微信实现自动告警
2019/09/17 Python
Windows下实现将Pascal VOC转化为TFRecords
2020/02/17 Python
python基于socket函数实现端口扫描
2020/05/28 Python
英国第二大营养品供应商:Vitabiotics
2016/10/01 全球购物
高性能装备提升营地:Kammok
2019/02/27 全球购物
巴西购物网站:Submarino
2020/01/19 全球购物
网络事业创业计划书范文
2014/01/09 职场文书
学习礼仪心得体会
2014/09/01 职场文书
园艺专业毕业生求职信
2014/09/02 职场文书
公安机关纪律作风整顿个人剖析材料材料
2014/10/10 职场文书
群众路线个人整改措施
2014/10/24 职场文书
幼儿园见习报告范文
2014/10/30 职场文书
平安家庭事迹材料
2014/12/20 职场文书
2015入党自传格式范文
2015/06/26 职场文书
创业计划之特色精品店
2019/08/12 职场文书
启迪人心的励志语录:脾气永远不要大于本事
2020/01/02 职场文书
MySQL 数据类型选择原则
2021/05/27 MySQL