| 45 | |
| 46 | template <typename data_type, typename idx_type = dt_int32> |
| 47 | void exec_set( |
| 48 | const TensorND& data, const TensorND& index, const TensorND& sub, |
| 49 | uint32_t axis) { |
| 50 | TensorND data_nomid = data; |
| 51 | data_nomid.layout.remove_axis_inplace(axis); |
| 52 | auto data_mid_stride = data.layout.stride[axis]; |
| 53 | int data_mid_shape = data.layout.shape[axis]; |
| 54 | |
| 55 | size_t nr_elems = data_nomid.layout.total_nr_elems(); |
| 56 | megdnn_assert( |
| 57 | nr_elems == index.layout.total_nr_elems() && |
| 58 | nr_elems == sub.layout.total_nr_elems()); |
| 59 | auto data_iter = tensor_iter_valonly<data_type>(data_nomid).begin(); |
| 60 | auto idx_iter = tensor_iter_valonly<idx_type>(index).begin(); |
| 61 | auto sub_iter = tensor_iter_valonly<data_type>(sub).begin(); |
| 62 | |
| 63 | data_type* dptr = data.ptr<data_type>(); |
| 64 | |
| 65 | for (size_t i = 0; i < nr_elems; ++i) { |
| 66 | auto idx = *idx_iter; |
| 67 | megdnn_assert(idx >= 0 && idx < data_mid_shape); |
| 68 | dptr[data_iter.offset() + *idx_iter * data_mid_stride] = *sub_iter; |
| 69 | ++data_iter; |
| 70 | ++sub_iter; |
| 71 | ++idx_iter; |
| 72 | } |
| 73 | } |
| 74 | |
| 75 | } // anonymous namespace |
| 76 |
nothing calls this directly
no test coverage detected