| 205 | } |
| 206 | |
| 207 | Tensor GpuPoolingBackward(const CudnnPoolingHandle &cph, const Tensor &dy, |
| 208 | const Tensor &x, const Tensor &y) { |
| 209 | CHECK_EQ(dy.device()->lang(), kCuda); |
| 210 | CHECK_EQ(dy.nDim(), 4u); |
| 211 | |
| 212 | Tensor dx; |
| 213 | dx.ResetLike(x); |
| 214 | |
| 215 | dx.device()->Exec( |
| 216 | [dx, dy, x, y, &cph](Context *ctx) mutable { |
| 217 | float alpha = 1.0f, beta = 0.0f; |
| 218 | cudnnPoolingBackward(ctx->cudnn_handle, cph.pool_desc, &alpha, |
| 219 | cph.y_desc, y.block()->data(), cph.y_desc, |
| 220 | dy.block()->data(), cph.x_desc, x.block()->data(), |
| 221 | &beta, cph.x_desc, dx.block()->mutable_data()); |
| 222 | }, |
| 223 | {dy.block(), y.block(), x.block()}, {dx.block()}, "GpuPoolingBackward"); |
| 224 | |
| 225 | return dx; |
| 226 | }; |
| 227 | #endif // USE_CUDNN |
| 228 | |
| 229 | } // namespace singa |