载入中...
搜索中...
未找到
OnnxLstm.cpp
浏览该文件的文档.
4
5#include <algorithm>
6#include <cmath>
7#include <limits>
8
10namespace {
11template <class T>
12T checked(Result<T> r) {
13 if (!r.ok()) throw Failure(r.error()->message(), r.error()->code());
14 return std::move(r.value());
15}
16void shape(const RuntimeTensor& t, const std::vector<int64_t>& expected) {
17 if (t.shape != expected) throw Failure("LSTM input shape mismatch");
18}
19std::vector<float> optionalFloats(const std::vector<const RuntimeTensor*>& in, size_t i,
20 const std::vector<int64_t>& dims) {
21 if (i >= in.size() || !in[i]) return std::vector<float>(count(dims), 0);
22 shape(*in[i], dims);
23 return floats(*in[i]);
24}
25struct Weights {
26 const RuntimeTensor& tensor;
27 std::vector<float> scales;
28 std::vector<int64_t> zeros;
30 Weights(const RuntimeTensor& t, const RuntimeTensor& s, const RuntimeTensor& z, size_t d, size_t g, size_t r)
32 if (t.element != OnnxElement::Int8 && t.element != OnnxElement::UInt8)
33 throw Failure("LSTM weights must be 8-bit");
34 if (t.element != z.element || s.shape != z.shape ||
35 (s.shape != std::vector<int64_t>{static_cast<int64_t>(d)} &&
36 s.shape != std::vector<int64_t>{static_cast<int64_t>(d), static_cast<int64_t>(g)}))
37 throw Failure("LSTM quantization parameter shape mismatch");
38 shape(t, {static_cast<int64_t>(d), static_cast<int64_t>(r), static_cast<int64_t>(g)});
39 for (float v : scales)
40 if (!(v > 0) || !std::isfinite(v)) throw Failure("Invalid LSTM weight scale");
41 }
42 std::vector<float> multiply(const affine::QuantizedActivation& a, size_t batch, size_t direction,
43 OnnxCompute* compute) const {
44 std::vector<float> out(batch * gates);
45 const bool sign = tensor.element == OnnxElement::Int8;
46 if (compute) {
47 std::vector<int32_t> z(scales.size() == directions ? 1 : gates);
48 for (size_t i = 0; i < z.size(); ++i) z[i] = static_cast<int32_t>(zeros[direction * z.size() + i]);
49 ByteStorage activation(a.values);
50 ByteStorage accumulator(gpuMatmulResident(*compute, activation.buffer(), tensor.bytes.buffer(), false, sign,
51 batch, rows, gates, a.zeroPoint, z, direction * rows * gates));
52 std::vector<int32_t> sums(batch * gates);
53 std::memcpy(sums.data(), static_cast<const ByteStorage&>(accumulator).data(),
54 sums.size() * sizeof(int32_t));
55 for (size_t b = 0; b < batch; ++b)
56 for (size_t g = 0; g < gates; ++g)
57 out[b * gates + g] = static_cast<float>(sums[b * gates + g]) * a.scale *
58 scales[direction * z.size() + (z.size() == 1 ? 0 : g)];
59 return out;
60 }
61 // Per-output-channel affine zero points do not require unpacking a weight matrix.
62 for (size_t b = 0; b < batch; ++b)
63 for (size_t g = 0; g < gates; ++g) {
64 const size_t p = scales.size() == directions ? direction : direction * gates + g;
65 int64_t sum = 0;
66 for (size_t k = 0; k < rows; ++k) {
67 const int byte = tensor.bytes[(direction * rows + k) * gates + g];
68 const int w = sign && byte >= 128 ? byte - 256 : byte;
69 sum += static_cast<int64_t>(int(a.values[b * rows + k]) - a.zeroPoint) * (w - zeros[p]);
70 }
71 if (sum < INT32_MIN || sum > INT32_MAX) throw Failure("LSTM integer accumulator overflow");
72 out[b * gates + g] = static_cast<float>(sum) * a.scale * scales[p];
73 }
74 return out;
75 }
76};
77} // namespace
78std::vector<RuntimeTensor> executeQuantLstm(const Node& n, const std::vector<const RuntimeTensor*>& in,
80 const auto& xt = required(in, 0);
81 if (xt.shape.size() != 3) throw Failure("LSTM X must be [time,batch,input]");
82 const auto x = floats(xt);
83 const int64_t time = xt.shape[0], batch = xt.shape[1], input = xt.shape[2], hidden = attr(n, "hidden_size", 0);
84 const auto direction = n.attrs.contains("direction") ? n.attrs.at("direction").text : "forward";
85 const int64_t dirs = direction == "bidirectional" ? 2 : 1;
86 if (time <= 0 || batch <= 0 || input <= 0 || hidden <= 0 || hidden > 16384)
87 throw Failure("Invalid LSTM dimensions");
88 const size_t gates = 4 * hidden;
89 const size_t projectionSize = count({time, batch, static_cast<int64_t>(gates)});
90 if (projectionSize > 128u * 1024u * 1024u) throw Failure("LSTM projection exceeds memory limit");
91 Weights w(required(in, 1), required(in, 8), required(in, 9), dirs, gates, input);
92 Weights r(required(in, 2), required(in, 10), required(in, 11), dirs, gates, hidden);
93 auto biases = optionalFloats(in, 3, {dirs, 8 * hidden});
94 auto h = optionalFloats(in, 5, {dirs, batch, hidden}), c = optionalFloats(in, 6, {dirs, batch, hidden});
95 auto p = optionalFloats(in, 7, {dirs, 3 * hidden});
96 std::vector<int64_t> lengths(batch, time);
97 if (in.size() > 4 && in[4]) {
98 if (in[4]->element != OnnxElement::Int32) throw Failure("LSTM sequence lengths require int32");
99 shape(*in[4], {batch});
100 lengths = ints(*in[4]);
101 }
102 for (auto len : lengths)
103 if (len < 0 || len > time) throw Failure("Invalid LSTM sequence length");
104 const int64_t coupled = attr(n, "input_forget", 0);
105 if (coupled != 0 && coupled != 1) throw Failure("Invalid coupled gate flag");
106 const float clip = n.attrs.contains("clip") ? n.attrs.at("clip").real : std::numeric_limits<float>::max();
107 if (!(clip > 0) || !std::isfinite(clip)) throw Failure("Invalid LSTM clipping threshold");
108 auto sigmoid = [clip](float z) { return 1.f / (1.f + std::exp(-std::clamp(z, -clip, clip))); };
109 auto tanh = [clip](float z) { return std::tanh(std::clamp(z, -clip, clip)); };
110 std::vector<float> y(count({time, dirs, batch, hidden}), 0);
111 const auto qx = checked(affine::dynamicQuantize(x));
112 for (int64_t d = 0; d < dirs; ++d) {
113 const auto projected = w.multiply(qx, time * batch, d, compute);
114 const bool reverse = direction == "reverse" || d == 1;
115 for (int64_t step = 0; step < time; ++step) {
116 const auto qh = checked(affine::dynamicQuantize(std::span(h).subspan(d * batch * hidden, batch * hidden)));
117 const auto recurrent = r.multiply(qh, batch, d, compute);
118 for (int64_t b = 0; b < batch; ++b) {
119 if (step >= lengths[b]) continue;
120 const int64_t t = reverse ? lengths[b] - 1 - step : step;
121 const size_t state = (d * batch + b) * hidden;
122 for (int64_t j = 0; j < hidden; ++j) {
123 auto gate = [&](int g) {
124 const size_t channel = g * hidden + j;
125 return projected[(t * batch + b) * gates + channel] + recurrent[b * gates + channel] +
126 biases[d * 8 * hidden + channel] + biases[d * 8 * hidden + gates + channel];
127 };
128 const float i = sigmoid(gate(0) + p[d * 3 * hidden + j] * c[state + j]);
129 const float f =
130 coupled ? 1 - i : sigmoid(gate(2) + p[d * 3 * hidden + 2 * hidden + j] * c[state + j]);
131 c[state + j] = f * c[state + j] + i * tanh(gate(3));
132 const float o = sigmoid(gate(1) + p[d * 3 * hidden + hidden + j] * c[state + j]);
133 h[state + j] = o * tanh(c[state + j]);
134 y[((t * dirs + d) * batch + b) * hidden + j] = h[state + j];
135 }
136 }
137 }
138 }
139 return {make(OnnxElement::Float32, {time, dirs, batch, hidden}, y),
140 make(OnnxElement::Float32, {dirs, batch, hidden}, h), make(OnnxElement::Float32, {dirs, batch, hidden}, c)};
141}
142} // namespace eve::tensor::onnx_detail
float w
Definition AnimClip.cpp:738
float y
Definition AnimClip.cpp:738
float x
Definition AnimClip.cpp:738
float z
Definition AnimClip.cpp:738
const std::string & s
Vec3 projected
Definition CaveMesh.cpp:122
glm::vec4 p[6]
EvpackChunkInput input
Definition Evpack.cpp:170
int rows
tensor::Graph g
Definition GpuGraph.cpp:7
ShaderImageInput shape
glm::vec4 clip
glm::vec3 n
Definition Grass.cpp:63
double r
float v
std::int32_t c
int h
bool required
MeleePoint3 b
Definition MeleeHit.cpp:41
MeleePoint3 a
Definition MeleeHit.cpp:40
std::vector< float > scales
Definition OnnxLstm.cpp:27
size_t directions
Definition OnnxLstm.cpp:29
const RuntimeTensor & tensor
Definition OnnxLstm.cpp:26
std::vector< int64_t > zeros
Definition OnnxLstm.cpp:28
size_t gates
Definition OnnxLstm.cpp:29
OnnxCompute * compute
float f
float d
float t
RoadLaneDirection direction
std::uint32_t count
std::string element
float step
Definition TreeMesh.cpp:314
float qx
GPU execution boundary for native ONNX; retains no model and retains compiled resources for the lifet...
Definition OnnxCompute.h:25
Result< QuantizedActivation > dynamicQuantize(std::span< const float > input)
Quantize finite FP32 activations with ONNX DynamicQuantizeLinear semantics.
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.
std::vector< float > floats(const RuntimeTensor &v)
Floats.
OnnxBuffer gpuMatmulResident(OnnxCompute &d, OnnxBuffer a, OnnxBuffer b, bool as, bool bs, size_t m, size_t k, size_t n, int az, std::span< const int32_t > bz, size_t bOffset)
Gpu matmul resident.
std::vector< RuntimeTensor > executeQuantLstm(const Node &node, const std::vector< const RuntimeTensor * > &inputs, OnnxCompute *compute=nullptr)
Execute quant lstm.
Definition OnnxLstm.cpp:78