@@ -20,139 +20,141 @@ BatchNormHandle::BatchNormHandle(const float momentum, const Tensor& input) {
2020
2121
2222#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- }
23+ if (input.device ()->lang () == kCpp ) {
24+ dtype = GetMKLDNNDataType (input.data_type ());
25+ epsilon = 1e-5f ;
26+ data_memory_format = is_2d ? mkldnn::memory::format::nc : mkldnn::memory::format::nchw;
27+ if (is_2d) {
28+ x_dims = {(int )batchsize, (int )channels};
29+ y_dims = {(int )batchsize, (int )channels};
30+ } else {
31+ x_dims = {(int )batchsize, (int )channels, (int )height, (int )width};
32+ y_dims = {(int )batchsize, (int )channels, (int )height, (int )width};
33+ }
3334
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);
35+ auto eng = *input.device ()->context (0 )->engine ;
36+ x_md = new mkldnn::memory::desc (x_dims, dtype, data_memory_format);
37+ dx_md = new mkldnn::memory::desc (x_dims, dtype, data_memory_format);
38+ bn_fwd_d = new mkldnn::batch_normalization_forward::desc (mkldnn::forward_training, *x_md, epsilon,
39+ mkldnn::use_scale_shift);
40+ bn_fwd_pd = new mkldnn::batch_normalization_forward::primitive_desc (*bn_fwd_d, eng);
41+ }
4042#endif // USE_MKLDNN
4143
4244};
4345
4446
45- BatchNormHandle::~BatchNormHandle () {
47+ BatchNormHandle::~BatchNormHandle () {
4648#ifdef USE_MKLDNN
49+ if (x_md != nullptr ) {
4750 delete (x_md);
4851 delete (dx_md);
4952 delete (bn_fwd_d);
5053 delete (bn_fwd_pd);
51- #endif // USE_MKLDNN
5254 }
55+ #endif // USE_MKLDNN
56+ }
5357
5458#ifdef USE_MKLDNN
5559
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-
60+ Tensor CpuBatchNormForwardInference (const BatchNormHandle &bnh, const Tensor& x, const Tensor& bnScale, const Tensor& bnBias,
61+ Tensor& running_mean, Tensor& running_var) {
6362
64- Tensor w = get_bn_weight_from (bnScale, bnBias);
63+ CHECK_EQ (x.device ()->lang (), kCpp );
64+ Tensor y;
65+ y.ResetLike (x);
6566
66- y.device ()->Exec ([&y, &x, &running_mean, &running_var, &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 ());
7267
73- // indicates using scale&bias and running mean&var
74- auto flags = use_scale_shift | use_global_stats;
75- auto bn_fwd_d = batch_normalization_forward::desc (forward_inference, *bnh.x_md , bnh.epsilon , flags);
76- auto bn_fwd_pd = batch_normalization_forward::primitive_desc (bn_fwd_d, eng);
77-
78- auto m_mem = memory (bn_fwd_pd.mean_primitive_desc (), running_mean.block ()->mutable_data ());
79- auto v_mem = memory (bn_fwd_pd.variance_primitive_desc (), running_var.block ()->mutable_data ());
80- auto w_mem = memory (bn_fwd_pd.weights_primitive_desc (), w.block ()->mutable_data ());
81-
82- // inputs require explicitly be indicated by casting according to
83- // https://intel.github.io/mkl-dnn/structmkldnn_1_1batch__normalization__forward.html
84- auto bn = batch_normalization_forward (bn_fwd_pd, x_mem, (const primitive::at)m_mem, (const primitive::at)v_mem, w_mem, y_mem);
68+ Tensor w = get_bn_weight_from (bnScale, bnBias);
8569
86- stream (stream::kind::eager).submit ({bn}).wait ();
87- }
88- catch (mkldnn::error &e) {
89- InitLogging (" " );
90- LOG (FATAL ) << " MKLDNN Batch Norm " << " Status: " << e.status << " Message: " << e.message ;
91- }
70+ y.device ()->Exec ([&y, &x, &running_mean, &running_var, &w, &bnh](Context * ctx) {
71+ try {
72+ auto eng = *ctx->engine ;
73+ using namespace mkldnn ;
74+ auto x_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng}, x.block ()->mutable_data ());
75+ auto y_mem = memory ({{{bnh.y_dims }, bnh.dtype , bnh.data_memory_format }, eng}, y.block ()->mutable_data ());
76+
77+ // indicates using scale&bias and running mean&var
78+ auto flags = use_scale_shift | use_global_stats;
79+ auto bn_fwd_d = batch_normalization_forward::desc (forward_inference, *bnh.x_md , bnh.epsilon , flags);
80+ auto bn_fwd_pd = batch_normalization_forward::primitive_desc (bn_fwd_d, eng);
81+
82+ auto m_mem = memory (bn_fwd_pd.mean_primitive_desc (), running_mean.block ()->mutable_data ());
83+ auto v_mem = memory (bn_fwd_pd.variance_primitive_desc (), running_var.block ()->mutable_data ());
84+ auto w_mem = memory (bn_fwd_pd.weights_primitive_desc (), w.block ()->mutable_data ());
85+
86+ // inputs require explicitly be indicated by casting according to
87+ // https://intel.github.io/mkl-dnn/structmkldnn_1_1batch__normalization__forward.html
88+ auto bn = batch_normalization_forward (bn_fwd_pd, x_mem, (const primitive::at)m_mem, (const primitive::at)v_mem, w_mem, y_mem);
89+
90+ stream (stream::kind::eager).submit ({bn}).wait ();
91+ } catch (mkldnn::error &e) {
92+ InitLogging (" " );
93+ LOG (FATAL ) << " MKLDNN Batch Norm " << " Status: " << e.status << " Message: " << e.message ;
94+ }
9295
93- }, {y.block (), x.block (), w.block ()}, {y.block ()});
96+ }, {y.block (), x.block (), w.block ()}, {y.block ()});
9497
95- return y;
98+ return y;
9699
97- }
100+ }
98101
99- const std::vector<Tensor>
100- CpuBatchNormForwardTraining (const BatchNormHandle &bnh, const Tensor &x, const Tensor &bnScale, const Tensor &bnBias,
101- Tensor &running_mean, Tensor &running_var) {
102+ const std::vector<Tensor>
103+ CpuBatchNormForwardTraining (const BatchNormHandle &bnh, const Tensor &x, const Tensor &bnScale, const Tensor &bnBias,
104+ Tensor &running_mean, Tensor &running_var) {
102105
103- Tensor y;
104- y.ResetLike (x);
106+ Tensor y;
107+ y.ResetLike (x);
105108
106- // mean and var for local batch
107- Tensor mean;
108- mean.ResetLike (running_mean);
109- Tensor var;
110- var.ResetLike (running_var);
109+ // mean and var for local batch
110+ Tensor mean;
111+ mean.ResetLike (running_mean);
112+ Tensor var;
113+ var.ResetLike (running_var);
111114
112- // combine scale and bias to construct weight tensor in required format for backward
113- Tensor w = get_bn_weight_from (bnScale, bnBias);
115+ // combine scale and bias to construct weight tensor in required format for backward
116+ Tensor w = get_bn_weight_from (bnScale, bnBias);
114117
115- y.device ()->Exec ([&x, &y, &mean, &var, &w, &bnh](Context *ctx) {
116- try {
117- auto eng = *ctx->engine ;
118- using namespace mkldnn ;
118+ y.device ()->Exec ([&x, &y, &mean, &var, &w, &bnh](Context * ctx) {
119+ try {
120+ auto eng = *ctx->engine ;
121+ using namespace mkldnn ;
119122
120- auto x_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng},
121- x.block ()->mutable_data ());
122- auto y_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng},
123- y.block ()->mutable_data ());
124- auto m_mem = memory (bnh.bn_fwd_pd ->mean_primitive_desc (), mean.block ()->mutable_data ());
123+ auto x_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng},
124+ x.block ()->mutable_data ());
125+ auto y_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng},
126+ y.block ()->mutable_data ());
127+ auto m_mem = memory (bnh.bn_fwd_pd ->mean_primitive_desc (), mean.block ()->mutable_data ());
125128
126- auto v_mem = memory (bnh.bn_fwd_pd ->variance_primitive_desc (), var.block ()->mutable_data ());
129+ auto v_mem = memory (bnh.bn_fwd_pd ->variance_primitive_desc (), var.block ()->mutable_data ());
127130
128- auto w_mem = memory (bnh.bn_fwd_pd ->weights_primitive_desc (),w.block ()->mutable_data ());
131+ auto w_mem = memory (bnh.bn_fwd_pd ->weights_primitive_desc (), w.block ()->mutable_data ());
129132
130- auto bn_fwd = batch_normalization_forward (*bnh.bn_fwd_pd , x_mem, w_mem, y_mem, m_mem, v_mem);
133+ auto bn_fwd = batch_normalization_forward (*bnh.bn_fwd_pd , x_mem, w_mem, y_mem, m_mem, v_mem);
131134
132- stream (stream::kind::eager).submit ({bn_fwd}).wait ();
133- }
134- catch (mkldnn::error &e) {
135- singa::InitLogging (" " );
136- LOG (FATAL ) << " MKLDNN Batch Norm Backward" << " Status: " << e.status << " Message: " << e.message ;
137- }
138- }, {x.block (), w.block ()}, {y.block (), mean.block (), var.block ()});
135+ stream (stream::kind::eager).submit ({bn_fwd}).wait ();
136+ } catch (mkldnn::error &e) {
137+ singa::InitLogging (" " );
138+ LOG (FATAL ) << " MKLDNN Batch Norm Backward" << " Status: " << e.status << " Message: " << e.message ;
139+ }
140+ }, {x.block (), w.block ()}, {y.block (), mean.block (), var.block ()});
139141
140142
141- // local implemented running mean as mkldnn does not support it yet:
142- // https://github.com/intel/mkl-dnn/issues/371
143- running_mean = running_mean* bnh.factor + mean*( 1 - bnh.factor );
144- running_var = running_var* bnh.factor + var*( 1 - bnh.factor );
143+ // local implemented running mean as mkldnn does not support it yet:
144+ // https://github.com/intel/mkl-dnn/issues/371
145+ running_mean = running_mean * bnh.factor + mean * ( 1 - bnh.factor );
146+ running_var = running_var * bnh.factor + var * ( 1 - bnh.factor );
145147
146148
147- return {y, running_mean, running_var};
149+ return {y, running_mean, running_var};
148150
149- }
151+ }
150152
151153const std::vector<Tensor> CpuBatchNormBackwardx (const BatchNormHandle &bnh,
152- const Tensor &y, const Tensor &dy,
153- const Tensor &x,
154- const Tensor &bnScale, const Tensor &bnBias,
155- const Tensor &mean, const Tensor &var){
154+ const Tensor &y, const Tensor &dy,
155+ const Tensor &x,
156+ const Tensor &bnScale, const Tensor &bnBias,
157+ const Tensor &mean, const Tensor &var) {
156158 Tensor dx;
157159 dx.ResetLike (dy);
158160
@@ -161,7 +163,7 @@ const std::vector<Tensor> CpuBatchNormBackwardx(const BatchNormHandle &bnh,
161163
162164 Tensor dw (Shape{bnScale.Size (), 2 });
163165
164- dx.device ()->Exec ([&dw, &x, &dx, &y, &dy, &w, &mean, &var, &bnh](Context *ctx) {
166+ dx.device ()->Exec ([&dw, &x, &dx, &y, &dy, &w, &mean, &var, &bnh](Context * ctx) {
165167
166168 try {
167169 auto eng = *ctx->engine ;
@@ -186,8 +188,7 @@ const std::vector<Tensor> CpuBatchNormBackwardx(const BatchNormHandle &bnh,
186188 auto bn_bwd = batch_normalization_backward (bn_bwd_pd, x_mem, m_mem, v_mem, dy_mem, w_mem, dx_mem, dw_mem);
187189
188190 stream (stream::kind::eager).submit ({bn_bwd}).wait ();
189- }
190- catch (mkldnn::error &e) {
191+ } catch (mkldnn::error &e) {
191192 singa::InitLogging (" " );
192193 LOG (FATAL ) << " MKLDNN Batch Norm Backward" << " Status: " << e.status << " Message: " << e.message ;
193194 }
@@ -206,7 +207,7 @@ const std::vector<Tensor> CpuBatchNormBackwardx(const BatchNormHandle &bnh,
206207 CHECK (dbnBias.shape ()[0 ] == bnBias.shape ()[0 ]) << " dbnBias shape not match bnBias" ;
207208
208209 return {dx, dbnScale, dbnBias};
209- }
210+ }
210211
211212
212213#endif // USE_MKLDNN
@@ -231,8 +232,8 @@ CudnnBatchNormHandle::CudnnBatchNormHandle(const float momentum,
231232};
232233
233234const std::vector<Tensor> GpuBatchNormForwardTraining (const CudnnBatchNormHandle &cbnh,
234- const Tensor& x, const Tensor& bnScale, const Tensor& bnBias,
235- Tensor& running_mean, Tensor& running_var) {
235+ const Tensor& x, const Tensor& bnScale, const Tensor& bnBias,
236+ Tensor& running_mean, Tensor& running_var) {
236237 CHECK_EQ (x.device ()->lang (), kCuda );
237238 CHECK_EQ (bnScale.device ()->lang (), kCuda );
238239 CHECK_EQ (bnBias.device ()->lang (), kCuda );
0 commit comments