| 24 | } |
| 25 | |
| 26 | WorkspaceBundle TopKImpl::make_bundle( |
| 27 | int k, const TensorLayout& data, const TensorLayout& values, |
| 28 | const TensorLayout& indices) { |
| 29 | auto handle = concrete_handle(this->handle()); |
| 30 | size_t topk_workspace = 0; |
| 31 | TopKCnnlDescs descs(data, values, param().mode); |
| 32 | bool largest = false; |
| 33 | if (k < 0) { |
| 34 | largest = true; |
| 35 | k = std::abs(k); |
| 36 | } |
| 37 | cnnl_check(cnnlGetTopKTensorWorkspaceSize( |
| 38 | /* handle */ handle->cnnl_handle(), |
| 39 | /* input_desc */ descs.data_desc.desc(), |
| 40 | /* k */ k, |
| 41 | /* dim */ descs.sort_dim, |
| 42 | /* largest */ largest, |
| 43 | /* output_desc */ descs.out_value_desc.desc(), |
| 44 | /* index_desc */ descs.out_indices_desc.desc(), |
| 45 | /* workspace_size */ &topk_workspace)); |
| 46 | size_t data_workspace = 0; |
| 47 | if (!data.is_contiguous()) { |
| 48 | data_workspace = data.access_bytes(); |
| 49 | } |
| 50 | return {nullptr, {data_workspace, topk_workspace}, handle->alignment_requirement()}; |
| 51 | } |
| 52 | |
| 53 | size_t TopKImpl::get_workspace_in_bytes( |
| 54 | int k, const TensorLayout& data, const TensorLayout& values, |
nothing calls this directly
no test coverage detected