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

Method exec

dnn/src/rocm/indexing_one_hot/opr_impl.cpp:26–49  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

24} // anonymous namespace
25
26void IndexingOneHotForwardImpl::exec(
27 _megdnn_tensor_in src, _megdnn_tensor_in index, _megdnn_tensor_out dst,
28 _megdnn_workspace workspace) {
29 check_exec(src.layout, index.layout, dst.layout, workspace.size);
30 ElemwiseOpParamN<0> ele_param{dst.layout.total_nr_elems()};
31 auto kern_param = make_kern_param(src.layout, m_param.axis);
32 auto stream = hip_stream(handle());
33 kern_param.error_tracker = m_error_tracker;
34 kern_param.error_info = async_error_info(handle());
35
36#define cb(_dt) \
37 case DTypeTrait<_dt>::enumv: { \
38 using ctype = DTypeTrait<_dt>::ctype; \
39 using Op = OpGet<DTypeTrait<_dt>::ctype, dt_int32>; \
40 Op op{src.ptr<ctype>(), index.ptr<dt_int32>(), dst.ptr<ctype>(), kern_param}; \
41 return run_elemwise<Op, void>(ele_param, stream, op); \
42 }
43 switch (src.layout.dtype.enumv()) {
44 MEGDNN_FOREACH_COMPUTING_DTYPE(cb)
45 default:
46 megdnn_throw("bad dtype");
47 }
48#undef cb
49}
50
51void IndexingSetOneHotForwardImpl::exec(
52 _megdnn_tensor_inout data, _megdnn_tensor_in index, _megdnn_tensor_in sub,

Callers

nothing calls this directly

Calls 5

hip_streamFunction · 0.85
make_kern_paramFunction · 0.70
async_error_infoFunction · 0.50
total_nr_elemsMethod · 0.45
enumvMethod · 0.45

Tested by

no test coverage detected