| 123 | } |
| 124 | |
| 125 | const std::pair<Tensor, vector<Tensor>> |
| 126 | Pooling::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 | |
| 162 | void Pooling::ForwardMaxPooling(const float* bottom, const int num, |
| 163 | const int channels, |