MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / make_bundle

Method make_bundle

dnn/src/cambricon/cond_take/opr_impl.cpp:36–62  ·  view source on GitHub ↗

namespace

Source from the content-addressed store, hash-verified

34
35} // namespace
36WorkspaceBundle 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
64size_t CondTakeImpl::get_workspace_in_bytes(
65 const TensorLayout& data, const TensorLayout& mask) {

Callers

nothing calls this directly

Calls 7

cnnl_handleMethod · 0.80
concrete_handleFunction · 0.50
handleMethod · 0.45
total_nr_elemsMethod · 0.45
setMethod · 0.45
descMethod · 0.45
alignment_requirementMethod · 0.45

Tested by

no test coverage detected