MCPcopy Create free account
hub / github.com/OAID/Tengine / BatchNormOps

Class BatchNormOps

executor/operator/common/batchnorm.cpp:39–166  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

37namespace BatchNormImpl {
38
39struct BatchNormOps : public NodeOps
40{
41 bool Prerun(Node* node)
42 {
43 const Tensor* input_tensor = node->GetInputTensor(0);
44 const TShape& shape = input_tensor->GetShape();
45
46 const std::vector<int> dims = shape.GetDim();
47
48 int channel_num = dims[1];
49
50 float* scale_mean = ( float* )mem_alloc(channel_num * sizeof(float));
51 float* scale_var_inv = ( float* )mem_alloc(channel_num * sizeof(float));
52
53 const Tensor* mean_tensor = node->GetInputTensor(3);
54 const Tensor* var_tensor = node->GetInputTensor(4);
55 const float* mean = ( const float* )get_tensor_mem(mean_tensor);
56 const float* var = ( const float* )get_tensor_mem(var_tensor);
57
58 BatchNorm* bn_op = dynamic_cast<BatchNorm*>(node->GetOp());
59 BatchNormParam* param = bn_op->GetParam();
60
61 float rescale_factor;
62 float eps = param->eps;
63
64 rescale_factor = param->rescale_factor ? 1 / param->rescale_factor : 0;
65 for(int c = 0; c < channel_num; c++)
66 {
67 scale_var_inv[c] = 1.f / sqrt(var[c] * rescale_factor + eps);
68 scale_mean[c] = -mean[c] * rescale_factor * scale_var_inv[c];
69 }
70
71 node->SetAttr("scale_mean", scale_mean);
72 node->SetAttr("scale_var_inv", scale_var_inv);
73
74 return true;
75 }
76
77 bool Run(Node* node)
78 {
79 const Tensor* input_tensor = node->GetInputTensor(0);
80 Tensor* output_tensor = node->GetOutputTensor(0);
81 const TShape& shape = input_tensor->GetShape();
82 const std::vector<int> dims = shape.GetDim();
83
84 int batch_number = dims[0];
85 int channel_num = dims[1];
86 int channel_size = dims[2] * dims[3];
87 int img_size = channel_num * channel_size;
88
89 BatchNorm* bn_op = dynamic_cast<BatchNorm*>(node->GetOp());
90 BatchNormParam* param = bn_op->GetParam();
91
92 const float* input = ( const float* )get_tensor_mem(input_tensor);
93 float* output = ( float* )get_tensor_mem(output_tensor);
94
95 if(param->caffe_flavor)
96 {

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected