| 25 | namespace singa { |
| 26 | |
| 27 | PoolingHandle::PoolingHandle(const Tensor &input, |
| 28 | const std::vector<int> &kernel_size, |
| 29 | const std::vector<int> &stride, |
| 30 | const std::vector<int> &padding, |
| 31 | const bool is_max) { |
| 32 | kernel_h = kernel_size[0]; |
| 33 | kernel_w = kernel_size[1]; |
| 34 | |
| 35 | pad_h = padding[0]; |
| 36 | pad_w = padding[1]; |
| 37 | |
| 38 | stride_h = stride[0]; |
| 39 | stride_w = stride[1]; |
| 40 | |
| 41 | batchsize = input.shape(0); |
| 42 | channels = input.shape(1); |
| 43 | height = input.shape(2); |
| 44 | width = input.shape(3); |
| 45 | |
| 46 | pooled_height = 1; |
| 47 | |
| 48 | if (stride_h > 0) |
| 49 | pooled_height = |
| 50 | std::floor(((height + 2 * pad_h - kernel_h) / stride_h)) + 1; |
| 51 | pooled_width = std::floor(((width + 2 * pad_w - kernel_w) / stride_w)) + 1; |
| 52 | is_max_pooling = is_max; |
| 53 | |
| 54 | #ifdef USE_DNNL |
| 55 | if (input.device()->lang() == kCpp) { |
| 56 | auto x_dims = |
| 57 | dnnl::memory::dims(input.shape().begin(), input.shape().end()); |
| 58 | auto y_dims = |
| 59 | dnnl::memory::dims({batchsize, channels, pooled_height, pooled_width}); |
| 60 | auto s_dims = dnnl::memory::dims(stride.begin(), stride.end()); |
| 61 | auto k_dims = dnnl::memory::dims(kernel_size.begin(), kernel_size.end()); |
| 62 | |
| 63 | auto p_dims = dnnl::memory::dims(padding.begin(), padding.end()); |
| 64 | |
| 65 | auto dtype_ = dnnl::memory::data_type::f32; |
| 66 | auto format_tag_ = get_dnnl_format_tag(input); |
| 67 | x_md = dnnl::memory::desc({x_dims}, dtype_, format_tag_); |
| 68 | y_md = dnnl::memory::desc({y_dims}, dtype_, format_tag_); |
| 69 | |
| 70 | // allow max or avg (follow cudnn implementation convention) |
| 71 | auto pooling_algo = dnnl::algorithm::pooling_avg_exclude_padding; |
| 72 | if (is_max_pooling) pooling_algo = dnnl::algorithm::pooling_max; |
| 73 | |
| 74 | auto pool_fwd_d = dnnl::pooling_forward::desc( |
| 75 | dnnl::prop_kind::forward_training, pooling_algo, x_md, y_md, s_dims, |
| 76 | k_dims, p_dims, p_dims); |
| 77 | auto pool_bwd_d = dnnl::pooling_backward::desc( |
| 78 | pooling_algo, x_md, y_md, s_dims, k_dims, p_dims, p_dims); |
| 79 | |
| 80 | auto eng = input.device()->context(0)->dnnl_engine; |
| 81 | pool_fwd_pd = dnnl::pooling_forward::primitive_desc(pool_fwd_d, eng); |
| 82 | pool_bwd_pd = |
| 83 | dnnl::pooling_backward::primitive_desc(pool_bwd_d, eng, pool_fwd_pd); |
| 84 | |