| 14 | using namespace megdnn::cuda::non_zero; |
| 15 | |
| 16 | WorkspaceBundle NonZeroImpl::make_bundle(const TensorLayout& data) { |
| 17 | size_t nr_item = data.total_nr_elems(); |
| 18 | cuda_check(cudaSetDevice(concrete_handle(handle())->device_id())); |
| 19 | auto gen_idx_wk_size = cuda::cond_take::gen_idx_get_workspace_size(nr_item); |
| 20 | SmallVector<size_t> sizes_in_bytes; |
| 21 | sizes_in_bytes.push_back((nr_item + 1) * sizeof(megdnn::cuda::cond_take::IdxType)); |
| 22 | sizes_in_bytes.push_back(gen_idx_wk_size); |
| 23 | // the two ele is the shape of arr and the reverse multiply arr of the shape |
| 24 | sizes_in_bytes.push_back(sizeof(TensorLayout::shape)); |
| 25 | sizes_in_bytes.push_back(sizeof(TensorLayout::shape)); |
| 26 | return {nullptr, sizes_in_bytes, handle()->alignment_requirement()}; |
| 27 | } |
| 28 | |
| 29 | size_t NonZeroImpl::get_workspace_in_bytes(const TensorLayout& data) { |
| 30 | return make_bundle(data).total_size_in_bytes(); |
nothing calls this directly
no test coverage detected