载入中...
搜索中...
未找到
OnnxIndexOps.cpp
浏览该文件的文档.
1#include <algorithm>
2#include <cmath>
3#include <limits>
4#include <numeric>
5#include <set>
7
9std::optional<std::vector<RuntimeTensor>> executeIndex(const Node& n, const std::vector<const RuntimeTensor*>& in) {
10 const auto& x = required(in, 0);
11 auto one = [](RuntimeTensor t) { return std::vector<RuntimeTensor>{std::move(t)}; };
12 if (n.op == "Range") {
13 const auto& end = required(in, 1);
14 const auto& delta = required(in, 2);
15 if (count(x.shape) != 1 || count(end.shape) != 1 || count(delta.shape) != 1 || x.element != end.element ||
16 x.element != delta.element)
17 throw Failure("Range requires matching scalars");
18 if (x.element == OnnxElement::Float32) {
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)
21 throw Failure("Invalid Range");
22 const double length = std::max(0., std::ceil((stop - start) / step));
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));
27 }
28 if (x.element != OnnxElement::Int64 && x.element != OnnxElement::Int32)
29 throw Failure("Range integer dtype unsupported");
30 const int64_t start = integer(x), stop = integer(end), step = integer(delta);
31 if (!step || step == INT64_MIN) throw Failure("Invalid Range step");
32 uint64_t length = 0;
33 if (step > 0 && stop > start) length = 1 + (uint64_t(stop) - uint64_t(start) - 1) / uint64_t(step);
34 if (step < 0 && stop < start) length = 1 + (uint64_t(start) - uint64_t(stop) - 1) / uint64_t(-step);
35 if (length > 128u * 1024u * 1024u) throw Failure("Range too large");
36 RuntimeTensor out{
37 x.element, {static_cast<int64_t>(length)}, std::vector<uint8_t>(length * elementSize(x.element))};
38 int64_t v = start;
39 for (size_t i = 0; i < length; ++i) {
40 std::memcpy(out.bytes.data() + i * elementSize(x.element), &v, elementSize(x.element));
41 if (i + 1 < length) v += step;
42 }
43 return one(std::move(out));
44 }
45 if (n.op == "TopK") {
46 const auto& kt = required(in, 1);
47 if (count(kt.shape) != 1) throw Failure("TopK requires one K");
48 int a = axis(attr(n, "axis", -1), x.shape.size());
49 const int64_t k = integer(kt), size = x.shape[a];
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];
54 auto shape = x.shape;
55 shape[a] = k;
56 RuntimeTensor values{x.element, shape, std::vector<uint8_t>(count(shape) * elementSize(x.element))};
57 std::vector<int64_t> indices(count(shape));
58 for (size_t o = 0; o < outer; ++o)
59 for (size_t i = 0; i < inner; ++i) {
60 std::vector<int64_t> order(size);
61 std::iota(order.begin(), order.end(), 0);
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;
65 return x.element == OnnxElement::Float32 ? compare(read<float>(x, ai), read<float>(x, bi))
66 : compare(integer(x, ai), integer(x, bi));
67 });
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;
70 indices[dst] = order[j];
71 std::memcpy(values.bytes.data() + dst * elementSize(x.element),
72 x.bytes.data() + src * elementSize(x.element), elementSize(x.element));
73 }
74 }
75 return std::vector<RuntimeTensor>{std::move(values), make(OnnxElement::Int64, shape, indices)};
76 }
77 if (n.op == "ScatterElements" || n.op == "ScatterND") {
78 const auto& idx = required(in, 1);
79 const auto& updates = required(in, 2);
80 if (updates.element != x.element || (idx.element != OnnxElement::Int64 && idx.element != OnnxElement::Int32))
81 throw Failure("Scatter dtype mismatch");
82 auto out = x;
83 const auto bytes = elementSize(x.element);
84 if (n.op == "ScatterElements") {
85 if (idx.shape != updates.shape || idx.shape.size() != x.shape.size())
86 throw Failure("ScatterElements shape mismatch");
87 int a = axis(attr(n, "axis", 0), x.shape.size());
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];
95 if (j - 1 == a) {
96 c = integer(idx, i);
97 if (c < 0) c += x.shape[j - 1];
98 }
99 if (c < 0 || c >= x.shape[j - 1]) throw Failure("Scatter index out of range");
100 source += c * stride;
101 stride *= x.shape[j - 1];
102 }
103 std::memcpy(out.bytes.data() + source * bytes, updates.bytes.data() + i * bytes, bytes);
104 }
105 } else {
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");
112 size_t chunk = 1;
113 for (size_t j = depth; j < x.shape.size(); ++j) chunk *= x.shape[j];
114 for (size_t i = 0; i < count(idx.shape) / depth; ++i) {
115 size_t target = 0;
116 for (size_t j = 0; j < depth; ++j) {
117 auto c = integer(idx, i * 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");
120 target = target * x.shape[j] + c;
121 }
122 if (chunk)
123 std::memcpy(out.bytes.data() + target * chunk * bytes, updates.bytes.data() + i * chunk * bytes,
124 chunk * bytes);
125 }
126 }
127 return one(std::move(out));
128 }
129 if (n.op == "ReduceProd" || n.op == "ReduceMax") {
130 auto axes = attrs(n, "axes", {});
131 if (axes.empty())
132 for (size_t i = 0; i < x.shape.size(); ++i) axes.push_back(i);
133 std::set<int> reduced;
134 for (auto a : axes)
135 if (!reduced.insert(axis(a, x.shape.size())).second) throw Failure("Duplicate reduction axis");
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))) {
140 mapShape[j] = 1;
141 if (attr(n, "keepdims", 1)) shape.push_back(1);
142 } else
143 shape.push_back(x.shape[j]);
144 }
145 if (x.element == OnnxElement::Float32) {
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) {
148 size_t t = broadcastIndex(i, mapShape, x.shape);
149 const float v = read<float>(x, i);
150 out[t] = n.op == "ReduceProd" ? out[t] * v : std::max(out[t], v);
151 }
152 return one(make(x.element, shape, out));
153 }
154 if (x.element != OnnxElement::Int64 && x.element != OnnxElement::Int32)
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) {
158 auto& a = out[broadcastIndex(i, mapShape, x.shape)];
159 const auto b = integer(x, i);
160 if (n.op == "ReduceMax")
161 a = std::max(a, b);
162 else {
163 if (a && b &&
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");
167 a *= b;
168 }
169 }
170 if (x.element == OnnxElement::Int64) return one(make(x.element, shape, out));
171 std::vector<int32_t> narrow;
172 for (auto v : out) {
173 if (v < INT32_MIN || v > INT32_MAX) throw Failure("Int32 reduction overflow");
174 narrow.push_back(static_cast<int32_t>(v));
175 }
176 return one(make(x.element, shape, narrow));
177 }
178 return std::nullopt;
179}
180} // namespace eve::tensor::onnx_detail
LogicalId target
Duration start
float x
Definition AnimClip.cpp:738
float length
Definition CaveMesh.cpp:94
std::map< std::string, Var > values
ShaderImageInput shape
glm::vec3 n
Definition Grass.cpp:63
std::uint32_t aa
std::vector< std::uint32_t > indices
float v
std::int32_t second
std::int32_t c
bool required
std::uint64_t bytes
MeleePoint3 b
Definition MeleeHit.cpp:41
MeleePoint3 a
Definition MeleeHit.cpp:40
std::vector< std::int32_t > order
int idx
float t
std::uint32_t count
float step
Definition TreeMesh.cpp:314
float size
Definition TreeMesh.cpp:156
const UnitySourceAsset & source
std::uint32_t depth
int64_t integer(const RuntimeTensor &v, size_t i=0)
Integer.
int64_t attr(const Node &n, const char *key, int64_t fallback)
Attr.
size_t broadcastIndex(size_t i, const std::vector< int64_t > &shape, const std::vector< int64_t > &output)
Broadcast index.
std::vector< int64_t > attrs(const Node &n, const char *key, std::vector< int64_t > fallback)
Attrs.
RuntimeTensor make(OnnxElement e, std::vector< int64_t > shape, const std::vector< T > &values)
Make.
int axis(int64_t a, size_t rank)
Axis.
size_t elementSize(OnnxElement e)
Element size.
std::optional< std::vector< RuntimeTensor > > executeIndex(const Node &n, const std::vector< const RuntimeTensor * > &in)
Execute index.