载入中...
搜索中...
未找到
SmrMeshRetNet.cpp
浏览该文件的文档.
2
3#include "common/Diagnostic.h"
4#include "tensor/Tensor.h"
5
6#include <algorithm>
7#include <cmath>
8#include <cstdint>
9#include <random>
10#include <utility>
11
12namespace eve::animation {
13namespace {
14
15void gemm(const float* a, const float* b, float* c, int m, int k, int n) {
16 for (int i = 0; i < m; ++i)
17 for (int j = 0; j < n; ++j) {
18 float sum = 0.f;
19 for (int t = 0; t < k; ++t) sum += a[i * k + t] * b[t * n + j];
20 c[i * n + j] = sum;
21 }
22}
23
24void addBias(float* x, const float* bias, int rows, int cols) {
25 for (int r = 0; r < rows; ++r)
26 for (int c = 0; c < cols; ++c) x[r * cols + c] += bias[c];
27}
28
29void gelu(float* x, int n) {
30 for (int i = 0; i < n; ++i) {
31 const float v = x[i];
32 x[i] = 0.5f * v * (1.f + std::tanh(0.7978845608f * (v + 0.044715f * v * v * v)));
33 }
34}
35
36void layernorm(float* x, int rows, int cols, const float* g, const float* b) {
37 for (int r = 0; r < rows; ++r) {
38 float* row = x + r * cols;
39 float mean = 0.f;
40 for (int c = 0; c < cols; ++c) mean += row[c];
41 mean /= static_cast<float>(cols);
42 float var = 0.f;
43 for (int c = 0; c < cols; ++c) {
44 const float d = row[c] - mean;
45 var += d * d;
46 }
47 var = 1.f / std::sqrt(var / static_cast<float>(cols) + 1e-5f);
48 for (int c = 0; c < cols; ++c) row[c] = (row[c] - mean) * var * g[c] + b[c];
49 }
50}
51
52void softmaxRows(float* x, int rows, int cols) {
53 for (int r = 0; r < rows; ++r) {
54 float* row = x + r * cols;
55 float mx = row[0];
56 for (int c = 1; c < cols; ++c) mx = std::max(mx, row[c]);
57 float sum = 0.f;
58 for (int c = 0; c < cols; ++c) {
59 row[c] = std::exp(row[c] - mx);
60 sum += row[c];
61 }
62 const float inv = 1.f / std::max(sum, 1e-8f);
63 for (int c = 0; c < cols; ++c) row[c] *= inv;
64 }
65}
66
67std::vector<float> xavier(std::mt19937& rng, int rows, int cols) {
68 std::normal_distribution<float> dist(0.f, std::sqrt(2.f / static_cast<float>(rows + cols)));
69 std::vector<float> out(static_cast<size_t>(rows * cols));
70 for (float& v : out) v = dist(rng);
71 return out;
72}
73
74std::vector<float> zeros(int n) { return std::vector<float>(static_cast<size_t>(n), 0.f); }
75std::vector<float> ones(int n) { return std::vector<float>(static_cast<size_t>(n), 1.f); }
76
77std::vector<float> mlp2(const float* x, int rows, int inDim, int mid, int outDim, const std::vector<float>& w1,
78 const std::vector<float>& b1, const std::vector<float>& w2, const std::vector<float>& b2) {
79 std::vector<float> h(static_cast<size_t>(rows * mid));
80 gemm(x, w1.data(), h.data(), rows, inDim, mid);
81 addBias(h.data(), b1.data(), rows, mid);
82 gelu(h.data(), rows * mid);
83 std::vector<float> y(static_cast<size_t>(rows * outDim));
84 gemm(h.data(), w2.data(), y.data(), rows, mid, outDim);
85 addBias(y.data(), b2.data(), rows, outDim);
86 return y;
87}
88
89std::vector<float> maxPool(const std::vector<float>& x, int groups, int count, int dim) {
90 std::vector<float> out(static_cast<size_t>(groups * dim), -1e30f);
91 for (int g = 0; g < groups; ++g)
92 for (int i = 0; i < count; ++i)
93 for (int d = 0; d < dim; ++d) {
94 const float v = x[static_cast<size_t>((g * count + i) * dim + d)];
95 float& slot = out[static_cast<size_t>(g * dim + d)];
96 slot = std::max(slot, v);
97 }
98 for (float& v : out)
99 if (v < -1e29f) v = 0.f;
100 return out;
101}
102
103void transformerLayer(std::vector<float>& x, int tokens, int dim, int heads, const std::vector<float>& wq,
104 const std::vector<float>& wk, const std::vector<float>& wv, const std::vector<float>& wo,
105 const std::vector<float>& w1, const std::vector<float>& b1, const std::vector<float>& w2,
106 const std::vector<float>& b2, const std::vector<float>& ln1g, const std::vector<float>& ln1b,
107 const std::vector<float>& ln2g, const std::vector<float>& ln2b, const std::vector<float>* memory,
108 int memoryTokens) {
109 const int headDim = std::max(1, dim / std::max(heads, 1));
110 const int srcTokens = memory ? memoryTokens : tokens;
111 std::vector<float> residual = x;
112 layernorm(x.data(), tokens, dim, ln1g.data(), ln1b.data());
113 std::vector<float> q(static_cast<size_t>(tokens * dim));
114 std::vector<float> k(static_cast<size_t>(srcTokens * dim));
115 std::vector<float> v(static_cast<size_t>(srcTokens * dim));
116 gemm(x.data(), wq.data(), q.data(), tokens, dim, dim);
117 const float* src = memory ? memory->data() : x.data();
118 gemm(src, wk.data(), k.data(), srcTokens, dim, dim);
119 gemm(src, wv.data(), v.data(), srcTokens, dim, dim);
120 std::vector<float> ctx(static_cast<size_t>(tokens * dim), 0.f);
121 const float scale = 1.f / std::sqrt(static_cast<float>(headDim));
122 for (int h = 0; h < heads; ++h) {
123 std::vector<float> scores(static_cast<size_t>(tokens * srcTokens));
124 for (int t = 0; t < tokens; ++t)
125 for (int s = 0; s < srcTokens; ++s) {
126 float sum = 0.f;
127 for (int d = 0; d < headDim; ++d)
128 sum += q[static_cast<size_t>(t * dim + h * headDim + d)] *
129 k[static_cast<size_t>(s * dim + h * headDim + d)];
130 scores[static_cast<size_t>(t * srcTokens + s)] = sum * scale;
131 }
132 softmaxRows(scores.data(), tokens, srcTokens);
133 for (int t = 0; t < tokens; ++t)
134 for (int d = 0; d < headDim; ++d) {
135 float sum = 0.f;
136 for (int s = 0; s < srcTokens; ++s)
137 sum += scores[static_cast<size_t>(t * srcTokens + s)] *
138 v[static_cast<size_t>(s * dim + h * headDim + d)];
139 ctx[static_cast<size_t>(t * dim + h * headDim + d)] = sum;
140 }
141 }
142 std::vector<float> projected(static_cast<size_t>(tokens * dim));
143 gemm(ctx.data(), wo.data(), projected.data(), tokens, dim, dim);
144 for (size_t i = 0; i < x.size(); ++i) x[i] = residual[i] + projected[i];
145 residual = x;
146 layernorm(x.data(), tokens, dim, ln2g.data(), ln2b.data());
147 auto ff = mlp2(x.data(), tokens, dim, static_cast<int>(b1.size()), dim, w1, b1, w2, b2);
148 for (size_t i = 0; i < x.size(); ++i) x[i] = residual[i] + ff[i];
149}
150
151} // namespace
152
153SmrMeshRetNet::SmrMeshRetNet(SmrMeshRetConfig config) : config_(config) { initWeights(); }
154
155void SmrMeshRetNet::initWeights() {
156 std::mt19937 rng(static_cast<uint32_t>(config_.seed));
157 const int d = config_.latentDim;
158 geomW1_ = xavier(rng, 7, d);
159 geomB1_ = zeros(d);
160 geomW2_ = xavier(rng, d, d);
161 geomB2_ = zeros(d);
162 dmiW1_ = xavier(rng, 10, d);
163 dmiB1_ = zeros(d);
164 dmiW2_ = xavier(rng, d, d);
165 dmiB2_ = zeros(d);
166 motionW_ = xavier(rng, 6, d);
167 motionB_ = zeros(d);
168 fuseW_ = xavier(rng, d * 3, d);
169 fuseB_ = zeros(d);
170 outW_ = xavier(rng, d, 6);
171 outB_ = zeros(6);
172 auto fill = [&](std::vector<LayerWeights>& layers) {
173 layers.assign(static_cast<size_t>(config_.numLayers), {});
174 for (auto& layer : layers) {
175 layer.wq = xavier(rng, d, d);
176 layer.wk = xavier(rng, d, d);
177 layer.wv = xavier(rng, d, d);
178 layer.wo = xavier(rng, d, d);
179 layer.w1 = xavier(rng, d, config_.ffSize);
180 layer.b1 = zeros(config_.ffSize);
181 layer.w2 = xavier(rng, config_.ffSize, d);
182 layer.b2 = zeros(d);
183 layer.ln1g = ones(d);
184 layer.ln1b = zeros(d);
185 layer.ln2g = ones(d);
186 layer.ln2b = zeros(d);
187 }
188 };
189 fill(enc_);
190 fill(dec_);
191 tensor::Tensor probe(d);
192 probe.fill(0.f);
193 (void)probe;
194}
195
197 if (features.frames <= 0 || features.joints <= 0) {
198 return Result<std::vector<float>>::failure(
199 Diagnostic::error(DiagnosticCode::InvalidArgument, "SmrMeshRetNet.forward: empty features"));
200 }
201 const int d = config_.latentDim;
202 const int T = features.frames;
203 const int J = features.joints;
204 const int S = std::max(features.sensors, 1);
205 const int P = std::max(features.pairs, 1);
206 const int heads = std::max(1, config_.numHeads);
207
208 auto srcGeom = mlp2(features.sourceGeom.data(), S, 7, d, d, geomW1_, geomB1_, geomW2_, geomB2_);
209 auto tgtGeom = mlp2(features.targetGeom.data(), S, 7, d, d, geomW1_, geomB1_, geomW2_, geomB2_);
210 auto srcPool = maxPool(srcGeom, 1, S, d);
211 auto tgtPool = maxPool(tgtGeom, 1, S, d);
212
213 std::vector<float> temporal(static_cast<size_t>(T * d), 0.f);
214 for (int t = 0; t < T; ++t) {
215 const float* dmiFrame = features.sourceDmi.data() + static_cast<size_t>(t) * static_cast<size_t>(P) * 10u;
216 auto dmiEnc = mlp2(dmiFrame, P, 10, d, d, dmiW1_, dmiB1_, dmiW2_, dmiB2_);
217 auto dmiPool = maxPool(dmiEnc, 1, P, d);
218 std::vector<float> meanRot(6, 0.f);
219 for (int j = 0; j < J; ++j) {
220 const float* r = features.sourceRot6d.data() + static_cast<size_t>((t * J + j) * 6);
221 for (int k = 0; k < 6; ++k) meanRot[static_cast<size_t>(k)] += r[k];
222 }
223 for (float& v : meanRot) v /= static_cast<float>(std::max(J, 1));
224 std::vector<float> motion(static_cast<size_t>(d), 0.f);
225 gemm(meanRot.data(), motionW_.data(), motion.data(), 1, 6, d);
226 addBias(motion.data(), motionB_.data(), 1, d);
227 std::vector<float> fusedIn(static_cast<size_t>(d * 3));
228 for (int i = 0; i < d; ++i) {
229 fusedIn[static_cast<size_t>(i)] = dmiPool[static_cast<size_t>(i)];
230 fusedIn[static_cast<size_t>(d + i)] = motion[static_cast<size_t>(i)];
231 fusedIn[static_cast<size_t>(2 * d + i)] = srcPool[static_cast<size_t>(i)];
232 }
233 std::vector<float> fused(static_cast<size_t>(d));
234 gemm(fusedIn.data(), fuseW_.data(), fused.data(), 1, d * 3, d);
235 addBias(fused.data(), fuseB_.data(), 1, d);
236 for (int i = 0; i < d; ++i) temporal[static_cast<size_t>(t * d + i)] = fused[static_cast<size_t>(i)];
237 }
238
239 std::vector<float> memory = temporal;
240 for (const auto& layer : enc_)
241 transformerLayer(memory, T, d, heads, layer.wq, layer.wk, layer.wv, layer.wo, layer.w1, layer.b1, layer.w2,
242 layer.b2, layer.ln1g, layer.ln1b, layer.ln2g, layer.ln2b, nullptr, 0);
243
244 std::vector<float> queries(static_cast<size_t>(T * d));
245 for (int t = 0; t < T; ++t)
246 for (int i = 0; i < d; ++i)
247 queries[static_cast<size_t>(t * d + i)] =
248 tgtPool[static_cast<size_t>(i)] + temporal[static_cast<size_t>(t * d + i)];
249 for (const auto& layer : dec_)
250 transformerLayer(queries, T, d, heads, layer.wq, layer.wk, layer.wv, layer.wo, layer.w1, layer.b1, layer.w2,
251 layer.b2, layer.ln1g, layer.ln1b, layer.ln2g, layer.ln2b, &memory, T);
252
253 std::vector<float> out(static_cast<size_t>(T * J * 6));
254 for (int t = 0; t < T; ++t) {
255 std::vector<float> delta(6, 0.f);
256 gemm(queries.data() + static_cast<size_t>(t * d), outW_.data(), delta.data(), 1, d, 6);
257 addBias(delta.data(), outB_.data(), 1, 6);
258 for (int j = 0; j < J; ++j) {
259 const float* src = features.sourceRot6d.data() + static_cast<size_t>((t * J + j) * 6);
260 float* dst = out.data() + static_cast<size_t>((t * J + j) * 6);
261 for (int k = 0; k < 6; ++k) dst[k] = src[k] + 0.05f * delta[static_cast<size_t>(k)];
262 }
263 }
264 return Result<std::vector<float>>::success(std::move(out));
265}
266
267} // namespace eve::animation
float y
Definition AnimClip.cpp:738
float x
Definition AnimClip.cpp:738
int mid
Definition AnimSmr.cpp:120
const std::string & s
Vec3 projected
Definition CaveMesh.cpp:122
Stable, structured diagnostics shared by engine modules.
int rows
int cols
tensor::Graph g
Definition GpuGraph.cpp:7
std::unordered_map< const Graphics *, std::shared_ptr< Lifetime > > tokens
vk::UniqueDeviceMemory memory
glm::vec3 n
Definition Grass.cpp:63
std::array< double, 10 > q
double r
float v
std::int32_t c
int h
std::array< float, 3 > scale
MeleePoint3 b
Definition MeleeHit.cpp:41
MeleePoint3 a
Definition MeleeHit.cpp:40
uint32_t groups
Definition OnnxGpgpu.cpp:39
std::vector< int64_t > zeros
Definition OnnxLstm.cpp:28
TileLayer * layer
float d
float t
float bias
std::uint32_t count
float m[16]
static Diagnostic error(DiagnosticCode code, std::string message, std::string path={}, DiagnosticDetails details={}, std::string source={})
Construct an error diagnostic with the standard error severity.
Definition Diagnostic.h:125
Move-only operation result carrying either a value or Status.
Definition Result.h:155
SmrMeshRetNet(SmrMeshRetConfig config={})
Constructs a SmrMeshRetNet.
Result< std::vector< float > > forward(const SmrFeatureBatch &features) const
Owning target rot6d [T*J*6], or structured failure.
void layernorm(const float *in, int rows, int cols, const float *scale, const float *bias, float eps, float *out)
Layernorm.
WidgetDesc row(std::vector< WidgetDesc > children, std::string id)
Horizontal elastic layout row.
Definition Widget.cpp:679
Packed MeshRet-style features for one retarget inference window.
Definition SmrFeatures.h:15
std::vector< float > targetGeom
[S*7]
Definition SmrFeatures.h:23
std::vector< float > sourceRot6d
[T*J*6]
Definition SmrFeatures.h:21
std::vector< float > sourceDmi
[T*P*10]
Definition SmrFeatures.h:24
std::vector< float > sourceGeom
[S*7]
Definition SmrFeatures.h:22
Hyperparameters for the builtin MeshRet-style graph.