载入中...
搜索中...
未找到
OnnxShapeOps.cpp
浏览该文件的文档.
1#include <algorithm>
2#include <cmath>
3#include <numeric>
4#include <set>
6
8std::vector<int64_t> broadcast(const std::vector<int64_t>& a, const std::vector<int64_t>& b) {
9 std::vector<int64_t> out(std::max(a.size(), b.size()), 1);
10 for (size_t i = 0; i < out.size(); ++i) {
11 const auto x = i < a.size() ? a[a.size() - 1 - i] : 1, y = i < b.size() ? b[b.size() - 1 - i] : 1;
12 if (x != y && x != 1 && y != 1) throw Failure("Broadcast dimension mismatch");
13 out[out.size() - 1 - i] = x == 1 ? y : x;
14 }
15 count(out);
16 return out;
17}
18size_t broadcastIndex(size_t i, const std::vector<int64_t>& shape, const std::vector<int64_t>& output) {
19 size_t index = 0, stride = 1;
20 for (size_t j = 0; j < output.size(); ++j) {
21 const auto coordinate = i % output[output.size() - 1 - j];
22 i /= output[output.size() - 1 - j];
23 if (j < shape.size()) {
24 const auto d = shape[shape.size() - 1 - j];
25 if (d != 1) index += coordinate * stride;
26 stride *= d;
27 }
28 }
29 return index;
30}
31std::optional<RuntimeTensor> executeShape(const Node& n, const std::vector<const RuntimeTensor*>& in) {
32 const auto& x = required(in, 0);
33 if (n.op == "Slice") {
34 auto starts = ints(required(in, 1)), ends = ints(required(in, 2));
35 auto axes = in.size() > 3 && in[3] ? ints(*in[3]) : std::vector<int64_t>(starts.size());
36 if (in.size() <= 3 || !in[3]) std::iota(axes.begin(), axes.end(), 0);
37 auto steps = in.size() > 4 && in[4] ? ints(*in[4]) : std::vector<int64_t>(starts.size(), 1);
38 if (starts.size() != ends.size() || starts.size() != axes.size() || starts.size() != steps.size())
39 throw Failure("Slice parameter lengths mismatch");
40 auto shape = x.shape;
41 std::vector<int64_t> begin(shape.size(), 0), step(shape.size(), 1);
42 std::set<int> seen;
43 for (size_t j = 0; j < axes.size(); ++j) {
44 int a = axis(axes[j], shape.size());
45 if (!seen.insert(a).second || !steps[j] || steps[j] == INT64_MIN) throw Failure("Invalid Slice axis/step");
46 const auto d = x.shape[a], st = steps[j];
47 auto norm = [d, st](int64_t v, bool end) {
48 if (v < 0) v = v < -d ? (st < 0 ? -1 : 0) : v + d;
49 return std::clamp(v, st < 0 && end ? int64_t(-1) : int64_t(0),
50 st < 0 ? std::max(int64_t(0), d - 1) : d);
51 };
52 begin[a] = norm(starts[j], false);
53 const auto e = norm(ends[j], true);
54 step[a] = st;
55 const int64_t distance = st > 0 ? e - begin[a] : begin[a] - e, magnitude = st > 0 ? st : -st;
56 shape[a] = d == 0 || distance <= 0 ? 0 : 1 + (distance - 1) / magnitude;
57 }
58 RuntimeTensor out{x.element, shape, std::vector<uint8_t>(count(shape) * elementSize(x.element))};
59 for (size_t i = 0; i < count(shape); ++i) {
60 size_t rest = i, source = 0, stride = 1;
61 for (size_t j = shape.size(); j > 0; --j) {
62 auto c = rest % shape[j - 1];
63 rest /= shape[j - 1];
64 source += (begin[j - 1] + c * step[j - 1]) * stride;
65 stride *= x.shape[j - 1];
66 }
67 std::memcpy(out.bytes.data() + i * elementSize(x.element), x.bytes.data() + source * elementSize(x.element),
68 elementSize(x.element));
69 }
70 return out;
71 }
72 if (n.op == "Concat") {
73 int a = axis(attr(n, "axis", 0), x.shape.size());
74 auto shape = x.shape;
75 shape[a] = 0;
76 for (auto* t : in) {
77 if (!t || t->element != x.element || t->shape.size() != shape.size())
78 throw Failure("Concat type/rank mismatch");
79 for (size_t j = 0; j < shape.size(); ++j)
80 if (j != a && t->shape[j] != shape[j]) throw Failure("Concat dimension mismatch");
81 shape[a] += t->shape[a];
82 }
83 RuntimeTensor out{x.element, shape, std::vector<uint8_t>(count(shape) * elementSize(x.element))};
84 size_t outer = 1, inner = elementSize(x.element);
85 for (int j = 0; j < a; ++j) outer *= shape[j];
86 for (size_t j = a + 1; j < shape.size(); ++j) inner *= shape[j];
87 size_t offset = 0;
88 for (size_t o = 0; o < outer; ++o)
89 for (auto* t : in) {
90 const size_t size = t->shape[a] * inner;
91 if (size) std::memcpy(out.bytes.data() + offset, t->bytes.data() + o * size, size);
92 offset += size;
93 }
94 return out;
95 }
96 if (n.op == "Expand") {
97 auto requested = ints(required(in, 1));
98 count(requested);
99 auto shape = broadcast(x.shape, requested);
100 RuntimeTensor out{x.element, shape, std::vector<uint8_t>(count(shape) * elementSize(x.element))};
101 for (size_t i = 0; i < count(shape); ++i)
102 std::memcpy(out.bytes.data() + i * elementSize(x.element),
103 x.bytes.data() + broadcastIndex(i, x.shape, shape) * elementSize(x.element),
104 elementSize(x.element));
105 return out;
106 }
107 if (n.op == "ConstantOfShape") {
108 auto shape = ints(x);
109 RuntimeTensor value = make(OnnxElement::Float32, {}, std::vector<float>{0});
110 if (n.attrs.contains("value")) {
111 if (!n.attrs.at("value").tensor) throw Failure("Missing fill value");
112 value = *n.attrs.at("value").tensor;
113 }
114 if (count(value.shape) != 1) throw Failure("ConstantOfShape requires one fill value");
115 RuntimeTensor out{value.element, shape, std::vector<uint8_t>(count(shape) * elementSize(value.element))};
116 for (size_t i = 0; i < out.bytes.size(); i += value.bytes.size())
117 std::memcpy(out.bytes.data() + i, value.bytes.data(), value.bytes.size());
118 return out;
119 }
120 if (n.op == "Where") {
121 const auto& a = required(in, 1);
122 const auto& b = required(in, 2);
123 if (x.element != OnnxElement::Bool || a.element != b.element) throw Failure("Where type mismatch");
124 auto shape = broadcast(broadcast(x.shape, a.shape), b.shape);
125 const auto size = elementSize(a.element);
126 RuntimeTensor out{a.element, shape, std::vector<uint8_t>(count(shape) * size)};
127 for (size_t i = 0; i < count(shape); ++i) {
128 const auto& v = x.bytes[broadcastIndex(i, x.shape, shape)] ? a : b;
129 std::memcpy(out.bytes.data() + i * size, v.bytes.data() + broadcastIndex(i, v.shape, shape) * size, size);
130 }
131 return out;
132 }
133 if (n.op == "Not") {
134 if (x.element != OnnxElement::Bool) throw Failure("Not requires boolean");
135 auto out = x;
136 for (auto& b : out.bytes) b = !b;
137 return out;
138 }
139 if (n.op == "Equal" || n.op == "Less" || n.op == "Greater" || n.op == "And") {
140 const auto& y = required(in, 1);
141 if (x.element != y.element || (n.op == "And" && x.element != OnnxElement::Bool))
142 throw Failure("Comparison type mismatch");
143 auto shape = broadcast(x.shape, y.shape);
144 RuntimeTensor out{OnnxElement::Bool, shape, std::vector<uint8_t>(count(shape))};
145 for (size_t i = 0; i < out.bytes.size(); ++i) {
146 auto a = broadcastIndex(i, x.shape, shape), b = broadcastIndex(i, y.shape, shape);
147 auto compare = [&](auto u, auto v) {
148 return n.op == "Equal" ? u == v
149 : n.op == "Less" ? u < v
150 : n.op == "Greater" ? u > v
151 : bool(u) && bool(v);
152 };
153 out.bytes[i] = x.element == OnnxElement::Float32 ? compare(read<float>(x, a), read<float>(y, b))
154 : compare(integer(x, a), integer(y, b));
155 }
156 return out;
157 }
158 if ((n.op == "Add" || n.op == "Sub" || n.op == "Mul" || n.op == "Div") && x.element != OnnxElement::Float32) {
159 const auto& y = required(in, 1);
160 if (x.element != y.element || (x.element != OnnxElement::Int64 && x.element != OnnxElement::Int32))
161 throw Failure("Integer arithmetic dtype mismatch");
162 auto shape = broadcast(x.shape, y.shape);
163 RuntimeTensor out{x.element, shape, std::vector<uint8_t>(count(shape) * elementSize(x.element))};
164 for (size_t i = 0; i < count(shape); ++i) {
165 int64_t a = integer(x, broadcastIndex(i, x.shape, shape)),
166 b = integer(y, broadcastIndex(i, y.shape, shape)), v = 0;
167 if (n.op == "Add") {
168 if ((b > 0 && a > INT64_MAX - b) || (b < 0 && a < INT64_MIN - b))
169 throw Failure("Integer addition overflow");
170 v = a + b;
171 }
172 if (n.op == "Sub") {
173 if ((b < 0 && a > INT64_MAX + b) || (b > 0 && a < INT64_MIN + b))
174 throw Failure("Integer subtraction overflow");
175 v = a - b;
176 }
177 if (n.op == "Mul") {
178 if (a && b &&
179 (a > 0 ? (b > 0 ? a > INT64_MAX / b : b < INT64_MIN / a)
180 : (b > 0 ? a < INT64_MIN / b : a < INT64_MAX / b)))
181 throw Failure("Integer multiplication overflow");
182 v = a * b;
183 }
184 if (n.op == "Div") {
185 if (!b || (a == INT64_MIN && b == -1)) throw Failure("Invalid integer division");
186 v = a / b;
187 }
188 if (x.element == OnnxElement::Int32 && (v < INT32_MIN || v > INT32_MAX))
189 throw Failure("Int32 arithmetic overflow");
190 std::memcpy(out.bytes.data() + i * elementSize(x.element), &v, elementSize(x.element));
191 }
192 return out;
193 }
194 return std::nullopt;
195}
196} // namespace eve::tensor::onnx_detail
double value
float y
Definition AnimClip.cpp:738
float x
Definition AnimClip.cpp:738
std::string output
ShaderImageInput shape
float u
Definition Grass.cpp:233
glm::vec3 n
Definition Grass.cpp:63
float v
std::int32_t c
size_t offset
bool required
MeleePoint3 b
Definition MeleeHit.cpp:41
MeleePoint3 a
Definition MeleeHit.cpp:40
float distance
float begin
float d
int steps
float t
std::uint32_t count
float step
Definition TreeMesh.cpp:314
float size
Definition TreeMesh.cpp:156
uint32_t index
const UnitySourceAsset & source
int64_t integer(const RuntimeTensor &v, size_t i=0)
Integer.
std::vector< int64_t > ints(const RuntimeTensor &v)
Ints.
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.
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::vector< int64_t > broadcast(const std::vector< int64_t > &a, const std::vector< int64_t > &b)
Broadcast.
std::optional< RuntimeTensor > executeShape(const Node &, const std::vector< const RuntimeTensor * > &)
Execute shape.