载入中...
搜索中...
未找到
OnnxNumeric.cpp
浏览该文件的文档.
1#include <algorithm>
2#include <cmath>
3#include <limits>
4#include <set>
5#include <sstream>
8
10namespace {
11std::string literal(float v) {
12 std::ostringstream s;
13 s.precision(9);
14 s << std::scientific << v;
15 return s.str();
16}
17std::string indexCode(const std::vector<int64_t>& input, const std::vector<int64_t>& output, const char* index) {
18 std::string result = "0u";
19 size_t src = 1, dst = 1;
20 for (size_t j = 0; j < output.size(); ++j) {
21 if (j < input.size()) {
22 const auto dim = input[input.size() - 1 - j];
23 if (dim != 1)
24 result += "+(" + std::string(index) + "/" + std::to_string(dst) + "u%" + std::to_string(dim) + "u)*" +
25 std::to_string(src) + "u";
26 src *= dim;
27 }
28 dst *= output[output.size() - 1 - j];
29 }
30 return result;
31}
32} // namespace
33RuntimeTensor dispatchFloat(OnnxCompute& compute, const std::vector<const RuntimeTensor*>& inputs,
34 const std::vector<int64_t>& shape, const std::string& body, size_t work) {
35 RuntimeTensor out{OnnxElement::Float32, shape, std::vector<uint8_t>(count(shape) * 4)};
36 if (out.bytes.empty()) return out;
37 if (!work) work = count(shape);
38 if (work > 65535u * 64u) throw Failure("GPU float dispatch limit exceeded", DiagnosticCode::Unsupported);
39 OnnxKernel k;
40 std::vector<OnnxBuffer> buffers;
41 k.outputBytes = out.bytes.size();
42 k.workItems = static_cast<uint32_t>(work);
43 k.source = "#version 450\nlayout(local_size_x=64)in;\n";
44 for (size_t i = 0; i < inputs.size(); ++i) {
46 throw Failure("GPU float input type mismatch");
47 buffers.push_back(inputs[i]->bytes.buffer());
48 k.source += "layout(std430,binding=" + std::to_string(i) + ")readonly buffer B" + std::to_string(i) +
49 (inputs[i]->element == OnnxElement::Int32 ? "{int x" : "{float x") + std::to_string(i) + "[];};\n";
50 }
51 k.source +=
52 "layout(std430,binding=" + std::to_string(inputs.size()) +
53 ")writeonly buffer Y{float y[];};\nvoid main(){uint i=gl_GlobalInvocationID.x;if(i>=" + std::to_string(work) +
54 "u)return;" + body + "}";
55 auto r = compute.enqueue(k.source, buffers, k.outputBytes, k.workItems);
56 if (!r.ok()) throw Failure(r.error()->message(), r.error()->code());
57 if (r.value().size != out.bytes.size()) throw Failure("GPU output size mismatch");
58 out.bytes = std::move(r.value());
59 return out;
60}
61std::optional<RuntimeTensor> executeNumeric(const Node& n, const std::vector<const RuntimeTensor*>& in,
63 const auto& x = required(in, 0);
64 if (x.element != OnnxElement::Float32) return std::nullopt;
65 if (n.op == "MatMul") {
66 const auto& b = required(in, 1);
67 if (b.element != x.element || x.shape.empty() || b.shape.empty()) throw Failure("MatMul type/rank mismatch");
68 auto as = x.shape, bs = b.shape;
69 const bool av = as.size() == 1, bv = bs.size() == 1;
70 if (av) as.insert(as.begin(), 1);
71 if (bv) bs.push_back(1);
72 const auto m = as[as.size() - 2], k = as.back(), cols = bs.back();
73 if (k != bs[bs.size() - 2]) throw Failure("MatMul inner dimension mismatch");
74 std::vector<int64_t> ab(as.begin(), as.end() - 2), bb(bs.begin(), bs.end() - 2);
75 auto batches = broadcast(ab, bb);
76 auto shape = batches;
77 shape.push_back(m);
78 shape.push_back(cols);
79 count(shape);
80 if (av) shape.erase(shape.end() - 2);
81 if (bv) shape.pop_back();
82 if (compute && m && cols) {
83 std::string body = "uint batch=i/" + std::to_string(m * cols) + "u,r=i/" + std::to_string(cols) + "u%" +
84 std::to_string(m) + "u,c=i%" + std::to_string(cols) +
85 "u;precise float v=0.0;for(uint j=0;j<" + std::to_string(k) + "u;++j)v+=x0[(" +
86 indexCode(ab, batches, "batch") + ")*" + std::to_string(m * k) + "u+r*" +
87 std::to_string(k) + "u+j]*x1[(" + indexCode(bb, batches, "batch") + ")*" +
88 std::to_string(k * cols) + "u+j*" + std::to_string(cols) + "u+c];y[i]=v;";
89 return dispatchFloat(*compute, {&x, &b}, shape, body);
90 }
91 std::vector<float> out(count(shape), 0);
92 for (size_t batch = 0; batch < count(batches); ++batch)
93 for (int64_t r = 0; r < m; ++r)
94 for (int64_t c = 0; c < cols; ++c) {
95 float sum = 0;
96 for (int64_t j = 0; j < k; ++j)
97 sum += read<float>(x, broadcastIndex(batch, ab, batches) * m * k + r * k + j) *
98 read<float>(b, broadcastIndex(batch, bb, batches) * k * cols + j * cols + c);
99 out[(batch * m + r) * cols + c] = sum;
100 }
101 return make(OnnxElement::Float32, shape, out);
102 }
103 static const std::set<std::string> binary{"Add", "Sub", "Mul", "Div", "Pow"};
104 if (binary.contains(n.op)) {
105 const auto& b = required(in, 1);
106 if (b.element != x.element) throw Failure("Float arithmetic dtype mismatch");
107 auto shape = broadcast(x.shape, b.shape);
108 if (compute) {
109 const std::string a = "x0[" + indexCode(x.shape, shape, "i") + "]",
110 v = "x1[" + indexCode(b.shape, shape, "i") + "]";
111 const std::string e = n.op == "Pow" ? "pow(" + a + "," + v + ")"
112 : a +
113 (n.op == "Add" ? "+"
114 : n.op == "Sub" ? "-"
115 : n.op == "Mul" ? "*"
116 : "/") +
117 v;
118 if (n.op == "Pow")
119 return dispatchFloat(*compute, {&x, &b}, shape,
120 "float a=" + a + ",b=" + v +
121 ";float value;if(b==0.0)value=1.0;else "
122 "if(a==0.0)value=b>0.0?0.0:uintBitsToFloat(0x7f800000u);else "
123 "if(a<0.0)value=b==floor(b)?pow(-a,b)*(mod(abs(b),2.0)==0.0?1.0:-1.0):"
124 "uintBitsToFloat(0x7fc00000u);else value=pow(a,b);y[i]=value;");
125 return dispatchFloat(*compute, {&x, &b}, shape, "y[i]=" + e + ";");
126 }
127 std::vector<float> out(count(shape));
128 for (size_t i = 0; i < out.size(); ++i) {
129 const float a = read<float>(x, broadcastIndex(i, x.shape, shape)),
130 v = read<float>(b, broadcastIndex(i, b.shape, shape));
131 out[i] = n.op == "Add" ? a + v
132 : n.op == "Sub" ? a - v
133 : n.op == "Mul" ? a * v
134 : n.op == "Div" ? a / v
135 : std::pow(a, v);
136 }
137 return make(OnnxElement::Float32, shape, out);
138 }
139 if (n.op == "ReduceMean" || n.op == "ReduceSum") {
140 auto axes = n.op == "ReduceMean" ? attrs(n, "axes", {})
141 : in.size() > 1 && in[1] ? ints(*in[1])
142 : std::vector<int64_t>{};
143 if (axes.empty() && n.op == "ReduceSum" && attr(n, "noop_with_empty_axes", 0)) return x;
144 if (axes.empty())
145 for (size_t i = 0; i < x.shape.size(); ++i) axes.push_back(i);
146 std::set<int> set;
147 for (auto a : axes)
148 if (!set.insert(axis(a, x.shape.size())).second) throw Failure("Duplicate reduction axis");
149 std::vector<int64_t> outerShape, reducedShape, shape;
150 std::vector<size_t> outerStrides, reducedStrides;
151 size_t stride = count(x.shape);
152 for (size_t j = 0; j < x.shape.size(); ++j) {
153 if (!x.shape[j]) throw Failure("Empty reduction unsupported", DiagnosticCode::Unsupported);
154 stride /= x.shape[j];
155 if (set.contains(static_cast<int>(j))) {
156 reducedShape.push_back(x.shape[j]);
157 reducedStrides.push_back(stride);
158 if (attr(n, "keepdims", 1)) shape.push_back(1);
159 } else {
160 outerShape.push_back(x.shape[j]);
161 outerStrides.push_back(stride);
162 shape.push_back(x.shape[j]);
163 }
164 }
165 auto offset = [](size_t i, const std::vector<int64_t>& dims, const std::vector<size_t>& strides) {
166 size_t p = 0;
167 for (size_t j = dims.size(); j > 0; --j) {
168 p += (i % dims[j - 1]) * strides[j - 1];
169 i /= dims[j - 1];
170 }
171 return p;
172 };
173 auto code = [](const char* var, const std::vector<int64_t>& dims, const std::vector<size_t>& strides) {
174 std::string s = "0u";
175 size_t divisor = 1;
176 for (size_t j = dims.size(); j > 0; --j) {
177 s += "+(" + std::string(var) + "/" + std::to_string(divisor) + "u%" + std::to_string(dims[j - 1]) +
178 "u)*" + std::to_string(strides[j - 1]) + "u";
179 divisor *= dims[j - 1];
180 }
181 return s;
182 };
183 const size_t terms = count(reducedShape);
184 if (compute)
185 return dispatchFloat(*compute, {&x}, shape,
186 "uint base=" + code("i", outerShape, outerStrides) +
187 ";precise float v=0.0;for(uint j=0;j<" + std::to_string(terms) +
188 "u;++j)v+=x0[base+(" + code("j", reducedShape, reducedStrides) + ")];y[i]=v" +
189 (n.op == "ReduceMean" ? "/" + std::to_string(terms) + ".0" : "") + ";");
190 std::vector<float> out(count(shape));
191 for (size_t i = 0; i < out.size(); ++i) {
192 float v = 0;
193 for (size_t j = 0; j < terms; ++j)
194 v += read<float>(x, offset(i, outerShape, outerStrides) + offset(j, reducedShape, reducedStrides));
195 out[i] = n.op == "ReduceMean" ? v / terms : v;
196 }
197 return make(OnnxElement::Float32, shape, out);
198 }
199 if (n.op == "Softmax" && compute) {
200 int a = axis(attr(n, "axis", -1), x.shape.size());
201 size_t inner = 1, outer = 1;
202 for (int j = 0; j < a; ++j) outer *= x.shape[j];
203 for (size_t j = a + 1; j < x.shape.size(); ++j) inner *= x.shape[j];
204 const auto channels = x.shape[a];
205 if (!channels || !inner) return RuntimeTensor{x.element, x.shape, {}};
206 auto ch = std::to_string(channels), ins = std::to_string(inner);
207 return dispatchFloat(*compute, {&x}, x.shape,
208 "uint base=i/" + ins + "u*" + ins + "u*" + ch + "u+i%" + ins +
209 "u;float m=-3.402823466e38;for(uint j=0;j<" + ch + "u;++j)m=max(m,x0[base+j*" + ins +
210 "u]);float sum=0.0;for(uint j=0;j<" + ch + "u;++j)sum+=exp(x0[base+j*" + ins +
211 "u]-m);for(uint j=0;j<" + ch + "u;++j)y[base+j*" + ins + "u]=exp(x0[base+j*" + ins +
212 "u]-m)/sum;",
213 outer * inner);
214 }
215 static const std::map<std::string, std::string> unary{{"Sqrt", "sqrt(v)"},
216 {"Exp", "exp(v)"},
217 {"Log", "log(v)"},
218 {"Sin", "sin(v)"},
219 {"Cos", "cos(v)"},
220 {"Relu", "max(v,0.0)"},
221 {"Sigmoid", "1.0/(1.0+exp(-v))"},
222 {"Tanh", "isnan(v)?v:tanh(clamp(v,-10.0,10.0))"},
223 {"Neg", "-v"},
224 {"Abs", "abs(v)"},
225 {"Reciprocal", "1.0/v"},
226 {"Floor", "floor(v)"},
227 {"Round", "roundEven(v)"},
228 {"Atan", "isinf(v)?sign(v)*1.5707963267948966:atan(v)"}};
229 if (n.op == "LeakyRelu" || unary.contains(n.op)) {
230 const float alpha = n.attrs.contains("alpha") ? n.attrs.at("alpha").real : 0.01f;
231 if (!std::isfinite(alpha)) throw Failure("Invalid LeakyRelu alpha");
232 if (compute)
233 return dispatchFloat(
234 *compute, {&x}, x.shape,
235 "float v=x0[i];y[i]=" + (n.op == "LeakyRelu" ? "v>=0.0?v:v*" + literal(alpha) : unary.at(n.op)) + ";");
236 if (n.op != "LeakyRelu" && n.op != "Reciprocal" && n.op != "Floor" && n.op != "Round" && n.op != "Atan")
237 return std::nullopt;
238 auto out = floats(x);
239 for (auto& v : out) {
240 if (n.op == "LeakyRelu")
241 v = v >= 0 ? v : alpha * v;
242 else if (n.op == "Reciprocal")
243 v = 1 / v;
244 else if (n.op == "Floor")
245 v = std::floor(v);
246 else if (n.op == "Atan")
247 v = std::atan(v);
248 else {
249 float lo = std::floor(v), f = v - lo;
250 v = lo + (f > 0.5f || (f == 0.5f && std::fmod(lo, 2.f) != 0));
251 }
252 }
253 return make(OnnxElement::Float32, x.shape, out);
254 }
255 return std::nullopt;
256}
257} // namespace eve::tensor::onnx_detail
float x
Definition AnimClip.cpp:738
std::map< std::string, std::vector< Key >, std::less<> > channels
std::string output
const std::string & s
glm::vec4 p[6]
EvpackChunkInput input
Definition Evpack.cpp:170
int cols
DiagnosticCode code
ShaderImageInput shape
glm::vec3 n
Definition Grass.cpp:63
int inputs
Definition GridGraph.cpp:23
std::uint32_t ab
double r
float v
std::int32_t c
size_t offset
bool required
std::uint64_t bytes
MeleePoint3 b
Definition MeleeHit.cpp:41
MeleePoint3 a
Definition MeleeHit.cpp:40
OnnxCompute * compute
float f
eve::Value literal
std::uint32_t count
std::string element
std::string body
uint32_t index
float m[16]
GPU execution boundary for native ONNX; retains no model and retains compiled resources for the lifet...
Definition OnnxCompute.h:25
std::vector< int64_t > ints(const RuntimeTensor &v)
Ints.
int64_t attr(const Node &n, const char *key, int64_t fallback)
Attr.
std::optional< RuntimeTensor > executeNumeric(const Node &, const std::vector< const RuntimeTensor * > &, OnnxCompute *)
Execute numeric.
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.
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.
std::vector< int64_t > broadcast(const std::vector< int64_t > &a, const std::vector< int64_t > &b)
Broadcast.
Synchronous GPU kernel request; all inputs are borrowed only until dispatch returns.
Definition OnnxCompute.h:12