@@ -41,198 +41,12 @@ BatchNormHandle::BatchNormHandle(const float momentum, const Tensor& input) {
4141 }
4242
4343
44- #ifdef USE_MKLDNN
45- if (input.device ()->lang () == kCpp ) {
46- dtype = GetMKLDNNDataType (input.data_type ());
47- epsilon = 1e-5f ;
48- data_memory_format = is_2d ? mkldnn::memory::format::nc : mkldnn::memory::format::nchw;
49- if (is_2d) {
50- x_dims = {(int )batchsize, (int )channels};
51- y_dims = {(int )batchsize, (int )channels};
52- } else {
53- x_dims = {(int )batchsize, (int )channels, (int )height, (int )width};
54- y_dims = {(int )batchsize, (int )channels, (int )height, (int )width};
55- }
56-
57- auto eng = *input.device ()->context (0 )->engine ;
58- x_md = new mkldnn::memory::desc (x_dims, dtype, data_memory_format);
59- dx_md = new mkldnn::memory::desc (x_dims, dtype, data_memory_format);
60- bn_fwd_d = new mkldnn::batch_normalization_forward::desc (mkldnn::forward_training, *x_md, epsilon,
61- mkldnn::use_scale_shift);
62- bn_fwd_pd = new mkldnn::batch_normalization_forward::primitive_desc (*bn_fwd_d, eng);
63- }
64- #endif // USE_MKLDNN
65-
6644};
6745
6846
6947BatchNormHandle::~BatchNormHandle () {
70- #ifdef USE_MKLDNN
71- if (x_md != nullptr ) {
72- delete (x_md);
73- delete (dx_md);
74- delete (bn_fwd_d);
75- delete (bn_fwd_pd);
76- }
77- #endif // USE_MKLDNN
78- }
79-
80- #ifdef USE_MKLDNN
81-
82- Tensor CpuBatchNormForwardInference (const BatchNormHandle &bnh, const Tensor& x, const Tensor& bnScale, const Tensor& bnBias,
83- Tensor& running_mean, Tensor& running_var) {
84-
85- CHECK_EQ (x.device ()->lang (), kCpp );
86- Tensor y;
87- y.ResetLike (x);
88-
89-
90- Tensor w = get_bn_weight_from (bnScale, bnBias);
91-
92- y.device ()->Exec ([&y, &x, &running_mean, &running_var, &w, &bnh](Context * ctx) {
93- try {
94- auto eng = *ctx->engine ;
95- using namespace mkldnn ;
96- auto x_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng}, x.block ()->mutable_data ());
97- auto y_mem = memory ({{{bnh.y_dims }, bnh.dtype , bnh.data_memory_format }, eng}, y.block ()->mutable_data ());
98-
99- // indicates using scale&bias and running mean&var
100- auto flags = use_scale_shift | use_global_stats;
101- auto bn_fwd_d = batch_normalization_forward::desc (forward_inference, *bnh.x_md , bnh.epsilon , flags);
102- auto bn_fwd_pd = batch_normalization_forward::primitive_desc (bn_fwd_d, eng);
103-
104- auto m_mem = memory (bn_fwd_pd.mean_primitive_desc (), running_mean.block ()->mutable_data ());
105- auto v_mem = memory (bn_fwd_pd.variance_primitive_desc (), running_var.block ()->mutable_data ());
106- auto w_mem = memory (bn_fwd_pd.weights_primitive_desc (), w.block ()->mutable_data ());
107-
108- // inputs require explicitly be indicated by casting according to
109- // https://intel.github.io/mkl-dnn/structmkldnn_1_1batch__normalization__forward.html
110- auto bn = batch_normalization_forward (bn_fwd_pd, x_mem, (const primitive::at)m_mem, (const primitive::at)v_mem, w_mem, y_mem);
111-
112- stream (stream::kind::eager).submit ({bn}).wait ();
113- } catch (mkldnn::error &e) {
114- InitLogging (" " );
115- LOG (FATAL ) << " MKLDNN Batch Norm " << " Status: " << e.status << " Message: " << e.message ;
116- }
117-
118- }, {y.block (), x.block (), w.block ()}, {y.block ()});
119-
120- return y;
121-
122- }
123-
124- const std::vector<Tensor>
125- CpuBatchNormForwardTraining (const BatchNormHandle &bnh, const Tensor &x, const Tensor &bnScale, const Tensor &bnBias,
126- Tensor &running_mean, Tensor &running_var) {
127-
128- Tensor y;
129- y.ResetLike (x);
130-
131- // mean and var for local batch
132- Tensor mean;
133- mean.ResetLike (running_mean);
134- Tensor var;
135- var.ResetLike (running_var);
136-
137- // combine scale and bias to construct weight tensor in required format for backward
138- Tensor w = get_bn_weight_from (bnScale, bnBias);
139-
140- y.device ()->Exec ([&x, &y, &mean, &var, &w, &bnh](Context * ctx) {
141- try {
142- auto eng = *ctx->engine ;
143- using namespace mkldnn ;
144-
145- auto x_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng},
146- x.block ()->mutable_data ());
147- auto y_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng},
148- y.block ()->mutable_data ());
149- auto m_mem = memory (bnh.bn_fwd_pd ->mean_primitive_desc (), mean.block ()->mutable_data ());
150-
151- auto v_mem = memory (bnh.bn_fwd_pd ->variance_primitive_desc (), var.block ()->mutable_data ());
152-
153- auto w_mem = memory (bnh.bn_fwd_pd ->weights_primitive_desc (), w.block ()->mutable_data ());
154-
155- auto bn_fwd = batch_normalization_forward (*bnh.bn_fwd_pd , x_mem, w_mem, y_mem, m_mem, v_mem);
156-
157- stream (stream::kind::eager).submit ({bn_fwd}).wait ();
158- } catch (mkldnn::error &e) {
159- singa::InitLogging (" " );
160- LOG (FATAL ) << " MKLDNN Batch Norm Backward" << " Status: " << e.status << " Message: " << e.message ;
161- }
162- }, {x.block (), w.block ()}, {y.block (), mean.block (), var.block ()});
163-
164-
165- // local implemented running mean as mkldnn does not support it yet:
166- // https://github.com/intel/mkl-dnn/issues/371
167- running_mean = running_mean * bnh.factor + mean * (1 - bnh.factor );
168- running_var = running_var * bnh.factor + var * (1 - bnh.factor );
169-
170-
171- return {y, running_mean, running_var};
172-
17348}
17449
175- const std::vector<Tensor> CpuBatchNormBackwardx (const BatchNormHandle &bnh,
176- const Tensor &y, const Tensor &dy,
177- const Tensor &x,
178- const Tensor &bnScale, const Tensor &bnBias,
179- const Tensor &mean, const Tensor &var) {
180- Tensor dx;
181- dx.ResetLike (dy);
182-
183- // combine scale and bias to construct weight tensor in required format for backward
184- Tensor w = get_bn_weight_from (bnScale, bnBias);
185-
186- Tensor dw (Shape{bnScale.Size (), 2 });
187-
188- dx.device ()->Exec ([&dw, &x, &dx, &y, &dy, &w, &mean, &var, &bnh](Context * ctx) {
189-
190- try {
191- auto eng = *ctx->engine ;
192- using namespace mkldnn ;
193-
194- auto x_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng}, x.block ()->mutable_data ());
195- auto dx_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng}, dx.block ()->mutable_data ());
196- auto y_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng}, y.block ()->mutable_data ());
197- auto dy_mem = memory ({{{bnh.x_dims }, bnh.dtype , bnh.data_memory_format }, eng}, dy.block ()->mutable_data ());
198-
199- auto m_mem = memory (bnh.bn_fwd_pd ->mean_primitive_desc (), mean.block ()->mutable_data ());
200- auto v_mem = memory (bnh.bn_fwd_pd ->variance_primitive_desc (), var.block ()->mutable_data ());
201- auto w_mem = memory (bnh.bn_fwd_pd ->weights_primitive_desc (), w.block ()->mutable_data ());
202-
203-
204- auto bn_bwd_d = batch_normalization_backward::desc (backward, *bnh.dx_md , *bnh.x_md , bnh.epsilon , use_scale_shift);
205- auto bn_bwd_pd = batch_normalization_backward::primitive_desc (bn_bwd_d, eng, *bnh.bn_fwd_pd );
206-
207-
208- auto dw_mem = memory (bn_bwd_pd.diff_weights_primitive_desc (), dw.block ()->mutable_data ());
209-
210- auto bn_bwd = batch_normalization_backward (bn_bwd_pd, x_mem, m_mem, v_mem, dy_mem, w_mem, dx_mem, dw_mem);
211-
212- stream (stream::kind::eager).submit ({bn_bwd}).wait ();
213- } catch (mkldnn::error &e) {
214- singa::InitLogging (" " );
215- LOG (FATAL ) << " MKLDNN Batch Norm Backward" << " Status: " << e.status << " Message: " << e.message ;
216- }
217-
218- }, {x.block (), dy.block (), mean.block (), var.block ()},
219- {dx.block (), dw.block ()});
220-
221- singa::Tensor dbnScale (bnScale.shape ());
222- CopyDataToFrom (&dbnScale, dw, bnScale.Size (), 0 , 0 );
223- singa::Tensor dbnBias (bnBias.shape ());
224- CopyDataToFrom (&dbnBias, dw, bnBias.Size (), 0 , bnScale.Size ());
225-
226- CHECK (dbnScale.nDim () == bnScale.nDim ()) << " dbnScale ndim not match bnScale" ;
227- CHECK (dbnBias.nDim () == bnBias.nDim ()) << " dbnScale ndim not match bnScale" ;
228- CHECK (dbnScale.shape ()[0 ] == bnScale.shape ()[0 ]) << " dbnScale shape not match bnScale" ;
229- CHECK (dbnBias.shape ()[0 ] == bnBias.shape ()[0 ]) << " dbnBias shape not match bnBias" ;
230-
231- return {dx, dbnScale, dbnBias};
232- }
233-
234-
235- #endif // USE_MKLDNN
23650
23751#ifdef USE_CUDNN
23852CudnnBatchNormHandle::CudnnBatchNormHandle (const float momentum,
0 commit comments