@@ -134,6 +134,14 @@ mkldnn::algorithm GetMKLDNNPoolAlgo(const PoolingParam ¶m) {
134134 }
135135}
136136
137+ static inline int GetPaddingSizeFull (int x, int padl, int padr, int k, int s) {
138+ if ((x + padl + padr - k) % s != 0 ) {
139+ return (padr + s - ((x + padl + padr - k) % s));
140+ } else {
141+ return padr;
142+ }
143+ }
144+
137145mkldnn::pooling_forward::primitive_desc GetPoolingFwdPdesc (
138146 const PoolingParam ¶m, const bool is_train, const memory::desc &data_md,
139147 const memory::desc &out_md) {
@@ -154,11 +162,17 @@ mkldnn::pooling_forward::primitive_desc GetPoolingFwdPdesc(
154162 int pad_l_ = param.pad [1 ], pad_r_ = param.pad [1 ];
155163 int stride_h_ = param.stride [0 ], stride_w_ = param.stride [1 ];
156164
165+ if (param.pooling_convention == pool_enum::kFull ) {
166+ pad_b_ = GetPaddingSizeFull (data_md.data .dims [2 ], pad_t_, pad_b_, kernel_h_, stride_h_);
167+ pad_r_ = GetPaddingSizeFull (data_md.data .dims [3 ], pad_l_, pad_r_, kernel_w_, stride_w_);
168+ }
169+
157170 const mkldnn::engine engine = CpuEngine::Get ()->get_engine ();
158171 if (param.global_pool ) {
159172 pad_t_ = pad_b_ = pad_l_ = pad_r_ = 0 ;
160173 stride_h_ = stride_w_ = 1 ;
161174 }
175+
162176 if (pad_t_ != 0 || pad_l_ != 0 ) {
163177 CHECK (param.pool_type == pool_enum::kAvgPooling ||
164178 param.pool_type == pool_enum::kMaxPooling )
@@ -167,7 +181,6 @@ mkldnn::pooling_forward::primitive_desc GetPoolingFwdPdesc(
167181 CHECK_LT (pad_t_, kernel_h_);
168182 }
169183
170-
171184 const mkldnn::algorithm alg = GetMKLDNNPoolAlgo (param);
172185 mkldnn::prop_kind kind = mkldnn::prop_kind::forward_scoring;
173186 if (is_train && alg != algorithm::pooling_avg) {
@@ -227,17 +240,22 @@ MKLDNNPoolingFwd &GetPoolingFwd(const PoolingParam ¶m,
227240 int pad_l_ = param.pad [1 ], pad_r_ = param.pad [1 ];
228241 int stride_h_ = param.stride [0 ], stride_w_ = param.stride [1 ];
229242
243+ if (param.pooling_convention == pool_enum::kFull ) {
244+ pad_b_ = GetPaddingSizeFull (data_md.data .dims [2 ], pad_t_, pad_b_, kernel_h_, stride_h_);
245+ pad_r_ = GetPaddingSizeFull (data_md.data .dims [3 ], pad_l_, pad_r_, kernel_w_, stride_w_);
246+ }
247+
230248 if (param.global_pool ) {
231- pad_t_ = pad_b_ = pad_l_ = pad_r_ = 0 ;
232- stride_h_ = stride_w_ = 1 ;
249+ pad_t_ = pad_b_ = pad_l_ = pad_r_ = 0 ;
250+ stride_h_ = stride_w_ = 1 ;
233251 }
234252
235253 if (pad_t_ != 0 || pad_l_ != 0 ) {
236- CHECK (param.pool_type == pool_enum::kAvgPooling ||
237- param.pool_type == pool_enum::kMaxPooling )
238- << " Padding implemented only for average and max pooling." ;
239- CHECK_LT (pad_l_, kernel_w_);
240- CHECK_LT (pad_t_, kernel_h_);
254+ CHECK (param.pool_type == pool_enum::kAvgPooling ||
255+ param.pool_type == pool_enum::kMaxPooling )
256+ << " Padding implemented only for average and max pooling." ;
257+ CHECK_LT (pad_l_, kernel_w_);
258+ CHECK_LT (pad_t_, kernel_h_);
241259 }
242260
243261 const mkldnn::algorithm alg = GetMKLDNNPoolAlgo (param);
@@ -353,6 +371,12 @@ MKLDNNPoolingBwd &GetPoolingBwd(const PoolingParam ¶m,
353371 int pad_t_ = param.pad [0 ], pad_b_ = param.pad [0 ];
354372 int pad_l_ = param.pad [1 ], pad_r_ = param.pad [1 ];
355373 int stride_h_ = param.stride [0 ], stride_w_ = param.stride [1 ];
374+
375+ if (param.pooling_convention == pool_enum::kFull ) {
376+ pad_b_ = GetPaddingSizeFull (data_md.data .dims [2 ], pad_t_, pad_b_, kernel_h_, stride_h_);
377+ pad_r_ = GetPaddingSizeFull (data_md.data .dims [3 ], pad_l_, pad_r_, kernel_w_, stride_w_);
378+ }
379+
356380 if (param.global_pool ) {
357381 pad_t_ = pad_b_ = pad_l_ = pad_r_ = 0 ;
358382 stride_h_ = stride_w_ = 1 ;
0 commit comments