torchnet+VGG16计算patch之间相似度

torch
VGG16
similarity

本来打算使用VGG实现siamese CNN的,但是没想明白怎么使用torchnet对模型进行微调。。。所以只好把VGG的卷积层单独做一个数据预处理模块,后面跟一个网络,将两个VGG输出的结果输入该网络中,仅训练这个浅层网络。

数据:使用了MOTChallenge数据库MOT16-02中的pedestrian

代码:

  1. -- --------------------------------------------------------------------------------------- 

  2. -- 读取MOT16-02数据集的groundtruth,分成训练集和测试集 

  3. -- --------------------------------------------------------------------------------------- 

  4. require 'torch' 

  5. require 'cutorch' 

  6. torch.setdefaulttensortype('torch.FloatTensor') 

  7. data_type = 'torch.CudaTensor' -- 设置数据类型,不适用GPU可以设置为torch.FloatTensor 


  8. require 'image' 

  9. local datapath = '/home/zwzhou/programFiles/2DMOT2015/MOT16/train/MOT16-02/' 

  10. local tmp = image.load(datapath .. 'img1/000001.jpg',3,'byte') 

  11. local width = tmp:size(3) 

  12. local height = tmp:size(2) 

  13. local num = 600 

  14. local imgs = torch.Tensor(num,3,height,width) 


  15. local file,_ = io.open('imgs.t7') 

  16. if not file then 

  17. for i=1,num do -- 读取视频帧 

  18. imgs[i]=image.load(datapath .. 'img1/' .. string.format('%06d.jpg',i)) 

  19. end 

  20. torch.save('imgs.t7',imgs) 

  21. else 

  22. imgs = torch.load('imgs.t7') 

  23. end 


  24. require'sys' 

  25. local gt_path = datapath .. 'gt/gt.txt' 

  26. local gt_info={} 

  27. local i=0 

  28. for line in io.lines(gt_path) do -- pedestrians的patch信息 

  29. local v=sys.split(line,',') 

  30. if tonumber(v[7]) ==1 and tonumber(v[9]) > 0.8 then -- 筛选有效的patch,是pedestrian且可见度>0.8 

  31. table.insert(gt_info,{tonumber(v[1]),tonumber(v[2]),tonumber(v[3]),tonumber(v[4]),tonumber(v[5]),tonumber(v[6])}) 

  32. -- 对应的是frame index,track index, x, y, w, h 

  33. end 

  34. end 

  35. -- 构建样本对,这里主要是为了正负样本个数相同,每个pedestrian选取25个相同id的patch,25个不同id的patch 

  36. local pairwise={} 

  37. for i=1,#gt_info do 

  38. local count=0 

  39. local iter=0 

  40. repeat  

  41. local j=torch.ceil(torch.rand(1)*(#gt_info))[1] 

  42. if gt_info[i][2] == gt_info[j][2] then  

  43. count=count+1 

  44. table.insert(pairwise,{i,j}) 

  45. end 

  46. iter=iter+1 

  47. until(count >25 or iter>100) 

  48. repeat  

  49. local j=torch.ceil(torch.rand(1)*#gt_info)[1] 

  50. if gt_info[i][2] ~= gt_info[j][2] then  

  51. count=count-1 

  52. table.insert(pairwise,{i,j}) 

  53. end 

  54. until(count <0) 

  55. end 


  56. local function cast(x) return x:type(data_type) end -- 类型转换 


  57. -- 加载pretrained VGG16 model 

  58. require 'nn' 

  59. require 'loadcaffe' 

  60. local function getPretrainedModel() 

  61. local proto = '/home/zwzhou/modelZoo/VGG_ILSVRC_16_layers_deploy.prototxt' 

  62. local caffemodel = '/home/zwzhou/modelZoo/VGG_ILSVRC_16_layers.caffemodel' 

  63. local VGG16 = loadcaffe.load(proto,caffemodel,'nn') 

  64. for i = 1,3 do 

  65. VGG16.modules[#VGG16.modules]=nil 

  66. end 

  67. return VGG16 

  68. end 


  69. -- 为了能够使用VGG,需要定义一些预处理方法 

  70. local loadSize = {3,256,256} 

  71. local sampleSize={3,224,224} 


  72. local function adjustScale(input) -- VGG需要先将输入图片的最小边缩放到256,另一边保持纵横比 

  73. if input:size(3) < input:size(2) then 

  74. input = image.scale(input,loadSize[2],loadSize[3]*input:size(2)/input:size(3)) 

  75. else 

  76. input = image.scale(input,loadSize[2]*input:size(3)/input:size(2),loadSize[3]) 

  77. end 

  78. return input 

  79. end 


  80. local bgr_means = {103.939,116.779,123.68} -- VGG使用的均值,注意是BGR通道,image.load()获得的是rgb 

  81. local function vggProcessing(img) 

  82. local img2 = img:clone() -- 深度拷贝 

  83. img2[{{1}}] = img[{{3}}] 

  84. img2[{{3}}] = img[{{1}}] -- rgb -> bgr 

  85. img2=img2:mul(255) 

  86. for i=1,3 do 

  87. img2[i]:add(-bgr_means[i]) 

  88. end 

  89. return img2 

  90. end 


  91. local function centerCrop(input) -- 截取224*224大小 

  92. local oH = sampleSize[2] 

  93. local oW = sampleSize[3] 

  94. local iW = input:size(3) 

  95. local iH = input:size(2) 

  96. local w1 = math.ceil((iW-oW)/2) 

  97. local h1 = math.ceil((iH-oH)/2) 

  98. local out = image.crop(input,w1,h1,w1+oW,h1+oH) 

  99. return out 

  100. end 


  101. local file,_ = io.open('vgg_info.t7') 

  102. local vgg_info={} 

  103. if not file then 

  104. local VGG16_model = getPretrainedModel() 

  105. if data_type:match'torch.Cuda.*Tensor' then 

  106. require 'cudnn' 

  107. require 'cunn' 

  108. cudnn.convert(VGG16_model,cudnn):cuda() 

  109. cudnn.benchmark = true 

  110. end 

  111. cast(VGG16_model) 

  112. for i=1, #gt_info do 

  113. local idx=gt_info[i] 

  114. local img = imgs[idx[1]] 

  115. local x1 = math.max(idx[3],1) 

  116. local y1 = math.max(idx[4],1) 

  117. local x2 = math.min(idx[3]+idx[5],width) 

  118. local y2 = math.min(idx[4]+idx[6],height) 

  119. local patch = image.crop(img,x1,y1,x2,y2) 

  120. patch = adjustScale(patch) 

  121. patch = vggProcessing(patch) 

  122. patch = centerCrop(patch) 

  123. patch=cast(patch) 

  124. table.insert(vgg_info,VGG16_model:forward(patch):float()) 

  125. end 

  126. torch.save('vgg_info.t7',vgg_info) 

  127. else  

  128. vgg_info=torch.load('vgg_info.t7') 

  129. end 


  130. local function getPatchPair(tmp) -- 获得patch 对 

  131. local pp = {} 

  132. pp[1] = vgg_info[tmp[1]] 

  133. pp[2] = vgg_info[tmp[2]] 

  134. local t=torch.cat(pp[1],pp[2],1) 

  135. return t 

  136. end 


  137. -- 定义datasetiterator 

  138. local tnt=require'torchnet' 

  139. local function getIterator(mode) 

  140. -- 创建model 

  141. local fc = nn.Sequential() 

  142. fc:add(nn.View(-1,4096*2)) 

  143. fc:add(nn.Linear(4096*2,500)) 

  144. fc:add(nn.ReLU(true)) 

  145. fc:add(nn.Normalize(2)) 

  146. fc:add(nn.Linear(500,500)) 

  147. fc:add(nn.ReLU(true)) 

  148. fc:add(nn.Linear(500,1)) 


  149. -- print(fc:forward(torch.randn(2,4096*2))) 

  150. if data_type:match'torch.Cuda.*Tensor' then 

  151. require 'cudnn' 

  152. require 'cunn' 

  153. cudnn.convert(fc,cudnn):cuda() 

  154. cudnn.benchmark = true 

  155. end 

  156. cast(fc) 


  157. -- 构建训练引擎,使用OptimEngine 

  158. require 'optim' 

  159. local engine = tnt.OptimEngine() 

  160. local criterion = cast(nn.MarginCriterion()) 


  161. -- 创建一些评估值 

  162. local train_timer = torch.Timer() 

  163. local test_timer = torch.Timer() 

  164. local data_timer = torch.Timer() 


  165. local meter = tnt.AverageValueMeter() -- 用于统计评估函数的输出 

  166. local confusion = optim.ConfusionMatrix(2) -- 2类混淆矩阵 

  167. local data_time_meter = tnt.AverageValueMeter() 

  168. -- log 

  169. local logtext=require 'torchnet.log.view.text' 

  170. log = tnt.Log{ 

  171. keys = {'train_loss','train_acc','data_loading_time','epoch','test_acc','train_time','test_time'}, 

  172. onFlush={ 

  173. logtext{keys={'train_loss','train_acc','data_loading_time','epoch','test_acc','train_time','test_time'}} 

  174. } 

  175. } 


  176. local inputs = cast(torch.Tensor()) 

  177. local targets = cast(torch.Tensor()) 


  178. -- 填一些hook函数,以便观察训练过程 

  179. engine.hooks.onSample = function(state) 

  180. if state.training then 

  181. data_time_meter:add(data_timer:time().real) 

  182. end 

  183. inputs:resize(state.sample.input:size()):copy(state.sample.input) 

  184. targets:resize(state.sample.target:size()):copy(state.sample.target) 

  185. state.sample.input = inputs 

  186. state.sample.target = targets 

  187. end 


  188. engine.hooks.onForwardCriterion = function(state) 

  189. meter:add(state.criterion.output) 

  190. confusion:batchAdd(state.network.output:gt(0):add(1),state.sample.target:gt(0):add(1)) 

  191. end 


  192. local function test() -- 用于测试 

  193. engine:test{ 

  194. network = fc, 

  195. iterator = getIterator('test'), 

  196. criterion=criterion,  

  197. } 

  198. confusion:updateValids() 

  199. end 


  200. engine.hooks.onStartEpoch = function(state) 

  201. local epoch = state.epoch + 1 

  202. print('===>' .. ' online epoch # ' .. epoch .. '[batchsize = 256]') 

  203. meter:reset() 

  204. confusion:zero() 

  205. train_timer:reset() 

  206. data_time_meter:reset() 

  207. end 


  208. engine.hooks.onEndEpoch = function(state) 

  209. local train_loss = meter:value() 

  210. confusion:updateValids() 

  211. local train_acc = confusion.totalValid*100 

  212. local train_time = train_timer:time().real 

  213. meter:reset() 

  214. print(confusion) 

  215. confusion:zero() 

  216. test_timer:reset() 


  217. local cache = state.params:clone() -- 保存现场 

  218. --state.params:copy(state.optim.ax) 

  219. test() 

  220. --state.params:copy(cache) -- 恢复现场 


  221. log:set{ 

  222. train_loss = train_loss, 

  223. train_acc = train_acc, 

  224. data_loading_time = data_time_meter:value(), 

  225. epoch = state.epoch, 

  226. test_acc = confusion.totalValid*100, 

  227. train_time = train_time, 

  228. test_time = test_timer:time().real, 

  229. } 

  230. log:flush() 

  231. end 


  232. engine.hooks.onUpdate = function(state) 

  233. data_timer:reset() 

  234. end 


  235. engine:train{ 

  236. network = fc, 

  237. criterion = criterion, 

  238. iterator = getIterator('train'), 

  239. optimMethod = optim.sgd, 

  240. config = {learningRate = 0.05, 

  241. --weightDecay = 0.05, 

  242. momentum = 0.9, 

  243. --t0 = 1e+4, 

  244. --eta0 =0.1 

  245. }, 

  246. maxepoch = 30,  

  247. } 


  248. -- 保存模型 

  249. local modelpath = 'SiaVGG16_model.t7' 

  250. print('Saving to ' .. modelpath) 

  251. torch.save(modelpath,fc:float():clearState()) 

  252. --]] 

输出:

1493386765674.jpg

发现网络太容易过拟合,主要一方面是数据太少,另一方面是视频中就那么几个人,所以patch之间的相关性太大,对网络提供的信息太少。所以使用更多的数据测试结果应该会好许多。

这个代码主要是为了熟悉torchnet package,感受呢,

  1. 对于数据的预处理,确实方便多了

  2. 如果使用提供的Engine,虽然训练过程简单了但是也太模块化了,比如某些层的微调,比如每层设置不同的学习率

  3. 使用Iterator时,尤其要小心

torchnet+VGG16计算patch之间相似度的更多相关文章

  1. (转)c# math 计算两点之间的角度公式

    计算两点之间的角度公式是: 假设点一(X1,Y1),点二(X2,Y2) double angleOfLine = Math.Atan2((Y2 - Y1), (X2 - X2)) * 180 / Ma ...

  2. python-Levenshtein几个计算字串相似度的函数解析

    linux环境下,没有首先安装python_Levenshtein,用法如下: 重点介绍几个该包中的几个计算字串相似度的几个函数实现. 1. Levenshtein.hamming(str1, str ...

  3. sql server2008根据经纬度计算两点之间的距离

    --通过经纬度计算两点之间的距离 create FUNCTION [dbo].[fnGetDistanceNew] --LatBegin 开始经度 --LngBegin 开始维度 --29.49029 ...

  4. C#面向对象思想计算两点之间距离

    题目为计算两点之间距离. 面向过程的思维方式,两点的横坐标之差,纵坐标之差,平方求和,再开跟,得到两点之间距离. using System; using System.Collections.Gene ...

  5. 2D和3D空间中计算两点之间的距离

    自己在做游戏的忘记了Unity帮我们提供计算两点之间的距离,在百度搜索了下. 原来有一个公式自己就写了一个方法O(∩_∩)O~,到僵尸到达某一个点之后就向另一个奔跑过去 /// <summary ...

  6. Jquery计算时间戳之间的差值,可返回年,月,日,小时等

    /** * 计算时间戳之间的差值 * @param startTime 开始时间戳 * @param endTime 结束时间戳 * @param type 返回指定类型差值(year, month, ...

  7. Levenshtein Distance莱文斯坦距离算法来计算字符串的相似度

    Levenshtein Distance莱文斯坦距离定义: 数学上,两个字符串a.b之间的莱文斯坦距离表示为levab(|a|, |b|). levab(i, j) = max(i, j)  如果mi ...

  8. <tf-idf + 余弦相似度> 计算文章的相似度

    背景知识: (1)tf-idf 按照词TF-IDF值来衡量该词在该文档中的重要性的指导思想:如果某个词比较少见,但是它在这篇文章中多次出现,那么它很可能就反映了这篇文章的特性,正是我们所需要的关键词. ...

  9. numpy :: 计算特征之间的余弦距离

    余弦距离在计算相似度的应用中经常使用,比如: 文本相似度检索 人脸识别检索 相似图片检索 原理简述 下面是余弦相似度的计算公式(图来自wikipedia): 但是,余弦相似度和常用的欧式距离的有所区别 ...

随机推荐

  1. javascript关闭网页的几种方法

    js关闭当前页面(窗口)的几种方式总结,需要的朋友可以参考一下: 1. 不带任何提示关闭窗口的js代码 <a href="javascript:window.opener=null;w ...

  2. python及numpy,pandas易混淆的点

    https://blog.csdn.net/happyhorizion/article/details/77894035 初接触python觉得及其友好(类似matlab),尤其是一些令人拍案叫绝不可 ...

  3. CentOS 6.4下Squid代理服务器的安装与配置(转)

    add by zhj: 其实我们主要还是关注它在服务器端使用时,充当反向代理和静态数据缓存.至于普通代理和透明代理,其实相当于客户端做的事,和服务端没有什么关系.另外,Squid的缓存主要是缓存在硬盘 ...

  4. 正则表达式python

    import re # re.match() 能够匹配出以xxx开头的字符串 ret = re.match(r"H", "Hello Python") # pr ...

  5. python技巧总结之set、日志、rsa加密

    一.日志模块logging模块调用 1.日志模块使用原理 #!/usr/bin/python # -*- coding:utf-8 -*- import logging # 方式一: "&q ...

  6. WebDriver API 实例详解(二)

    十一.双击某个元素 被测试网页的html源码: <html> <head> <meta charset="UTF-8"> </head&g ...

  7. Rails的HashWithIndifferentAccess

    ruby 2.0 引入了keyword arguments,方法的参数可以这么声明 def foo(bar: 'default') puts bar end foo # => 'default' ...

  8. 587. Erect the Fence(凸包算法)

    问题 给定一群树的坐标点,画个围栏把所有树围起来(凸包). 至少有一棵树,输入和输出没有顺序. Input: [[1,1],[2,2],[2,0],[2,4],[3,3],[4,2]] Output: ...

  9. IDFA踩坑记录

    IDFA踩坑记录: 1.iOS10.0 以下,即使打开“限制广告跟踪”,依然可以读取idfa: 2.打开“限制广告跟踪”,然后再关闭“限制广告跟踪”,idfa会改变: 3.越狱机器安装开发证书打的包, ...

  10. TED #06# Questioning the universe

    Stephen Hawking: Questioning the universe 1. 第一段: There is nothing bigger or older than the universe ...