MCPcopy Create free account
hub / github.com/apache/singa / Backward

Method Backward

src/model/layer/pooling.cc:125–160  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

123}
124
125const std::pair<Tensor, vector<Tensor>>
126Pooling::Backward(int flag, const Tensor& grad) {
127 CHECK_EQ(grad.device()->lang(), kCpp);
128 CHECK_EQ(grad.nDim(), 4u);
129
130 vector<Tensor> param_grad;
131
132 auto batchsize = grad.shape(0);
133 auto dtype = grad.data_type();
134 auto dev = grad.device();
135 Shape shape{batchsize, channels_, height_, width_};
136
137 Tensor dx(shape, dev, dtype);
138 auto gradptr = grad.data<float>();
139 float* dxptr = new float[dx.Size()];
140
141 if (pool_ == PoolingConf_PoolMethod_MAX) {
142 CHECK(!buf_.empty());
143 Tensor mask = buf_.top();
144 buf_.pop();
145 auto maskptr = mask.data<float>();
146 BackwardMaxPooling(gradptr, maskptr, batchsize, channels_, height_, width_,
147 pooled_height_, pooled_width_, kernel_h_, kernel_w_,
148 pad_h_, pad_w_, stride_h_, stride_w_, dxptr);
149 } else if (pool_ == PoolingConf_PoolMethod_AVE) {
150 BackwardAvgPooling(gradptr, batchsize, channels_, height_, width_,
151 pooled_height_, pooled_width_, kernel_h_, kernel_w_,
152 pad_h_, pad_w_, stride_h_, stride_w_, dxptr);
153 } else {
154 LOG(FATAL) << "Unknown pooling method";
155 }
156
157 dx.CopyDataFromHostPtr(dxptr, dx.Size());
158 delete[] dxptr;
159 return std::make_pair(dx, param_grad);
160}
161
162void Pooling::ForwardMaxPooling(const float* bottom, const int num,
163 const int channels,

Callers

nothing calls this directly

Calls 8

langMethod · 0.80
deviceMethod · 0.80
nDimMethod · 0.80
shapeMethod · 0.80
data_typeMethod · 0.80
emptyMethod · 0.80
SizeMethod · 0.45
CopyDataFromHostPtrMethod · 0.45

Tested by

no test coverage detected