| 23 | }; |
| 24 | |
| 25 | WorkspaceBundle ArgmaxForwardImpl::make_bundle( |
| 26 | const TensorLayout& src, const TensorLayout& dst) { |
| 27 | auto handle = concrete_handle(this->handle()); |
| 28 | size_t topk_ws = 0; |
| 29 | ArgmxxCnnlDescs descs(src, dst, param().axis); |
| 30 | cnnl_check(cnnlGetTopKTensorWorkspaceSize( |
| 31 | /* handle */ handle->cnnl_handle(), |
| 32 | /* input_desc */ descs.src_desc.desc(), |
| 33 | /* k */ descs.k, |
| 34 | /* dim */ param().axis, |
| 35 | /* largest */ true, |
| 36 | /* output_desc */ descs.out_value_desc.desc(), |
| 37 | /* index_desc */ descs.out_indices_desc.desc(), |
| 38 | /* workspace_size */ &topk_ws)); |
| 39 | size_t value_ws = dst.span().dist_elem() * src.dtype.size(); |
| 40 | size_t src_ws = 0; |
| 41 | if (!src.is_contiguous()) { |
| 42 | src_ws = src.access_bytes(); |
| 43 | } |
| 44 | return {nullptr, {topk_ws, value_ws, src_ws}, handle->alignment_requirement()}; |
| 45 | } |
| 46 | |
| 47 | size_t ArgmaxForwardImpl::get_workspace_in_bytes( |
| 48 | const TensorLayout& src, const TensorLayout& dst) { |
nothing calls this directly
no test coverage detected