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

Method exec

dnn/src/cambricon/group_norm/opr_impl.cpp:35–61  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

33}
34
35void GroupNormForwardImpl::exec(
36 _megdnn_tensor_in data, _megdnn_tensor_in weight, _megdnn_tensor_in bias,
37 _megdnn_tensor_out dst, _megdnn_tensor_out mean, _megdnn_tensor_out rstd,
38 _megdnn_workspace workspace) {
39 check_exec(
40 data.layout, weight.layout, bias.layout, dst.layout, mean.layout,
41 rstd.layout, workspace.size);
42 megdnn_assert(data.layout.dtype == dtype::Float32());
43 auto handle = concrete_handle(this->handle());
44 auto bundle = get_workspace_bundle(data.layout, workspace.raw_ptr);
45 size_t channel = data.layout.shape[1];
46 void *weight_ptr = nullptr, *bias_ptr = nullptr;
47 CnnlTensorDescriptor x_desc, scale_bias_desc, y_desc, mean_rstd_desc;
48 x_desc.set(data.layout, CNNL_LAYOUT_NCHW);
49 if (param().affine) {
50 scale_bias_desc.set(weight.layout, CNNL_LAYOUT_ARRAY);
51 weight_ptr = weight.raw_ptr();
52 bias_ptr = bias.raw_ptr();
53 }
54 y_desc.set(dst.layout, CNNL_LAYOUT_NCHW);
55 mean_rstd_desc.set(mean.layout, CNNL_LAYOUT_ARRAY);
56 cnnl_check(cnnlGroupNormForward_v3(
57 handle->cnnl_handle(), param().eps, static_cast<int>(param().group),
58 x_desc.desc(), data.raw_ptr(), scale_bias_desc.desc(), weight_ptr, bias_ptr,
59 bundle.get(0), bundle.get_size(0), y_desc.desc(), dst.raw_ptr(),
60 mean_rstd_desc.desc(), mean.raw_ptr(), rstd.raw_ptr()));
61}
62
63size_t GroupNormBackwardImpl::get_workspace_in_bytes(
64 const TensorLayout&, const TensorLayout& data, const TensorLayout&,

Callers

nothing calls this directly

Calls 12

cnnl_handleMethod · 0.80
concrete_handleFunction · 0.50
get_workspace_bundleFunction · 0.50
paramFunction · 0.50
handleMethod · 0.45
setMethod · 0.45
raw_ptrMethod · 0.45
descMethod · 0.45
getMethod · 0.45
get_sizeMethod · 0.45
enumvMethod · 0.45
sizeMethod · 0.45

Tested by

no test coverage detected