namespace
| 34 | |
| 35 | } // namespace |
| 36 | WorkspaceBundle CondTakeImpl::make_bundle( |
| 37 | const TensorLayout& data, const TensorLayout& mask) { |
| 38 | auto handle = concrete_handle(this->handle()); |
| 39 | size_t elem_size = mask.total_nr_elems(); |
| 40 | size_t assign_sub_workspace = 0, logicop_workspace = 0, num_true_workspace = 0, |
| 41 | where_workspace = 0; |
| 42 | CnnlTensorDescriptor num_true_desc, mask_desc, one_elem_desc; |
| 43 | ShapeInfo one_elem_shape = {1}; |
| 44 | num_true_desc.set(1, one_elem_shape, CNNL_DTYPE_INT32, CNNL_LAYOUT_ARRAY); |
| 45 | mask_desc.set(mask.ndim, mask.shape, CNNL_DTYPE_FLOAT, CNNL_LAYOUT_ARRAY); |
| 46 | one_elem_desc.set(1, one_elem_shape, CNNL_DTYPE_FLOAT, CNNL_LAYOUT_ARRAY); |
| 47 | cnnl_check(cnnlGetAssignSubWorkspaceSize( |
| 48 | handle->cnnl_handle(), one_elem_desc.desc(), mask_desc.desc(), |
| 49 | &assign_sub_workspace)); |
| 50 | cnnl_check(cnnlGetLogicOpWorkspaceSize( |
| 51 | handle->cnnl_handle(), mask_desc.desc(), one_elem_desc.desc(), |
| 52 | mask_desc.desc(), &logicop_workspace)); |
| 53 | cnnl_check(cnnlGetNumTrueWorkspaceSize_v2( |
| 54 | handle->cnnl_handle(), num_true_desc.desc(), &num_true_workspace)); |
| 55 | cnnl_check(cnnlGetWhereWorkspaceSize( |
| 56 | handle->cnnl_handle(), num_true_desc.desc(), &where_workspace)); |
| 57 | return {nullptr, |
| 58 | {/*float_mask*/ elem_size * sizeof(float), |
| 59 | /*assign_sub, logic, num_true*/ 1 * sizeof(float), assign_sub_workspace, |
| 60 | logicop_workspace, num_true_workspace, where_workspace}, |
| 61 | handle->alignment_requirement()}; |
| 62 | } |
| 63 | |
| 64 | size_t CondTakeImpl::get_workspace_in_bytes( |
| 65 | const TensorLayout& data, const TensorLayout& mask) { |
nothing calls this directly
no test coverage detected