| 38 | namespace BatchNormImpl64 { |
| 39 | |
| 40 | void batchnorm_kernel(int i, int id, void* data, const float* input, float* output, float* scale_mean, float* scale_var, int channel_size) |
| 41 | { |
| 42 | int step = ((int*)data)[0]; |
| 43 | for(int c = 0; c < step; c++) |
| 44 | { |
| 45 | int cur_c = id * step + c; |
| 46 | float s_mean = scale_mean[cur_c]; |
| 47 | float s_var = scale_var[cur_c]; |
| 48 | float32x4_t _mean = vdupq_n_f32(s_mean); |
| 49 | float32x4_t _var = vdupq_n_f32(s_var); |
| 50 | int offset = cur_c * channel_size; |
| 51 | const float* input_ptr = input + offset; |
| 52 | float* output_ptr = output + offset; |
| 53 | |
| 54 | // output[offset]= input[offset]*scale_var_inv[c] - scale_mean[c]; |
| 55 | for(int l = 0; l < (channel_size & -4); l += 4) |
| 56 | { |
| 57 | float32x4_t _input = vld1q_f32(input_ptr); |
| 58 | vst1q_f32(output_ptr, vmlaq_f32(_mean, _input, _var)); |
| 59 | input_ptr += 4; |
| 60 | output_ptr += 4; |
| 61 | } |
| 62 | for(int l = channel_size & ~3; l < channel_size; l++) |
| 63 | { |
| 64 | *output_ptr = (*input_ptr) * s_var + s_mean; |
| 65 | input_ptr++; |
| 66 | output_ptr++; |
| 67 | } |
| 68 | } |
| 69 | } |
| 70 | |
| 71 | struct BNOps : public NodeOps |
| 72 | { |