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