#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