关于MatConvNet的反向传播原理及源码解析

关于MatConvNet的反向传播原理及源码解析

MatConvNet深度学习库的反向传播原理如下图所示:

 

上图是简化版的反向传播原理推导,具体的张量表示法参考相关手册

为了方便对比及具体详尽的原理推导,在手册中,关于反向传播的推导部分截图如下:

其中在上图中,重点注意的地方就是式(2.2),它只是关于x的导数。为了引起重视,下面有一个练习:

 

[html] view plain copy
 
  1. %下面实现一个由一个卷积层和ReLU层构成的两层卷积网络的反向传播实现示例  
  2.   
  3. %前向模式:计算卷积和卷积后的ReLU输出  
  4. y = vl_nnconv(x, w, []) ;  
  5. z = vl_nnrelu(y) ;  
  6.   
  7. %初始化一个随机投影张量  
  8. p = randn(size(z), 'single') ;  
  9.   
  10. %反向传播模式:映射偏导  
  11. dy = vl_nnrelu(z, p) ;  
  12. [dx,dw] = vl_nnconv(x, w, [], dy) ;  
那么,修改上述代码,实现一个级联的Conv+ReLU+Conv网络,该如何实现呢?尤其是在反向模式下,vl_nnconv的第四个参数偏导该传入什么呢?会在后面给出完整参考代码。

下面进行反向传播的MatConvNet的反向传播实现方式:

 

[plain] view plain copy
 
  1. res = struct(...  
  2.     'x', cell(1,n+1), ...  
  3.     'dzdx', cell(1,n+1), ...  
  4.     'dzdw', cell(1,n+1), ...  
  5.     'aux', cell(1,n+1), ...  
  6.     'stats', cell(1,n+1), ...  
  7.     'time', num2cell(zeros(1,n+1)), ...  
  8.     'backwardTime', num2cell(zeros(1,n+1))) ;  
  9.   
  10.   
  11. % -------------------------------------------------------------------------  
  12. %                                                             Backward pass  
  13. % -------------------------------------------------------------------------  
  14.   
  15. if doder  
  16.   res(n+1).dzdx = dzdy ;  
  17.   for i=n:-1:backPropLim  
  18.     l = net.layers{i} ;  
  19.     res(i).backwardTime = tic ;  
  20.     switch l.type  
  21.   
  22.       case 'conv'  
  23.         [res(i).dzdx, dzdw{1}, dzdw{2}] = ...  
  24.           vl_nnconv(res(i).x, l.weights{1}, l.weights{2}, res(i+1).dzdx, ...  
  25.           'pad', l.pad, ...  
  26.           'stride', l.stride, ...  
  27.           'dilate', l.dilate, ...  
  28.           l.opts{:}, ...  
  29.           cudnn{:}) ;  
  30.   
  31.       case 'convt'  
  32.         [res(i).dzdx, dzdw{1}, dzdw{2}] = ...  
  33.           vl_nnconvt(res(i).x, l.weights{1}, l.weights{2}, res(i+1).dzdx, ...  
  34.           'crop', l.crop, ...  
  35.           'upsample', l.upsample, ...  
  36.           'numGroups', l.numGroups, ...  
  37.           l.opts{:}, ...  
  38.           cudnn{:}) ;  
  39.   
  40.       case 'pool'  
  41.         res(i).dzdx = vl_nnpool(res(i).x, l.pool, res(i+1).dzdx, ...  
  42.                                 'pad', l.pad, 'stride', l.stride, ...  
  43.                                 'method', l.method, ...  
  44.                                 l.opts{:}, ...  
  45.                                 cudnn{:}) ;  
  46.   
  47.       case {'normalize', 'lrn'}  
  48.         res(i).dzdx = vl_nnnormalize(res(i).x, l.param, res(i+1).dzdx) ;  
  49.   
  50.       case 'softmax'  
  51.         res(i).dzdx = vl_nnsoftmax(res(i).x, res(i+1).dzdx) ;  
  52.   
  53.       case 'loss'  
  54.         res(i).dzdx = vl_nnloss(res(i).x, l.class, res(i+1).dzdx) ;  
  55.   
  56.       case 'softmaxloss'  
  57.         res(i).dzdx = vl_nnsoftmaxloss(res(i).x, l.class, res(i+1).dzdx) ;  
  58.   
  59.       case 'relu'  
  60.         if l.leak > 0, leak = {'leak', l.leak} ; else leak = {} ; end  
  61.         if ~isempty(res(i).x)  
  62.           res(i).dzdx = vl_nnrelu(res(i).x, res(i+1).dzdx, leak{:}) ;  
  63.         else  
  64.           % if res(i).x is empty, it has been optimized away, so we use this  
  65.           % hack (which works only for ReLU):  
  66.           res(i).dzdx = vl_nnrelu(res(i+1).x, res(i+1).dzdx, leak{:}) ;  
  67.         end  
  68.   
  69.       case 'sigmoid'  
  70.         res(i).dzdx = vl_nnsigmoid(res(i).x, res(i+1).dzdx) ;  
  71.   
  72.       case 'noffset'  
  73.         res(i).dzdx = vl_nnnoffset(res(i).x, l.param, res(i+1).dzdx) ;  
  74.   
  75.       case 'spnorm'  
  76.         res(i).dzdx = vl_nnspnorm(res(i).x, l.param, res(i+1).dzdx) ;  
  77.   
  78.       case 'dropout'  
  79.         if testMode  
  80.           res(i).dzdx = res(i+1).dzdx ;  
  81.         else  
  82.           res(i).dzdx = vl_nndropout(res(i).x, res(i+1).dzdx, ...  
  83.                                      'mask', res(i+1).aux) ;  
  84.         end  
  85.   
  86.       case 'bnorm'  
  87.         [res(i).dzdx, dzdw{1}, dzdw{2}, dzdw{3}] = ...  
  88.           vl_nnbnorm(res(i).x, l.weights{1}, l.weights{2}, res(i+1).dzdx, ...  
  89.                      'epsilon', l.epsilon, ...  
  90.                      bnormCudnn{:}) ;  
  91.         % multiply the moments update by the number of images in the batch  
  92.         % this is required to make the update additive for subbatches  
  93.         % and will eventually be normalized away  
  94.         dzdw{3} = dzdw{3} * size(res(i).x,4) ;  
  95.   
  96.       case 'pdist'  
  97.         res(i).dzdx = vl_nnpdist(res(i).x, l.class, ...  
  98.           l.p, res(i+1).dzdx, ...  
  99.           'noRoot', l.noRoot, ...  
  100.           'epsilon', l.epsilon, ...  
  101.           'aggregate', l.aggregate, ...  
  102.           'instanceWeights', l.instanceWeights) ;  
  103.   
  104.       case 'custom'  
  105.         res(i) = l.backward(l, res(i), res(i+1)) ;  
  106.   
  107.     end % layers  
  108.   
  109.     switch l.type  
  110.       case {'conv', 'convt', 'bnorm'}  
  111.         if ~opts.accumulate  
  112.           res(i).dzdw = dzdw ;  
  113.         else  
  114.           for j=1:numel(dzdw)  
  115.             res(i).dzdw{j} = res(i).dzdw{j} + dzdw{j} ;  
  116.           end  
  117.         end  
  118.         dzdw = [] ;  
  119.         if ~isempty(opts.parameterServer) && ~opts.holdOn  
  120.           for j = 1:numel(res(i).dzdw)  
  121.             opts.parameterServer.push(sprintf('l%d_%d',i,j),res(i).dzdw{j}) ;  
  122.             res(i).dzdw{j} = [] ;  
  123.           end  
  124.         end  
  125.     end  
  126.     if opts.conserveMemory && ~net.layers{i}.precious && i ~= n  
  127.       res(i+1).dzdx = [] ;  
  128.       res(i+1).x = [] ;  
  129.     end  
  130.     if gpuMode && opts.sync  
  131.       wait(gpuDevice) ;  
  132.     end  
  133.     res(i).backwardTime = toc(res(i).backwardTime) ;  
  134.   end  
  135.   if i > 1 && i == backPropLim && opts.conserveMemory && ~net.layers{i}.precious  
  136.     res(i).dzdx = [] ;  
  137.     res(i).x = [] ;  
  138.   end  
  139. end  

由上述实现方式,练习题的参考答案如下:

 

[plain] view plain copy
 
  1. %下面实现一个由Conv+ReLU+Conv构成的三层卷积网络的反向传播实现示例  
  2.   
  3. %前向模式:计算卷积和卷积后的ReLU输出  
  4. y = vl_nnconv(x, w, []) ;  
  5. z = vl_nnrelu(y) ;  
  6. y1=vl_nnconv(z,w,[]);  
  7.   
  8. %初始化一个随机投影张量  
  9. p = randn(size(y1), 'single') ;  
  10.   
  11. %反向传播模式:映射偏导  
  12. [dz,dw]=vl_nnconv(z,w,[],p);  
  13. dy = vl_nnrelu(y, dz) ;  
  14. [dx,dw] = vl_nnconv(x, w, [], dy) ;  


版权声明:本文为博主原创文章,未经博主允许不得转载。 https://blog.csdn.net/u011501388/article/details/79381773

 

posted @ 2018-03-30 22:36  菜鸡一枚  阅读(359)  评论(0)    收藏  举报