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

Function do_exec

dnn/src/naive/indexing_multi_axis_vec/opr_impl.cpp:15–72  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

13
14template <typename data_type, class Opr, typename idx_type = dt_int32>
15void 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}

Callers

nothing calls this directly

Calls 7

idxMethod · 0.80
applyFunction · 0.50
sizeMethod · 0.45
push_backMethod · 0.45
broadcastMethod · 0.45
beginMethod · 0.45
total_nr_elemsMethod · 0.45

Tested by

no test coverage detected