| 33 | } |
| 34 | |
| 35 | void 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 | |
| 63 | size_t GroupNormBackwardImpl::get_workspace_in_bytes( |
| 64 | const TensorLayout&, const TensorLayout& data, const TensorLayout&, |
nothing calls this directly
no test coverage detected