torch 深度学习 (2)

torch
ConvNet

前面我们完成了数据的下载和预处理,接下来就该搭建网络模型了,CNN网络的东西可以参考博主 zouxy09的系列文章Deep Learning (深度学习) 学习笔记整理系列之 (七)

  1. 加载包

  1. require 'torch' 

  2. require 'image' 

  3. require 'nn' 

  1. 函数运行参数的设置

  1. if not opt then 

  2. print "==> processing options" 

  3. cmd = torch.CmdLine() 

  4. cmd:text() 

  5. cmd:text('options:') 

  6. -- 选择构建何种结构:线性|MLP|ConvNet。默认:convnet 

  7. cmd:option('-model','convnet','type of model to construct: linear | mlp | convnet') 

  8. -- 是否需要可视化 

  9. cmd:option('-visualize',true,'visualize input data and weights during training') 

  10. -- 参数 

  11. opt = cmd:parse(arg or {}) 

  12. end 

  1. 设置网络模型用到的一些参数

  1. -- 输出类别数,也就是输出节点个数 

  2. noutputs =10 

  3. -- 输入节点的个数 

  4. nfeats = 3 -- YUV三个通道,可以认为是3个features map 

  5. width =32 

  6. height =32 

  7. -- Linear 和 mlp model下的输入节点个数,就是将输入图像拉成列向量 

  8. ninputs = nfeats*width*height 


  9. -- 为mlp定义隐层节点的个数 

  10. nhiddens = ninputs/2  


  11. -- 为convnet定义隐层feature maps的个数以及滤波器的尺寸 

  12. nstates = {16,256,128} --第一个隐层有16个feature map,第二个隐层有256个特征图,第三个隐层有128个节点 

  13. fanin = {1,4} -- 定义了卷积层的输入和输出对应关系,以fanin[2]举例,表示该卷积层有16个map输入,256个map输出,每个输出map是有fanin[2]个输入map对应filters卷积得到的结果 

  14. filtsize =5 --滤波器的大小,方形滤波器 

  15. poolsize = 2 -- 池化池尺寸 

  16. normkernel = image.gaussian1D(7) --长度为7的一维高斯模板,用来local contrast normalization 

  1. 构建模型

  1. if opt.model == linear then  

  2. -- 线性模型 

  3. model = nn.Sequntial() 

  4. model:add(nn.Reshape(ninputs)) -- 输入层 

  5. model:add(nn.Linear(ninputs,noutputs)) -- 线性模型 y=Wx+b 

  6. elseif opt.model == mlp then  

  7. -- 多层感知器 

  8. model = nn.Sequential() 

  9. model:add(nn.Reshape(ninputs)) --输入层 

  10. model:add(nn.Linear(ninputs,nhiddens)) --线性层 

  11. model:add(nn.Tanh()) -- 非线性层 

  12. model:add(nn.Linear(nhiddens,noutputs)) -- 线性层 

  13. -- MLP 目标: `!$y=W_2 f(W_1X+b) + b $` 这里的激活函数采用的是Tanh(),MLP后面还可以接一层输出层Tanh() 

  14. elseif opt.model == convnet then 

  15. -- 卷积神经网络 

  16. model = nn.Sequential() 

  17. -- 第一阶段 

  18. model:add(nn.SpatialConvolutionMap(nn.tables.random(nfeats,nstates[1],fanin[1]),filtsize,filtsize)) 

  19. -- 这一步直接输入的是图像进行卷积,所以没有了 nn.Reshape(ninputs)输入层。 参数:nn.tables.random(nfeats,nstates[1],fanin[1])指定了卷积层中输入maps和输出maps之间的对应关系,这里表示bstates[1]个输出maps的每一map都是由fanin[1]个输入maps得到的。filtsize则是卷积算子的大小 

  20. -- 所以该层的连接个数为(filtsize*filtsize*fanin[1]+1)*nstates[1],1是偏置。这里的fanin[1]连接是随机的,也可以采用全连接 nn.tables.full(nfeats,nstates[1]), 当输入maps和输出maps个数相同时,还可以采用一对一连接 nn.tables.oneToOne(nfeats). 

  21. -- 参见解释文档 [Convolutional layers](https://github.com/torch/nn/blob/master/doc/concolution.md#nn.convlayers.dok) 


  22. model:add(nn.Tanh()) --非线性变换层 

  23. model:SpatialLPPooling(nstates[1],2,poolsize,poolsize,poolsize,poolsize) 

  24. -- 参数(feature maps个数,Lp范数,池化尺寸大小(w,h), 滑动窗步长(dw,dh)) 

  25. model:SpatialSubtractiveNormalization(nstates[1],normalkernel) 

  26. -- local contrast normalization 

  27. -- 具体操作是先在每个map的local邻域进行减法归一化,然后在不同的feature map上进行除法归一化。类似与图像点的均值化和方差归一化。参考[1^x][Nonlinear Image Representation Using Divisive Normalization], [Gaussian Scale Mixtures](stats.stackexchange.com/174502/what-are gaussian-scale-mixtures-and-how-to-generate-samples-of-gaussian-scale),还有解释文档 [Convolutional layers](https://github.com/torch/nn/blob/master/doc/concolution.md#nn.convlayers.dok) 


  28. --[[ 

  29. 这里需要说的一点是传统的CNN一般是先卷积再池化再非线性变换,但这两年CNN一般都是先非线性变换再池化了 

  30. --]] 

  31. -- 第二阶段 

  32. model:add(nn.SpatialConvolutionMap(nn.tables.random(nstates[1],nstates[2],fanin[2]),filtsize,filtsize)) 

  33. model:add(nn.Tanh()) 

  34. model:add(nn.SpatialLPPooling(nstates[2],2,poolsize,poolsize)) 

  35. model:add(nn.SpatialSubtractiveNormalization(nstates[2],kernel)) 


  36. --第三阶段 

  37. model:add(nn.Reshape(nstates[2]*filtsize*filtsize)) --矢量化,全连接 

  38. model:add(nn.Linear(nstates[2]*filtsize*filtsize,nstates[3])) 

  39. model:add(nn.Tanh()) 

  40. model:add(nn.Linear(nstates[3],noutputs)) 

  41. else 

  42. error('unknown -model') 

  43. end 

  1. 显示网络结构以及参数

  1. print('==> here is the model') 

  2. print(model) 

结果如下图

model.png

可以发现,可训练参数分别在1,5部分,所以可以观察权重矩阵的大小

  1. print('==> 权重矩阵的大小 ') 

  2. print(model:get(1).weight:size()) 

  3. print('==> 偏置的大小') 

  4. print(model:get(1).bias:numel()) 

weights numel.png
  1. 参数的可视化

  1. if opt.visualize then 

  2. image.display(image=model:get(1).weight, padding=2,zoom=4,legend='filters@ layer 1') 

  3. image.diaplay(image=model:get(5).weight,padding=2,zoom=4,legend='filters @ layer 2') 

  4. end 

weights visualization.png

torch 深度学习 (2)的更多相关文章

  1. torch 深度学习(5)

    torch 深度学习(5) mnist torch siamese deep-learning 这篇文章主要是想使用torch学习并理解如何构建siamese network. siamese net ...

  2. torch 深度学习(4)

    torch 深度学习(4) test doall files 经过数据的预处理.模型创建.损失函数定义以及模型的训练,现在可以使用训练好的模型对测试集进行测试了.测试模块比训练模块简单的多,只需调用模 ...

  3. torch 深度学习(3)

    torch 深度学习(3) 损失函数,模型训练 前面我们已经完成对数据的预处理和模型的构建,那么接下来为了训练模型应该定义模型的损失函数,然后使用BP算法对模型参数进行调整 损失函数 Criterio ...

  4. 深度学习菜鸟的信仰地︱Supervessel超能云服务器、深度学习环境全配置

    并非广告~实在是太良心了,所以费时间给他们点赞一下~ SuperVessel云平台是IBM中国研究院和中国系统与技术中心基于POWER架构和OpenStack技术共同构建的, 支持开发者远程开发的免费 ...

  5. 深度学习框架caffe/CNTK/Tensorflow/Theano/Torch的对比

    在单GPU下,所有这些工具集都调用cuDNN,因此只要外层的计算或者内存分配差异不大其性能表现都差不多. Caffe: 1)主流工业级深度学习工具,具有出色的卷积神经网络实现.在计算机视觉领域Caff ...

  6. 小白学习之pytorch框架(2)-动手学深度学习(begin-random.shuffle()、torch.index_select()、nn.Module、nn.Sequential())

    在这向大家推荐一本书-花书-动手学深度学习pytorch版,原书用的深度学习框架是MXNet,这个框架经过Gluon重新再封装,使用风格非常接近pytorch,但是由于pytorch越来越火,个人又比 ...

  7. [深度学习] Pytorch学习(一)—— torch tensor

    [深度学习] Pytorch学习(一)-- torch tensor 学习笔记 . 记录 分享 . 学习的代码环境:python3.6 torch1.3 vscode+jupyter扩展 #%% im ...

  8. 【深度学习Deep Learning】资料大全

    最近在学深度学习相关的东西,在网上搜集到了一些不错的资料,现在汇总一下: Free Online Books  by Yoshua Bengio, Ian Goodfellow and Aaron C ...

  9. [深度学习大讲堂]从NNVM看2016年深度学习框架发展趋势

    本文为微信公众号[深度学习大讲堂]特约稿,转载请注明出处 虚拟框架杀入 从发现问题到解决问题 半年前的这时候,暑假,我在SIAT MMLAB实习. 看着同事一会儿跑Torch,一会儿跑MXNet,一会 ...

随机推荐

  1. 编辑器——vscode

    1.编辑器个人工作配置 // 将设置放入此文件中以覆盖默认设置 { "editor.tabSize": 2, "workbench.iconTheme": &q ...

  2. 字典的fromkeys的用法

    fromkeys方法语法 dict.fromkeys(iterable[,value=None]) iterable 用于创建新的字典的键的可迭代对象(字符串,列表,元组,字典) value 可选参数 ...

  3. Linux系统——本地yum仓库安装

    一.yum仓库概述 yum是基于rpm包管理,能够从指定的服务器自动下载rpm包并且安装,可以自动处理依赖性关系,并且一次安装所有依赖的软件包,无需繁琐地一次次下载.安装. 二.yum仓库安装的方式 ...

  4. 101. Symmetric Tree(判断二叉树是否对称)

      Given a binary tree, check whether it is a mirror of itself (ie, symmetric around its center). For ...

  5. RPC细节

    服务化有什么好处? 服务化的一个好处就是,不限定服务的提供方使用什么技术选型,能够实现大公司跨团队的技术解耦,如下图所示: 服务A:欧洲团队维护,技术背景是Java 服务B:美洲团队维护,用C++实现 ...

  6. Ubuntu安装dlib后import出现libstdc++.so.6: version `GLIBCXX_3.4.21' not found

    1 问题描述 先安装依赖包cmake,libboost,再安装dlib sudo apt-get install cmake sudo apt-get install libboost-python- ...

  7. 多路选择I/O

    多路选择I/O提供另一种处理I/O的方法,相比于传统的I/O方法,这种方法更好,更具有效率.多路选择是一种充分利用系统时间的典型. 1.多路选择I/O的概念 当用户需要从网络设备上读数据时,会发生的读 ...

  8. 关于HttpRuntime.Cache的运用

    存Cache方法: HttpRuntime.Cache.Add( KeyName,//缓存名 KeyValue,//要缓存的对象 Dependencies,//依赖项 AbsoluteExpirati ...

  9. 20145328 《Java程序设计》第0周学习总结

    20145328 <Java程序设计>第0周学习总结 阅读心得 从总体上来说,这几篇文章都是围绕着软件工程专业的一些现象来进行描述的,但深入了解之后就可以发现,无论是软件工程专业还是我们现 ...

  10. SublimeText3 编辑器使用小结

    1. 快捷键: Command + shift + D : 复制当前行 Command + shift + K : 删除当前行 Command + J : 合并一行 Command + Enter : ...