@@ -17,8 +17,187 @@ BatchNormHandle::BatchNormHandle(const float momentum, const Tensor& input) {
1717 } else {
1818 LOG (FATAL ) << " The dimension of input should either be 4D or 2D." ;
1919 }
20+
21+
22+ #ifdef USE_MKLDNN
23+ dtype = GetMKLDNNDataType (input.data_type ());
24+ epsilon =1e-5f ;
25+ data_memory_format = is_2d ? mkldnn::memory::format::nc : mkldnn::memory::format::nchw;
26+ if (is_2d) {
27+ x_dims = {(int )batchsize, (int )channels};
28+ y_dims = {(int )batchsize, (int )channels};
29+ } else {
30+ x_dims = {(int )batchsize, (int )channels, (int )height, (int )width};
31+ y_dims = {(int )batchsize, (int )channels, (int )height, (int )width};
32+ }
33+
34+ auto eng = *input.device ()->context (0 )->engine ;
35+ x_md = new mkldnn::memory::desc (x_dims, dtype, data_memory_format);
36+ dx_md = new mkldnn::memory::desc (x_dims, dtype, data_memory_format);
37+ bn_fwd_d = new mkldnn::batch_normalization_forward::desc (mkldnn::forward_training, *x_md, epsilon,
38+ mkldnn::use_scale_shift);
39+ bn_fwd_pd = new mkldnn::batch_normalization_forward::primitive_desc (*bn_fwd_d, eng);
40+ #endif // USE_MKLDNN
41+
2042};
2143
44+
45+ BatchNormHandle::~BatchNormHandle () {
46+ #ifdef USE_MKLDNN
47+ delete (x_md);
48+ delete (dx_md);
49+ delete (bn_fwd_d);
50+ delete (bn_fwd_pd);
51+ #endif // USE_MKLDNN
52+ }
53+
54+ #ifdef USE_MKLDNN
55+
56+ Tensor CpuBatchNormForwardInference (const BatchNormHandle &bnh, const Tensor& x, const Tensor& bnScale, const Tensor& bnBias,
57+ Tensor& running_mean, Tensor& running_var){
58+
59+ CHECK_EQ (x.device ()->lang (), kCpp );
60+ Tensor y;
61+ y.ResetLike (x);
62+
63+
64+ Tensor w = get_bn_weight_from (bnScale, bnBias);
65+
66+ y.device ()->Exec ([&y, &x, &w, &bnh](Context *ctx) {
67+ try {
68+ auto eng = *ctx->engine ;
69+ using namespace mkldnn ;
70+ auto x_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng}, x.block ()->mutable_data ());
71+ auto y_mem = memory ({{{bnh.y_dims }, bnh.dtype , bnh.data_memory_format }, eng}, y.block ()->mutable_data ());
72+
73+ auto bn_fwd_d = batch_normalization_forward::desc (forward_inference, *bnh.x_md , bnh.epsilon , use_scale_shift);
74+ auto bn_fwd_pd = batch_normalization_forward::primitive_desc (bn_fwd_d, eng);
75+
76+ auto w_mem = memory (bn_fwd_pd.weights_primitive_desc (), w.block ()->mutable_data ());
77+
78+ auto bn = batch_normalization_forward (bn_fwd_pd, x_mem, w_mem, y_mem);
79+
80+ stream (stream::kind::eager).submit ({bn}).wait ();
81+ }
82+ catch (mkldnn::error &e) {
83+ InitLogging (" " );
84+ LOG (FATAL ) << " MKLDNN Batch Norm" << " Status: " << e.status << " Message: " << e.message ;
85+ }
86+
87+ }, {y.block (), x.block (), w.block ()}, {y.block ()});
88+
89+ return y;
90+
91+ }
92+
93+ const std::vector<Tensor>
94+ CpuBatchNormForwardTraining (const BatchNormHandle &bnh, const Tensor &x, const Tensor &bnScale, const Tensor &bnBias,
95+ Tensor &running_mean, Tensor &running_var) {
96+
97+ Tensor y;
98+ y.ResetLike (x);
99+
100+ // mean and var for local batch
101+ Tensor mean;
102+ mean.ResetLike (running_mean);
103+ Tensor var;
104+ var.ResetLike (running_var);
105+
106+ // combine scale and bias to construct weight tensor in required format for backward
107+ Tensor w = get_bn_weight_from (bnScale, bnBias);
108+
109+ y.device ()->Exec ([&x, &y, &mean, &var, &w, &bnh](Context *ctx) {
110+ try {
111+ auto eng = *ctx->engine ;
112+ using namespace mkldnn ;
113+
114+ auto x_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng},
115+ x.block ()->mutable_data ());
116+ auto y_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng},
117+ y.block ()->mutable_data ());
118+ auto m_mem = memory (bnh.bn_fwd_pd ->mean_primitive_desc (), mean.block ()->mutable_data ());
119+
120+ auto v_mem = memory (bnh.bn_fwd_pd ->variance_primitive_desc (), var.block ()->mutable_data ());
121+
122+ auto w_mem = memory (bnh.bn_fwd_pd ->weights_primitive_desc (),w.block ()->mutable_data ());
123+
124+ auto bn_fwd = batch_normalization_forward (*bnh.bn_fwd_pd , x_mem, w_mem, y_mem, m_mem, v_mem);
125+
126+ stream (stream::kind::eager).submit ({bn_fwd}).wait ();
127+ }
128+ catch (mkldnn::error &e) {
129+ singa::InitLogging (" " );
130+ LOG (FATAL ) << " MKLDNN Batch Norm Backward" << " Status: " << e.status << " Message: " << e.message ;
131+ }
132+ }, {x.block (), w.block ()}, {y.block (), mean.block (), var.block ()});
133+
134+
135+ // local implemented running mean as mkldnn does not support it yet:
136+ // https://github.com/intel/mkl-dnn/issues/371
137+ running_mean = running_mean*bnh.factor + mean*(1 -bnh.factor );
138+ running_var = running_var*bnh.factor + var*(1 -bnh.factor );
139+
140+
141+ return {y, running_mean, running_var};
142+
143+ }
144+
145+ const std::vector<Tensor> CpuBatchNormBackwardx (const BatchNormHandle &bnh,
146+ const Tensor &y, const Tensor &dy,
147+ const Tensor &x,
148+ const Tensor &bnScale, const Tensor &bnBias,
149+ const Tensor &mean, const Tensor &var){
150+ Tensor dx;
151+ dx.ResetLike (dy);
152+
153+ // combine scale and bias to construct weight tensor in required format for backward
154+ Tensor w = get_bn_weight_from (bnScale, bnBias);
155+
156+ Tensor dw (Shape{bnScale.Size (),bnBias.Size ()});
157+
158+ dx.device ()->Exec ([&dw, &x, &dx, &y, &dy, &w, &mean, &var, &bnh](Context *ctx) {
159+
160+ try {
161+ auto eng = *ctx->engine ;
162+ using namespace mkldnn ;
163+
164+ auto x_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng}, x.block ()->mutable_data ());
165+ auto dx_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng}, dx.block ()->mutable_data ());
166+ auto y_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng}, y.block ()->mutable_data ());
167+ auto dy_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng}, dy.block ()->mutable_data ());
168+
169+ auto m_mem = memory (bnh.bn_fwd_pd ->mean_primitive_desc (), mean.block ()->mutable_data ());
170+ auto v_mem = memory (bnh.bn_fwd_pd ->variance_primitive_desc (), var.block ()->mutable_data ());
171+ auto w_mem = memory (bnh.bn_fwd_pd ->weights_primitive_desc (), w.block ()->mutable_data ());
172+
173+
174+ auto bn_bwd_d = batch_normalization_backward::desc (backward, *bnh.dx_md , *bnh.x_md , bnh.epsilon , use_scale_shift);
175+ auto bn_bwd_pd = batch_normalization_backward::primitive_desc (bn_bwd_d, eng, *bnh.bn_fwd_pd );
176+
177+
178+ auto dw_mem = memory (bn_bwd_pd.diff_weights_primitive_desc (), dw.block ()->mutable_data ());
179+
180+ auto bn_bwd = batch_normalization_backward (bn_bwd_pd, x_mem, m_mem, v_mem, dy_mem, w_mem, dx_mem, dw_mem);
181+
182+ stream (stream::kind::eager).submit ({bn_bwd}).wait ();
183+ }
184+ catch (mkldnn::error &e) {
185+ singa::InitLogging (" " );
186+ LOG (FATAL ) << " MKLDNN Batch Norm Backward" << " Status: " << e.status << " Message: " << e.message ;
187+ }
188+
189+ }, {x.block (), dy.block (), mean.block (), var.block ()},
190+ {dx.block (), dw.block ()});
191+
192+ Tensor dbnScale = CopyRows (dw, 0 , bnScale.Size ());
193+ Tensor dbnBias = CopyRows (dw, 1 , bnBias.Size ());
194+
195+ return {dx, dbnScale, dbnBias};
196+ }
197+
198+
199+ #endif // USE_MKLDNN
200+
22201#ifdef USE_CUDNN
23202CudnnBatchNormHandle::CudnnBatchNormHandle (const float momentum,
24203 const Tensor& input): BatchNormHandle(momentum, input) {
0 commit comments