9std::optional<std::vector<RuntimeTensor>>
executeIndex(
const Node&
n,
const std::vector<const RuntimeTensor*>& in) {
11 auto one = [](
RuntimeTensor t) {
return std::vector<RuntimeTensor>{std::move(
t)}; };
12 if (
n.op ==
"Range") {
16 x.element != delta.element)
17 throw Failure(
"Range requires matching scalars");
19 const double start = read<float>(
x, 0), stop = read<float>(
end, 0),
step = read<float>(delta, 0);
20 if (!std::isfinite(
start) || !std::isfinite(stop) || !std::isfinite(
step) ||
step == 0)
23 if (
length > 128u * 1024u * 1024u)
throw Failure(
"Range too large");
24 std::vector<float>
v(
static_cast<size_t>(
length));
25 for (
size_t i = 0; i <
v.size(); ++i)
v[i] =
static_cast<float>(
start + i *
step);
26 return one(
make(
x.element, {static_cast<int64_t>(v.size())},
v));
29 throw Failure(
"Range integer dtype unsupported");
35 if (
length > 128u * 1024u * 1024u)
throw Failure(
"Range too large");
39 for (
size_t i = 0; i <
length; ++i) {
43 return one(std::move(out));
47 if (
count(kt.shape) != 1)
throw Failure(
"TopK requires one K");
50 if (k < 0 || k >
size)
throw Failure(
"Invalid TopK K");
51 size_t outer = 1, inner = 1;
52 for (
int j = 0; j <
a; ++j) outer *=
x.shape[j];
53 for (
size_t j =
a + 1; j <
x.shape.size(); ++j) inner *=
x.shape[j];
58 for (
size_t o = 0; o < outer; ++o)
59 for (
size_t i = 0; i < inner; ++i) {
62 std::stable_sort(
order.begin(),
order.end(), [&](int64_t
aa, int64_t bb) {
63 auto compare = [&](auto u, auto v) { return attr(n,
"largest", 1) ? u > v : u < v; };
64 const size_t ai = (o *
size +
aa) * inner + i, bi = (o *
size + bb) * inner + i;
68 for (int64_t j = 0; j < k; ++j) {
69 size_t dst = (o * k + j) * inner + i, src = (o *
size +
order[j]) * inner + i;
77 if (
n.op ==
"ScatterElements" ||
n.op ==
"ScatterND") {
79 const auto& updates =
required(in, 2);
81 throw Failure(
"Scatter dtype mismatch");
84 if (
n.op ==
"ScatterElements") {
85 if (
idx.shape != updates.shape ||
idx.shape.size() !=
x.shape.size())
86 throw Failure(
"ScatterElements shape mismatch");
88 for (
size_t j = 0; j <
x.shape.size(); ++j)
89 if (j !=
a &&
idx.shape[j] >
x.shape[j])
throw Failure(
"Scatter index shape exceeds data");
90 for (
size_t i = 0; i <
count(
idx.shape); ++i) {
91 size_t rest = i,
source = 0, stride = 1;
92 for (
size_t j =
x.shape.size(); j > 0; --j) {
93 int64_t
c = rest %
idx.shape[j - 1];
94 rest /=
idx.shape[j - 1];
97 if (
c < 0)
c +=
x.shape[j - 1];
99 if (c < 0 || c >=
x.shape[j - 1])
throw Failure(
"Scatter index out of range");
101 stride *=
x.shape[j - 1];
106 if (
idx.shape.empty() ||
idx.shape.back() < 1 ||
static_cast<size_t>(
idx.shape.back()) >
x.shape.size())
107 throw Failure(
"ScatterND index rank mismatch");
108 const size_t depth =
idx.shape.back();
109 std::vector<int64_t> expected(
idx.shape.begin(),
idx.shape.end() - 1);
110 expected.insert(expected.end(),
x.shape.begin() +
depth,
x.shape.end());
111 if (expected != updates.shape)
throw Failure(
"ScatterND updates shape mismatch");
113 for (
size_t j =
depth; j <
x.shape.size(); ++j) chunk *=
x.shape[j];
116 for (
size_t j = 0; j <
depth; ++j) {
118 if (
c < 0)
c +=
x.shape[j];
119 if (c < 0 || c >=
x.shape[j])
throw Failure(
"ScatterND index out of range");
123 std::memcpy(out.bytes.data() +
target * chunk *
bytes, updates.bytes.data() + i * chunk *
bytes,
127 return one(std::move(out));
129 if (
n.op ==
"ReduceProd" ||
n.op ==
"ReduceMax") {
130 auto axes =
attrs(
n,
"axes", {});
132 for (
size_t i = 0; i <
x.shape.size(); ++i) axes.push_back(i);
133 std::set<int> reduced;
136 auto mapShape =
x.shape;
137 std::vector<int64_t>
shape;
138 for (
size_t j = 0; j <
x.shape.size(); ++j) {
139 if (reduced.contains(
static_cast<int>(j))) {
141 if (
attr(
n,
"keepdims", 1))
shape.push_back(1);
143 shape.push_back(
x.shape[j]);
146 std::vector<float> out(
count(
shape),
n.op ==
"ReduceProd" ? 1.f : -
std::numeric_limits<float>::infinity());
147 for (
size_t i = 0; i <
count(
x.shape); ++i) {
149 const float v = read<float>(
x, i);
150 out[
t] =
n.op ==
"ReduceProd" ? out[
t] *
v : std::max(out[
t],
v);
155 throw Failure(
"Integer reduction dtype unsupported");
156 std::vector<int64_t> out(
count(
shape),
n.op ==
"ReduceProd" ? 1 : INT64_MIN);
157 for (
size_t i = 0; i <
count(
x.shape); ++i) {
160 if (
n.op ==
"ReduceMax")
164 (
a > 0 ? (
b > 0 ?
a > INT64_MAX /
b :
b < INT64_MIN /
a)
165 : (
b > 0 ?
a < INT64_MIN /
b :
a < INT64_MAX /
b)))
166 throw
Failure(
"ReduceProd overflow");
171 std::vector<int32_t> narrow;
173 if (v < INT32_MIN || v > INT32_MAX)
throw Failure(
"Int32 reduction overflow");
174 narrow.push_back(
static_cast<int32_t
>(
v));