[ Web Proxy ]
URL:
Viewing: https://raw.githubusercontent.com/2php/CaffeOnACL/master/src/caffe/layers/local_connect_layer.cpp [Back]  [Original]

#include 

#include "caffe/filler.hpp"
#include "caffe/layer.hpp"
#include "caffe/util/im2col.hpp"
#include "caffe/util/math_functions.hpp"
#include "caffe/layers/local_connect_layer.hpp"

namespace caffe {

template 
void LocalConnectLayer::LayerSetUp(const vector& bottom,
      const vector& top) {
  CHECK_EQ(bottom.size(), 1) layer_param_.local_param().stride();
  pad_ = this->layer_param_.local_param().pad();
  num_ = bottom[0]->num();
  channels_ = bottom[0]->channels();
  height_ = bottom[0]->height();
  width_ = bottom[0]->width();
  num_output_ = this->layer_param_.local_param().num_output();
  dilation_=1;

  height_out_ = (height_ + 2 * pad_ - kernel_size_) / stride_ + 1;
  width_out_ = (width_ + 2 * pad_ - kernel_size_) / stride_ + 1;

  M_ = num_output_;
  K_ = channels_ * kernel_size_ * kernel_size_;
  N_ = height_out_ * width_out_;

  CHECK_GT(num_output_, 0); 
  CHECK_GE(height_, kernel_size_) blobs_.size() > 0) {
    LOG(INFO) blobs_.resize(2);
    } else {
      this->blobs_.resize(1);
    }
    // Intialize the weight
    this->blobs_[0].reset(new Blob(
        num_output_, 1, K_, N_));
    // fill the weights
    shared_ptr weight_filler(GetFiller(
        this->layer_param_.local_param().weight_filler()));
    weight_filler->Fill(this->blobs_[0].get());
    // If necessary, intiialize and fill the bias term
    if (bias_term_) {
      this->blobs_[1].reset(new Blob(1, 1, M_, N_));
      shared_ptr bias_filler(GetFiller(
          this->layer_param_.local_param().bias_filler()));
      bias_filler->Fill(this->blobs_[1].get());  
    }
  }
}

template 
void LocalConnectLayer::Reshape(const vector& bottom,
      const vector& top) {
  CHECK_EQ(bottom[0]->channels(), channels_) num()) channels())
        height())
        width())
        Reshape(num_, num_output_, height_out_, width_out_);
  }

  // The im2col result buffer would only hold one image at a time to avoid
  // overly large memory usage.
  col_buffer_.Reshape(
      1, channels_ * kernel_size_ * kernel_size_, height_out_, width_out_);

  for (int top_id = 0; top_id < top.size(); ++top_id) {
    top[top_id]->Reshape(num_, num_output_, height_out_, width_out_);
  }
}

template 
void LocalConnectLayer::Forward_cpu(const vector& bottom,
      const vector& top) {

  Dtype* x_data = col_buffer_.mutable_cpu_data();
  const Dtype* weight = this->blobs_[0]->cpu_data();
  const Dtype* bottom_data = bottom[0]->cpu_data();
  Dtype* top_data = top[0]->mutable_cpu_data();

  Blob E;
  E.Reshape(1, 1, 1, K_);
  FillerParameter filler_param;
  filler_param.set_value(1);
  ConstantFiller filler(filler_param);
  filler.Fill(&E);

  Blob intermediate;
  intermediate.Reshape(1, 1, K_, N_);
  for (int n=0; noffset(n), channels_, height_,
               width_, kernel_size_, kernel_size_, pad_, pad_, stride_, stride_,dilation_,dilation_, x_data);

    for (int m=0; mblobs_[0]->offset(m),
                intermediate.mutable_cpu_data());

      caffe_cpu_gemm(CblasNoTrans, CblasNoTrans, 1, N_, K_,
                            (Dtype)1., E.cpu_data(),
                            intermediate.cpu_data(),
                            (Dtype)0., top_data + top[0]->offset(n, m));
    }

    if (bias_term_) {
      caffe_add(M_ * N_, this->blobs_[1]->cpu_data(),
                top_data + top[0]->offset(n),
                top_data + top[0]->offset(n));
    }
  }
}

template 
void LocalConnectLayer::Backward_cpu(const vector& top,
      const vector& propagate_down, const vector& bottom) {

  const Dtype* top_diff = top[0]->cpu_diff();
  const Dtype* bottom_data = bottom[0]->cpu_data();
  Dtype* bottom_diff = bottom[0]->mutable_cpu_diff();
  Dtype* x_data = col_buffer_.mutable_cpu_data();
  Dtype* x_diff = col_buffer_.mutable_cpu_diff();
  const Dtype* weight = this->blobs_[0]->cpu_data();
  Dtype* weight_diff = this->blobs_[0]->mutable_cpu_diff();
  Dtype* bias_diff = NULL;

  Blob intermediate;
  intermediate.Reshape(1, 1, 1, N_);

  Blob xt;
  xt.Reshape(1, 1, K_, N_);
  Dtype* xt_data = xt.mutable_cpu_data();

  if (bias_term_) {
    bias_diff = this->blobs_[1]->mutable_cpu_diff();
    memset(bias_diff, 0, sizeof(Dtype) * this->blobs_[1]->count());
    for (int n = 0; n < num_; ++n) {
      caffe_add(M_ * N_, bias_diff,
                top_diff + top[0]->offset(n),
                bias_diff);
    }
  }

  memset(weight_diff, 0, sizeof(Dtype) * this->blobs_[0]->count());
  for (int n=0; noffset(n), channels_, height_,
               width_, kernel_size_, kernel_size_, pad_, pad_, stride_, stride_, dilation_,dilation_,x_data);

    // gradient wrt weight
    for (int m=0; mblobs_[0]->offset(m);
      for (int k=0; koffset(n, m),  
                  x_data+col_buffer_.offset(0,k), xt_data+xt.offset(0,0,k));
      }
      caffe_cpu_axpby(K_*N_, Dtype(1.0), xt_data, Dtype(1.0), filter_weight_diff);
    }
      
    // gradient wrt bottom data
    if (propagate_down[0]) {
      memset(x_diff, 0, col_buffer_.count() * sizeof(Dtype));
      for (int m=0; mblobs_[0]->offset(m,0,k),
                    intermediate.mutable_cpu_data());

          caffe_cpu_axpby(N_, Dtype(1.0),
                          intermediate.cpu_data(), Dtype(1.0),
                          x_diff+col_buffer_.offset(0,k));
        }
      }

      // col2im back to the data
      col2im_cpu(x_diff, channels_, height_, width_, kernel_size_, kernel_size_,
                 pad_, pad_, stride_, stride_, dilation_,dilation_,bottom_diff + bottom[0]->offset(n));

    }
  }

}

#ifdef CPU_ONLY
STUB_GPU(LocalConnectLayer);
#endif

INSTANTIATE_CLASS(LocalConnectLayer);

}  // namespace caffe

Web Proxy Viewer  |  New URL  |  Original Page