| 115 | |
| 116 | template <typename T0, typename T1 = T0> |
| 117 | void bn_backward_exec( |
| 118 | _megdnn_tensor_in x_in, _megdnn_tensor_in dy_in, |
| 119 | _megdnn_tensor_in saved_batch_mean, _megdnn_tensor_in saved_batch_inv_variance, |
| 120 | _megdnn_tensor_in bn_scale, _megdnn_tensor_out d_bn_scale, |
| 121 | _megdnn_tensor_out d_bn_bias, _megdnn_tensor_out dx_out, |
| 122 | const WorkspaceBundle& bundle) { |
| 123 | size_t src_shape[4], dim_offset[4], param_pos = 0, src_pos = 0, batch_size = 1, |
| 124 | param_size = bn_scale.layout.total_nr_elems(); |
| 125 | |
| 126 | float denominator = 1.f; |
| 127 | |
| 128 | T0 *x = x_in.ptr<T0>(), *dx = dx_out.ptr<T0>(), *dy = dy_in.ptr<T0>(); |
| 129 | T1 *gamma = bn_scale.ptr<T1>(), *mu = saved_batch_mean.ptr<T1>(), |
| 130 | *ivar = saved_batch_inv_variance.ptr<T1>(), *dgamma = d_bn_scale.ptr<T1>(), |
| 131 | *dbeta = d_bn_bias.ptr<T1>(); |
| 132 | |
| 133 | rep(i, dy_in.layout.ndim) { |
| 134 | src_shape[i] = dy_in.layout.shape[i]; |
| 135 | if (bn_scale.layout.shape[i] == 1) { |
| 136 | dim_offset[i] = 0; |
| 137 | batch_size *= src_shape[i]; |
| 138 | } else { |
| 139 | dim_offset[i] = 1; |
| 140 | } |
| 141 | } |
| 142 | |
| 143 | int curr_stride = 0; |
| 144 | for (int i = 3; i >= 0; --i) { |
| 145 | if (dim_offset[i] != 0) { |
| 146 | if (curr_stride == 0) { |
| 147 | dim_offset[i] = 1; |
| 148 | curr_stride = src_shape[i]; |
| 149 | } else { |
| 150 | dim_offset[i] = curr_stride; |
| 151 | curr_stride *= src_shape[i]; |
| 152 | } |
| 153 | } |
| 154 | } |
| 155 | denominator = 1.0 / batch_size; |
| 156 | |
| 157 | // step1. dbeta, dgamma |
| 158 | memset(dbeta, 0, param_size * sizeof(T1)); |
| 159 | memset(dgamma, 0, param_size * sizeof(T1)); |
| 160 | rep_4d(src_shape, dim_offset) float xhat = |
| 161 | (x[src_pos] - mu[param_pos]) * ivar[param_pos]; |
| 162 | dbeta[param_pos] += dy[src_pos]; |
| 163 | dgamma[param_pos] += dy[src_pos] * xhat; |
| 164 | rep_4d_end |
| 165 | |
| 166 | // step2. dxhat = dy * gamma |
| 167 | float* dxhat = static_cast<float*>(bundle.get(0)); |
| 168 | rep_4d(src_shape, dim_offset) dxhat[src_pos] = dy[src_pos] * gamma[param_pos]; |
| 169 | rep_4d_end |
| 170 | |
| 171 | // step3. dvar = sigma[dxhat * xmu] * [-1/2 * (ivar)^(3/2)] |
| 172 | // dmu = sigma[dxhat] * (-ivar) |
| 173 | float* dvar = static_cast<float*>(bundle.get(1)); |
| 174 | float* dmu = static_cast<float*>(bundle.get(2)); |
nothing calls this directly
no test coverage detected