| 13 | |
| 14 | template <typename data_type, class Opr, typename idx_type = dt_int32> |
| 15 | void do_exec( |
| 16 | const TensorND& data, const TensorND& value, |
| 17 | const IndexingMultiAxisVec::IndexDesc& index, |
| 18 | const IndexingMultiAxisVec::ExecInfo& exec_info) { |
| 19 | size_t nonidx_axes[TensorLayout::MAX_NDIM], |
| 20 | nr_nonidx_axes = IndexingMultiAxisVec::get_nonindex_axes( |
| 21 | data.layout.ndim, index, nonidx_axes); |
| 22 | |
| 23 | auto data_layout = data.layout; |
| 24 | auto data_ptr = data.ptr<data_type>(); |
| 25 | std::tuple<size_t, const idx_type*, TensorLayout> index_raw[TensorLayout::MAX_NDIM]; |
| 26 | size_t nr_index = index.size(); |
| 27 | TensorShape idx_shape; |
| 28 | { |
| 29 | TensorShapeArray idx_shapes; |
| 30 | for (size_t i = 0; i < nr_index; ++i) { |
| 31 | idx_shapes.push_back(index[i].vec.layout); |
| 32 | } |
| 33 | Elemwise::deduce_shape(idx_shapes, idx_shape); |
| 34 | } |
| 35 | for (size_t i = 0; i < nr_index; ++i) { |
| 36 | auto&& s = index[i]; |
| 37 | index_raw[i] = std::make_tuple( |
| 38 | s.axis, s.vec.ptr<idx_type>(), s.vec.layout.broadcast(idx_shape)); |
| 39 | } |
| 40 | |
| 41 | auto value_iter = tensor_iter<data_type>(value).begin(); |
| 42 | for (size_t _ = 0, _t = value.layout.total_nr_elems(); _ < _t; ++_) { |
| 43 | ptrdiff_t offset = 0; |
| 44 | auto* index_idx = value_iter.idx() + exec_info.idx_axis; |
| 45 | for (size_t i = 0; i < nr_index; ++i) { |
| 46 | size_t axis = std::get<0>(index_raw[i]), |
| 47 | data_shape = data_layout.shape[axis]; |
| 48 | ptrdiff_t data_stride = data_layout.stride[axis]; |
| 49 | size_t index_offset = 0; |
| 50 | TensorLayout& index_layout = std::get<2>(index_raw[i]); |
| 51 | for (size_t i = 0; i < index_layout.ndim; ++i) { |
| 52 | index_offset += index_idx[i] * index_layout.stride[i]; |
| 53 | } |
| 54 | idx_type data_idx = std::get<1>(index_raw[i])[index_offset]; |
| 55 | if (data_idx < 0) |
| 56 | data_idx += data_shape; |
| 57 | megdnn_assert( |
| 58 | data_idx >= 0 && static_cast<size_t>(data_idx) < data_shape, |
| 59 | "invalid advanced indexing: " |
| 60 | "input index %d is out of bounds for axis %zu with size %zu", |
| 61 | data_idx, i, data_shape); |
| 62 | offset += data_stride * data_idx; |
| 63 | } |
| 64 | for (size_t i = 0; i < nr_nonidx_axes; ++i) { |
| 65 | auto stride = data_layout.stride[nonidx_axes[i]]; |
| 66 | auto idx = value_iter.idx()[i + (i >= exec_info.idx_axis) * idx_shape.ndim]; |
| 67 | offset += stride * idx; |
| 68 | } |
| 69 | Opr::apply(data_ptr[offset], *value_iter); |
| 70 | ++value_iter; |
| 71 | } |
| 72 | } |