MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / exec

Method exec

dnn/src/naive/batch_normalization/opr_impl.cpp:199–239  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

197}; // anonymous namespace
198
199void BNForwardImpl::exec(
200 _megdnn_tensor_in src, _megdnn_tensor_in bn_scale, _megdnn_tensor_in bn_bias,
201 _megdnn_tensor_inout mean, _megdnn_tensor_inout variance,
202 _megdnn_tensor_out batch_mean, _megdnn_tensor_out batch_inv_variance,
203 _megdnn_tensor_out, _megdnn_tensor_out dst, _megdnn_workspace workspace) {
204#if !MGE_BUILD_WITHOUT_NAIVE_EXEC
205 check_exec(
206 src.layout, bn_scale.layout, bn_bias.layout, mean.layout, variance.layout,
207 batch_mean.layout, batch_inv_variance.layout, dst.layout, workspace.size);
208
209 DNN_INC_FLOAT16(
210 if (src.layout.dtype == dtype::Float16() &&
211 bn_scale.layout.dtype == dtype::Float32()) {
212 MEGDNN_DISPATCH_CPU_KERN_OPR(({
213 using T0 = typename DTypeTrait<dtype::Float16>::ctype;
214 using T1 = typename DTypeTrait<dtype::Float32>::ctype;
215 bn_forward_exec<T0, T1>(
216 src, bn_scale, bn_bias, mean, variance, batch_mean,
217 batch_inv_variance, dst, m_param);
218 }));
219 } else) {
220 megdnn_assert(src.layout.dtype == bn_scale.layout.dtype);
221 switch (src.layout.dtype.enumv()) {
222#define cb(_dt) \
223 case DTypeTrait<_dt>::enumv: { \
224 using T = typename DTypeTrait<_dt>::ctype; \
225 MEGDNN_DISPATCH_CPU_KERN_OPR((bn_forward_exec<T>( \
226 src, bn_scale, bn_bias, mean, variance, batch_mean, \
227 batch_inv_variance, dst, m_param))); \
228 break; \
229 }
230 MEGDNN_FOREACH_COMPUTING_DTYPE_FLOAT(cb)
231#undef cb
232 default:
233 megdnn_assert_internal(0);
234 }
235 }
236#else
237 __builtin_trap();
238#endif
239}
240
241WorkspaceBundle BNBackwardImpl::get_workspace_bundle(
242 size_t x_size, size_t param_size, void* raw_ptr) {

Callers

nothing calls this directly

Calls 4

ifFunction · 0.50
get_workspace_bundleFunction · 0.50
enumvMethod · 0.45
total_nr_elemsMethod · 0.45

Tested by

no test coverage detected