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

Method exec

dnn/src/cuda/cond_take/opr_impl.cpp:28–76  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

26}
27
28CondTakeImpl::Output CondTakeImpl::exec(
29 _megdnn_tensor_in data, _megdnn_tensor_in mask, _megdnn_workspace workspace,
30 DynOutMallocPolicyCall malloc_policy) {
31 size_t size = check_exec_get_size(data.layout, mask.layout, workspace.size);
32 auto wk_bundle = make_bundle(size);
33 wk_bundle.set(workspace.raw_ptr);
34
35 auto idx_tmp = static_cast<IdxType*>(wk_bundle.get(0));
36
37 KParam kparam(param());
38 auto stream = cuda_stream(handle());
39 size_t out_size;
40 switch (mask.layout.dtype.enumv()) {
41#define cb(_dt) \
42 case DTypeTrait<_dt>::enumv: { \
43 using ctype = DTypeTrait<_dt>::ctype; \
44 out_size = gen_idx( \
45 wk_bundle.get(1), wk_bundle.get_size(1), idx_tmp, mask.ptr<ctype>(), \
46 size, static_cast<uint32_t>(param().mode), kparam, stream); \
47 break; \
48 }
49 MEGDNN_FOREACH_COMPUTING_DTYPE(cb)
50 cb(::megdnn::dtype::Bool)
51#undef cb
52 default : megdnn_throw("bad mask dtype");
53 }
54
55 auto out_data = malloc_policy.alloc_output(0, data.layout.dtype, {out_size});
56 auto out_idx = malloc_policy.alloc_output(1, dtype::Int32(), {out_size});
57 auto out_idx_ptr = out_idx.ptr<dt_int32>();
58
59 switch (data.layout.dtype.enumv()) {
60#define cb(_dt) \
61 case DTypeTrait<_dt>::enumv: { \
62 using ctype = DTypeTrait<_dt>::ctype; \
63 auto out_data_ptr = out_data.ptr<ctype>(); \
64 auto data_ptr = data.ptr<ctype>(); \
65 copy_output<ctype>( \
66 out_data_ptr, out_idx_ptr, data_ptr, idx_tmp, size, stream); \
67 break; \
68 }
69 MEGDNN_FOREACH_COMPUTING_DTYPE(cb)
70 cb(::megdnn::dtype::Bool)
71#undef cb
72 default : megdnn_throw("bad data dtype");
73 }
74
75 return {{out_data, out_idx}};
76}
77
78// vim: syntax=cpp.doxygen

Callers

nothing calls this directly

Calls 8

make_bundleFunction · 0.85
cuda_streamFunction · 0.85
paramFunction · 0.50
cbFunction · 0.50
setMethod · 0.45
getMethod · 0.45
enumvMethod · 0.45
alloc_outputMethod · 0.45

Tested by

no test coverage detected