载入中...
搜索中...
未找到
OnnxMisc.cpp
浏览该文件的文档.
1#include <algorithm>
2#include <limits>
3#include <sstream>
5
7std::optional<RuntimeTensor> executeMisc(const Node& n, const std::vector<const RuntimeTensor*>& in,
9 const auto& x = required(in, 0);
10 if (n.op == "Clip") {
11 if (x.element != OnnxElement::Float32)
12 throw Failure("Clip currently requires FP32", DiagnosticCode::Unsupported);
13 auto low = make(x.element, {}, std::vector<float>{-std::numeric_limits<float>::infinity()}),
14 high = make(x.element, {}, std::vector<float>{std::numeric_limits<float>::infinity()});
15 const auto& a = in.size() > 1 && in[1] ? *in[1] : low;
16 const auto& b = in.size() > 2 && in[2] ? *in[2] : high;
17 if (a.element != x.element || b.element != x.element || count(a.shape) != 1 || count(b.shape) != 1)
18 throw Failure("Clip bounds must be matching scalars");
19 if (compute) return dispatchFloat(*compute, {&x, &a, &b}, x.shape, "y[i]=min(max(x0[i],x1[0]),x2[0]);");
20 auto v = floats(x);
21 for (auto& f : v) f = std::min(std::max(f, read<float>(a, 0)), read<float>(b, 0));
22 return make(x.element, x.shape, v);
23 }
24 if (n.op == "CumSum") {
25 const auto& at = required(in, 1);
26 if (count(at.shape) != 1) throw Failure("CumSum axis must be scalar");
27 const int a = axis(integer(at), x.shape.size());
28 if (x.element != OnnxElement::Float32)
29 throw Failure("CumSum currently requires FP32", DiagnosticCode::Unsupported);
30 const bool reverse = attr(n, "reverse", 0) != 0, exclusive = attr(n, "exclusive", 0) != 0;
31 size_t outer = 1, inner = 1;
32 for (int j = 0; j < a; ++j) outer *= x.shape[j];
33 for (size_t j = a + 1; j < x.shape.size(); ++j) inner *= x.shape[j];
34 const auto width = x.shape[a];
35 if (compute && outer && inner && width) {
36 std::ostringstream s;
37 s << "uint base=i/" << inner << "u*" << inner * width << "u+i%" << inner
38 << "u;precise float sum=0.0;for(uint j=0;j<" << width << "u;++j){uint index=base+"
39 << (reverse ? "(" + std::to_string(width - 1) + "u-j)" : "j") << "*" << inner << "u;"
40 << (exclusive ? "y[index]=sum;sum+=x0[index];" : "sum+=x0[index];y[index]=sum;") << "}";
41 return dispatchFloat(*compute, {&x}, x.shape, s.str(), outer * inner);
42 }
43 auto out = floats(x);
44 for (size_t o = 0; o < outer; ++o)
45 for (size_t i = 0; i < inner; ++i) {
46 float sum = 0;
47 for (int64_t j = 0; j < width; ++j) {
48 size_t p = (o * width + (reverse ? width - 1 - j : j)) * inner + i;
49 const float v = out[p];
50 if (exclusive) {
51 out[p] = sum;
52 sum += v;
53 } else {
54 sum += v;
55 out[p] = sum;
56 }
57 }
58 }
59 return make(x.element, x.shape, out);
60 }
61 if (n.op == "Pad") {
62 const auto pads = ints(required(in, 1));
63 if (pads.size() != 2 * x.shape.size()) throw Failure("Pad rank mismatch");
64 auto shape = x.shape;
65 const auto mode = n.attrs.contains("mode") ? n.attrs.at("mode").text : "constant";
66 for (size_t j = 0; j < shape.size(); ++j) {
67 if (pads[j] < -INT32_MAX || pads[j] > INT32_MAX || pads[j + shape.size()] < -INT32_MAX ||
68 pads[j + shape.size()] > INT32_MAX)
69 throw Failure("Pad extent out of range");
70 shape[j] += pads[j] + pads[j + shape.size()];
71 if (mode != "constant" && x.shape[j] <= 0)
72 throw Failure("Edge/reflect padding requires nonempty dimensions");
73 if (mode == "reflect" && (pads[j] >= x.shape[j] || pads[j + shape.size()] >= x.shape[j]))
74 throw Failure("Reflect padding exceeds input dimension");
75 }
76 const auto size = elementSize(x.element);
77 std::vector<uint8_t> fill(size, 0);
78 if (in.size() > 2 && in[2]) {
79 if (in[2]->element != x.element || count(in[2]->shape) != 1)
80 throw Failure("Pad value dtype/shape mismatch");
81 fill = in[2]->bytes;
82 }
83 RuntimeTensor out{x.element, shape, std::vector<uint8_t>(count(shape) * size)};
84 for (size_t i = 0; i < count(shape); ++i) {
85 size_t rest = i, source = 0, stride = 1;
86 bool outside = false;
87 for (size_t j = shape.size(); j > 0; --j) {
88 int64_t p = static_cast<int64_t>(rest % shape[j - 1]) - pads[j - 1];
89 rest /= shape[j - 1];
90 if (p < 0 || p >= x.shape[j - 1]) {
91 if (mode == "constant")
92 outside = true;
93 else if (mode == "edge")
94 p = std::clamp(p, int64_t(0), x.shape[j - 1] - 1);
95 else
96 p = p < 0 ? -p : 2 * x.shape[j - 1] - 2 - p;
97 }
98 if (!outside) source += p * stride;
99 stride *= x.shape[j - 1];
100 }
101 std::memcpy(out.bytes.data() + i * size, outside ? fill.data() : x.bytes.data() + source * size, size);
102 }
103 return out;
104 }
105 return std::nullopt;
106}
107} // namespace eve::tensor::onnx_detail
float x
Definition AnimClip.cpp:738
const std::string & s
glm::vec4 p[6]
ShaderImageInput shape
glm::vec3 n
Definition Grass.cpp:63
float v
std::uint32_t width
bool required
MeleePoint3 b
Definition MeleeHit.cpp:41
MeleePoint3 a
Definition MeleeHit.cpp:40
OnnxCompute * compute
float f
std::uint32_t count
std::string element
float size
Definition TreeMesh.cpp:156
const UnitySourceAsset & source
std::size_t at
GPU execution boundary for native ONNX; retains no model and retains compiled resources for the lifet...
Definition OnnxCompute.h:25
int64_t integer(const RuntimeTensor &v, size_t i=0)
Integer.
std::optional< RuntimeTensor > executeMisc(const Node &, const std::vector< const RuntimeTensor * > &, OnnxCompute *)
Execute misc.
Definition OnnxMisc.cpp:7
std::vector< int64_t > ints(const RuntimeTensor &v)
Ints.
int64_t attr(const Node &n, const char *key, int64_t fallback)
Attr.
RuntimeTensor make(OnnxElement e, std::vector< int64_t > shape, const std::vector< T > &values)
Make.
int axis(int64_t a, size_t rank)
Axis.
std::vector< float > floats(const RuntimeTensor &v)
Floats.
RuntimeTensor dispatchFloat(OnnxCompute &, const std::vector< const RuntimeTensor * > &, const std::vector< int64_t > &, const std::string &, size_t work=0)
Dispatches float.
size_t elementSize(OnnxElement e)
Element size.