| 26 | } |
| 27 | |
| 28 | CondTakeImpl::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 |
nothing calls this directly
no test coverage detected