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

Method PoolingHandle

src/model/operation/pooling.cc:27–89  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

25namespace singa {
26
27PoolingHandle::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

Callers 5

re_new_handleFunction · 0.80
initializeMethod · 0.80
test_poolingMethod · 0.80
test_dnnl_pooling_maxMethod · 0.80
test_dnnl_pooling_avgMethod · 0.80

Calls 8

get_dnnl_format_tagFunction · 0.85
shapeMethod · 0.80
langMethod · 0.80
deviceMethod · 0.80
contextMethod · 0.80
floorFunction · 0.50
beginMethod · 0.45
endMethod · 0.45

Tested by 3

test_poolingMethod · 0.64
test_dnnl_pooling_maxMethod · 0.64
test_dnnl_pooling_avgMethod · 0.64