载入中...
搜索中...
未找到
KernelGenWgsl.cpp
浏览该文件的文档.
1#include "common/Exception.h"
2#include "tensor/CpuKernels.h"
4#include "tensor/Quant.h"
5
6#include <algorithm>
7#include <cmath>
8#include <sstream>
9
10namespace eve::tensor {
11namespace wgsl_detail {
12
13
14std::string header(int localX, int localY) {
15 return "const workgroupX = " + std::to_string(localX) + "u;\nconst workgroupY = " + std::to_string(localY) + "u;\n";
16}
17std::string bufferDecl(int binding, const char *name) {
18 return "@group(0) @binding(" + std::to_string(binding) + ") var<storage, read_write> " + name + ": array<f32>;\n";
19}
20std::string bufferDeclUint(int binding, const char *name) {
21 return "@group(0) @binding(" + std::to_string(binding) + ") var<storage, read_write> " + name + ": array<u32>;\n";
22}
23
29std::string emitQuantizedBVal(DType dt, int group, const char *bufName = "b") {
30 std::ostringstream os;
31 const int g = group > 0 ? group : 1;
32 switch (dt) {
33 case DType::Int8:
34 os << "fn bval(idx: u32) -> f32 {\n"
35 << " var word: u32 = " << bufName << "[idx >> 2u];\n"
36 << " var byte: u32 = (word >> ((idx & 3u) * 8u)) & 0xFFu;\n"
37 << " var v: i32 = i32(byte);\n"
38 << " if (v >= 128) { v -= 256; }\n"
39 << " return f32(v) * bs[idx / " << g << "u];\n"
40 << "}\n";
41 break;
42 case DType::Int4:
43 os << "fn bval(idx: u32) -> f32 {\n"
44 << " var word: u32 = " << bufName << "[idx >> 3u];\n"
45 << " var byte: u32 = (word >> (((idx >> 1u) & 3u) * 8u)) & 0xFFu;\n"
46 << " var nib: u32 = select(byte >> 4u, byte & 0xFu, (idx & 1u) == 0u);\n"
47 << " var v: i32 = i32(nib);\n"
48 << " if (v >= 8) { v -= 16; }\n"
49 << " return f32(v) * bs[idx / " << g << "u];\n"
50 << "}\n";
51 break;
52 case DType::Fp16:
53 os << "fn bval(idx: u32) -> f32 {\n"
54 << " var word: u32 = " << bufName << "[idx >> 1u];\n"
55 << " var hb: u32 = (word >> ((idx & 1u) * 16u)) & 0xFFFFu;\n"
56 << " return unpack2x16float(hb).x;\n"
57 << "}\n";
58 break;
59 case DType::Fp8E4M3:
60 os << "fn bval(idx: u32) -> f32 {\n"
61 << " var word: u32 = " << bufName << "[idx >> 2u];\n"
62 << " var byte: u32 = (word >> ((idx & 3u) * 8u)) & 0xFFu;\n"
63 << " var s: i32 = select(1, -1, (byte & 0x80u) != 0u);\n"
64 << " var e: i32 = i32((byte >> 3u) & 0xFu);\n"
65 << " var m: i32 = i32(byte & 0x7u);\n"
66 << " var v: f32 = select(exp2(f32(e - 7)) * (1.0 + f32(m) / 8.0),\n"
67 << " exp2(-6.0) * f32(m) / 8.0, e == 0);\n"
68 << " return f32(s) * v * bs[idx / " << g << "u];\n"
69 << "}\n";
70 break;
71 case DType::Fp4E2M1:
72 os << "fn bval(idx: u32) -> f32 {\n"
73 << " var word: u32 = " << bufName << "[idx >> 3u];\n"
74 << " var byte: u32 = (word >> (((idx >> 1u) & 3u) * 8u)) & 0xFFu;\n"
75 << " var nib: u32 = select(byte >> 4u, byte & 0xFu, (idx & 1u) == 0u);\n"
76 << " var s: i32 = select(1, -1, (nib & 8u) != 0u);\n"
77 << " var e: i32 = i32((nib >> 1u) & 3u);\n"
78 << " var m: i32 = i32(nib & 1u);\n"
79 << " var v: f32 = select(exp2(f32(e - 1)) * (1.0 + 0.5 * f32(m)),\n"
80 << " 0.5 * f32(m), e == 0);\n"
81 << " return f32(s) * v * bs[idx / " << g << "u];\n"
82 << "}\n";
83 break;
84 default: break;
85 }
86 return os.str();
87}
88
89std::string pushConstant() { return ""; }
90
91int groupsFor(int count) { return (count + kLocalSize - 1) / kLocalSize; }
92
93std::string scalarStr(float v) {
94 if (v == int(v) && std::fabs(v) < 1e9f) return std::to_string(int(v)) + ".0";
95 std::ostringstream os;
96 os << v << "f";
97 return os.str();
98}
99
106 const Graph &graph;
108 const std::vector<std::string> &indexExprs; // per group input
109 std::string rootVar;
110 std::ostringstream &os;
111 int temp = 0;
112 std::string biasIndexExpr; // e.g. "jj" (matmul) / "f" (conv), empty = none
113
114 std::string operand(int nodeId) {
115 if (nodeId == group.biasNode && !biasIndexExpr.empty()) return "bias[" + biasIndexExpr + "]";
116 const auto it = std::find(group.inputs.begin(), group.inputs.end(), nodeId);
117 if (it != group.inputs.end()) {
118 const size_t idx = static_cast<size_t>(it - group.inputs.begin());
119 return "a" + std::to_string(idx) + "[" + indexExprs[idx] + "]";
120 }
121 return vars[static_cast<size_t>(nodeId)];
122 }
123
124 std::vector<std::string> vars;
125
126 std::string emit(int nodeId) {
127 const GraphNode &nd = graph.node(nodeId);
128 const std::string x =
129 (nd.in0 == group.nodes.front() && rootVar != "") ? rootVar : (nd.in0 >= 0 ? operand(nd.in0) : rootVar);
130 std::string expr;
131 switch (nd.type) {
132 case OpType::Add: expr = "(" + x + " + " + operand(nd.in1) + ")"; break;
133 case OpType::Sub: expr = "(" + x + " - " + operand(nd.in1) + ")"; break;
134 case OpType::Multiply: expr = "(" + x + " * " + operand(nd.in1) + ")"; break;
135 case OpType::Divide: expr = "(" + x + " / " + operand(nd.in1) + ")"; break;
136 case OpType::AddScalar: expr = "(" + x + " + " + scalarStr(nd.s0) + ")"; break;
137 case OpType::SubScalar: expr = "(" + x + " - " + scalarStr(nd.s0) + ")"; break;
138 case OpType::MulScalar: expr = "(" + x + " * " + scalarStr(nd.s0) + ")"; break;
139 case OpType::DivScalar: expr = "(" + x + " / " + scalarStr(nd.s0) + ")"; break;
140 case OpType::PowScalar: expr = "pow(" + x + ", " + scalarStr(nd.s0) + ")"; break;
141 case OpType::Neg: expr = "(-" + x + ")"; break;
142 case OpType::Abs: expr = "abs(" + x + ")"; break;
143 case OpType::Sqrt: expr = "sqrt(" + x + ")"; break;
144 case OpType::Exp: expr = "exp(" + x + ")"; break;
145 case OpType::Log: expr = "log(" + x + ")"; break;
146 case OpType::Sin: expr = "sin(" + x + ")"; break;
147 case OpType::Cos: expr = "cos(" + x + ")"; break;
148 case OpType::Tanh: expr = "tanh(" + x + ")"; break;
149 case OpType::Relu: expr = "max(" + x + ", 0.0)"; break;
150 case OpType::Sigmoid: expr = "(1.0 / (1.0 + exp(-" + x + ")))"; break;
151 case OpType::Gelu:
152 expr = "(0.5 * " + x + " * (1.0 + tanh(0.7978845608028654 * (" + x + " + 0.044715 * " + x + " * " + x +
153 " * " + x + "))))";
154 break;
155 case OpType::Silu: expr = "(" + x + " / (1.0 + exp(-" + x + ")))"; break;
156 case OpType::Clamp: {
157 float lo = nd.s0, hi = nd.s1;
158 if (lo > hi) std::swap(lo, hi);
159 expr = "clamp(" + x + ", " + scalarStr(lo) + ", " + scalarStr(hi) + ")";
160 break;
161 }
162 case OpType::MaximumScalar: expr = "max(" + x + ", " + scalarStr(nd.s0) + ")"; break;
163 case OpType::MinimumScalar: expr = "min(" + x + ", " + scalarStr(nd.s0) + ")"; break;
164 case OpType::Where:
165 expr = "select(" + operand(nd.in2) + ", " + operand(nd.in1) + ", " + operand(nd.in0) + " > 0.5)";
166 break;
167 default: return "";
168 }
169 const std::string var = "t" + std::to_string(temp++);
170 os << " var " << var << ": f32 = " << expr << ";\n";
171 if (vars.size() <= static_cast<size_t>(nodeId)) vars.resize(static_cast<size_t>(nodeId) + 1);
172 vars[static_cast<size_t>(nodeId)] = var;
173 return var;
174 }
175};
176
182void emitInputIndexExprs(std::ostringstream &os, const Graph &g, const FusedGroup &grp,
183 std::vector<std::string> &indexExprs) {
184 const GraphNode &on = g.node(grp.outputNode);
185 const int rank = on.rank;
186 std::vector<int> S(static_cast<size_t>(rank), 1);
187 if (rank > 0) {
188 S[static_cast<size_t>(rank - 1)] = 1;
189 for (int k = rank - 2; k >= 0; --k) S[static_cast<size_t>(k)] = S[static_cast<size_t>(k + 1)] * on.dims[k + 1];
190 }
191 indexExprs.resize(grp.inputs.size());
192 for (size_t k = 0; k < grp.inputs.size(); ++k) {
193 const GraphNode &inN = g.node(grp.inputs[k]);
194 bool identical = inN.rank == rank;
195 if (identical) {
196 for (int d = 0; d < rank; ++d)
197 if (inN.dims[d] != on.dims[d]) {
198 identical = false;
199 break;
200 }
201 }
202 if (identical) {
203 indexExprs[k] = "i_";
204 continue;
205 }
206 const int pad = rank - inN.rank;
207 std::vector<int> stride(static_cast<size_t>(rank), 0);
208 for (int d = 0; d < rank; ++d) {
209 const int dim = d < pad ? 1 : inN.dims[d - pad];
210 if (dim == 1) continue;
211 int s = 1;
212 for (int t = d + 1; t < rank; ++t) {
213 const int dt = t < pad ? 1 : inN.dims[t - pad];
214 if (dt != 1) s *= dt;
215 }
216 stride[static_cast<size_t>(d)] = s;
217 }
218 const std::string var = "idx" + std::to_string(k);
219 os << " var " << var << ": u32 = 0u;\n";
220 for (int d = 0; d < rank; ++d) {
221 if (stride[static_cast<size_t>(d)] == 0) continue;
222 os << " " << var << " += ((i_ / " << S[static_cast<size_t>(d)] << "u) % " << on.dims[d] << "u) * "
223 << stride[static_cast<size_t>(d)] << "u;\n";
224 }
225 indexExprs[k] = var;
226 }
227}
228
229void genElementwise(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
230 if (grp.inputs.size() > size_t(kMaxKernelBindings - 1))
231 throw eve::Exception("Tensor WGSL: unsupported kernel variant or binding count");
232 const GraphNode &on = g.node(grp.outputNode);
233 std::ostringstream os;
234 os << header(kLocalSize);
235 for (size_t k = 0; k < grp.inputs.size(); ++k) os << bufferDecl(int(k), ("a" + std::to_string(k)).c_str());
236 os << bufferDecl(int(grp.inputs.size()), "o");
237 os << pushConstant();
238 os << "@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
239 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) {\n";
240 os << " var i_: u32 = globalId.x;\n";
241 os << " if (i_ >= " << on.size << "u) { return; }\n";
242 std::vector<std::string> indexExprs;
243 emitInputIndexExprs(os, g, grp, indexExprs);
244 ChainContext ctx{g, grp, indexExprs, "", os, 0, ""};
245 for (int u : grp.nodes) ctx.emit(u);
246 os << " o[i_] = " << ctx.vars[static_cast<size_t>(grp.outputNode)] << ";\n";
247 os << "}\n";
248 out.pass1.clear();
249 out.pass2 = os.str();
250 out.groupsX2 = groupsFor(on.size);
251 out.inputCount = int(grp.inputs.size());
252 return;
253}
254
256void genMatMul(const Graph &g, const FusedGroup &grp, bool tiled, KernelSpec &out) {
257 const GraphNode &mm = g.node(grp.nodes.front());
258 const GraphNode &A = g.node(mm.in0);
259 const GraphNode &B = g.node(mm.in1);
260 const bool batched = mm.rank == 3;
261 const int batch = batched ? A.dims[0] : 1;
262 const int m = A.dims[batched ? 1 : 0];
263 const int k = A.dims[batched ? 2 : 1];
264 const int n = B.dims[batched ? 2 : 1];
265 const bool hasBias = grp.biasNode >= 0;
266 const bool bQuant = q::isQuantDType(static_cast<DType>(B.dtype)) && !B.constBytes.empty();
267 if (bQuant && tiled)
268 throw eve::Exception("Tensor WGSL: unsupported kernel variant or binding count"); // quantized weights use the
269 // naive variant only
270 const int scalesBinding = hasBias ? 3 : 2;
271 const int bindingOut = bQuant ? (hasBias ? 4 : 3) : (hasBias ? 3 : 2);
272
273 std::ostringstream os;
274 if (!tiled) {
275 os << header(kLocalSize);
276 } else {
277 os << header(16, 16);
278 }
279 os << bufferDecl(0, "a");
280 if (bQuant)
281 os << bufferDeclUint(1, "b");
282 else
283 os << bufferDecl(1, "b");
284 if (hasBias) os << bufferDecl(2, "bias");
285 if (bQuant && B.dtype != static_cast<int>(DType::Fp16)) os << bufferDecl(scalesBinding, "bs");
286 os << bufferDecl(bindingOut, "o");
287 os << pushConstant();
288 if (bQuant) os << emitQuantizedBVal(static_cast<DType>(B.dtype), B.qGroup);
289
290 std::vector<std::string> indexExprs(grp.inputs.size(), "i_");
291 ChainContext ctx{g, grp, indexExprs, "r", os, 0, ""};
292 ctx.biasIndexExpr = "jj";
293
294 if (!tiled) {
295 os << "@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
296 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) "
297 "{\n";
298 os << " var i_: u32 = globalId.x;\n";
299 os << " if (i_ >= " << batch * m * n << "u) { return; }\n";
300 if (batched) {
301 os << " var bb: u32 = i_ / " << m * n << "u;\n";
302 os << " var rem: u32 = i_ % " << m * n << "u;\n";
303 os << " var ii: u32 = rem / " << n << "u;\n";
304 os << " var jj: u32 = rem % " << n << "u;\n";
305 } else {
306 os << " var ii: u32 = i_ / " << n << "u;\n";
307 os << " var jj: u32 = i_ % " << n << "u;\n";
308 }
309 os << " var r: f32 = 0.0;\n";
310 os << " for (var t: u32 = 0u; t < " << k << "u; t++) {\n";
311 if (batched) {
312 os << " r += a[bb * " << m * k << "u + ii * " << k << "u + t] * "
313 << (bQuant ? "bval(bb * " + std::to_string(k * n) + "u + t * " + std::to_string(n) + "u + jj)"
314 : "b[bb * " + std::to_string(k * n) + "u + t * " + std::to_string(n) + "u + jj]")
315 << ";\n";
316 } else {
317 os << " r += a[ii * " << k << "u + t] * "
318 << (bQuant ? "bval(t * " + std::to_string(n) + "u + jj)" : "b[t * " + std::to_string(n) + "u + jj]")
319 << ";\n";
320 }
321 os << " }\n";
322 for (int u : grp.epilogue) ctx.emit(u);
323 const std::string finalExpr =
324 grp.epilogue.empty() ? std::string("r") : ctx.vars[static_cast<size_t>(grp.outputNode)];
325 os << " o[i_] = " << finalExpr << ";\n";
326 os << "}\n";
327 out.groupsX2 = groupsFor(batch * m * n);
328 } else {
329 ctx.biasIndexExpr = "col";
330 const int gx = (n + 15) / 16;
331 const int gy = (m + 15) / 16;
332 os << "var<workgroup> As: array<array<f32, 17>, 16>;\n";
333 os << "var<workgroup> Bs: array<array<f32, 17>, 16>;\n";
334 os << "@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
335 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) "
336 "{\n";
337 os << " var tx: u32 = localId.x;\n";
338 os << " var ty: u32 = localId.y;\n";
339 os << " var row: u32 = globalId.y;\n";
340 os << " var col: u32 = globalId.x;\n";
341 os << " var acc: f32 = 0.0;\n";
342 os << " for (var tile: u32 = 0u; tile < " << (k + 15) / 16 << "u; tile++) {\n";
343 os << " var aCol: u32 = tile * 16u + tx;\n";
344 os << " var bRow: u32 = tile * 16u + ty;\n";
345 os << " As[ty][tx] = 0.0; Bs[ty][tx] = 0.0;\n";
346 os << " if (aCol < " << k << "u && row < " << m << "u) { As[ty][tx] = a[row * " << k << "u + aCol]; }\n";
347 os << " if (bRow < " << k << "u && col < " << n << "u) { Bs[ty][tx] = b[bRow * " << n << "u + col]; }\n";
348 os << " workgroupBarrier();\n";
349 os << " for (var t: u32 = 0u; t < 16u; t++) { acc += As[ty][t] * Bs[t][tx]; }\n";
350 os << " workgroupBarrier();\n";
351 os << " }\n";
352 os << " if (row >= " << m << "u || col >= " << n << "u) { return; }\n";
353 os << " var i_: u32 = row * " << n << "u + col;\n";
354 os << " var r: f32 = acc;\n";
355 for (int u : grp.epilogue) ctx.emit(u);
356 const std::string finalExpr =
357 grp.epilogue.empty() ? std::string("r") : ctx.vars[static_cast<size_t>(grp.outputNode)];
358 os << " o[i_] = " << finalExpr << ";\n}\n";
359 out.groupsX2 = gx;
360 out.groupsY2 = gy;
361 }
362 out.pass1.clear();
363 out.pass2 = os.str();
364 out.inputCount = hasBias ? 3 : 2;
365 out.qDtype = bQuant ? B.dtype : 0;
366 out.qGroup = B.qGroup;
367 out.scalesBinding = bQuant && B.dtype != static_cast<int>(DType::Fp16) ? scalesBinding : -1;
368 out.outputBinding = bQuant ? bindingOut : -1;
369 if (tiled && batched)
370 throw eve::Exception(
371 "Tensor WGSL: unsupported kernel variant or binding count"); // batched matmul uses the naive variant
372 return;
373}
374
375void genConv(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
376 const GraphNode &cn = g.node(grp.nodes.front());
377 const GraphNode &X = g.node(cn.in0);
378 const GraphNode &Wt = g.node(cn.in1);
379 const bool is1d = cn.type == OpType::Conv1d;
380 const int stride = cn.i0, pad = cn.i1;
381 const bool hasBias = grp.biasNode >= 0;
382 const bool hasConvBias = cn.in2 >= 0;
383 const int bindingOut = 2 + int(hasConvBias) + int(hasBias);
384
385 std::ostringstream os;
386 os << header(kLocalSize);
387 os << bufferDecl(0, "x");
388 os << bufferDecl(1, "w");
389 if (hasConvBias) os << bufferDecl(2, "convBias");
390 if (hasBias) os << bufferDecl(hasConvBias ? 3 : 2, "bias");
391 os << bufferDecl(bindingOut, "o");
392 os << pushConstant();
393 std::vector<std::string> indexExprs(grp.inputs.size(), "i_");
394 ChainContext ctx{g, grp, indexExprs, "r", os, 0, ""};
395 ctx.biasIndexExpr = is1d ? "ol" : "ow";
396
397 os << "@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
398 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) {\n";
399 os << " var i_: u32 = globalId.x;\n";
400 os << " if (i_ >= " << g.node(grp.outputNode).size << "u) { return; }\n";
401 if (is1d) {
402 const int C = X.dims[1], L = X.dims[2];
403 const int F = Wt.dims[0], K = Wt.dims[2];
404 os << " var n_: u32 = i_ / (" << F << "u * " << cn.dims[2] << "u);\n";
405 os << " var rem: u32 = i_ % (" << F << "u * " << cn.dims[2] << "u);\n";
406 os << " var f: u32 = rem / " << cn.dims[2] << "u;\n";
407 os << " var ol: u32 = rem % " << cn.dims[2] << "u;\n";
408 os << " var r: f32 = " << (hasConvBias ? "convBias[f]" : "0.0") << ";\n";
409 os << " for (var c: u32 = 0u; c < " << C << "u; c++) {\n";
410 os << " for (var kk: u32 = 0u; kk < " << K << "u; kk++) {\n";
411 os << " var il: i32 = i32(ol) * " << stride << " + i32(kk) - " << pad << ";\n";
412 os << " if (il < 0 || il >= " << L << ") { continue; }\n";
413 os << " r += x[(n_ * " << C << "u + c) * " << L << "u + u32(il)] * w[(f * " << C << "u + c) * " << K
414 << "u + kk];\n";
415 os << " }\n";
416 os << " }\n";
417 } else {
418 const int C = X.dims[1], H = X.dims[2], W = X.dims[3];
419 const int F = Wt.dims[0], KH = Wt.dims[2], KW = Wt.dims[3];
420 os << " var n_: u32 = i_ / (" << F << "u * " << cn.dims[2] << "u * " << cn.dims[3] << "u);\n";
421 os << " var rem: u32 = i_ % (" << F << "u * " << cn.dims[2] << "u * " << cn.dims[3] << "u);\n";
422 os << " var f: u32 = rem / (" << cn.dims[2] << "u * " << cn.dims[3] << "u);\n";
423 os << " var rem2: u32 = rem % (" << cn.dims[2] << "u * " << cn.dims[3] << "u);\n";
424 os << " var oh: u32 = rem2 / " << cn.dims[3] << "u;\n";
425 os << " var ow: u32 = rem2 % " << cn.dims[3] << "u;\n";
426 os << " var r: f32 = " << (hasConvBias ? "convBias[f]" : "0.0") << ";\n";
427 os << " for (var c: u32 = 0u; c < " << C << "u; c++) {\n";
428 os << " for (var kh: u32 = 0u; kh < " << KH << "u; kh++) {\n";
429 os << " var ih: i32 = i32(oh) * " << stride << " + i32(kh) - " << pad << ";\n";
430 os << " if (ih < 0 || ih >= " << H << ") { continue; }\n";
431 os << " for (var kw: u32 = 0u; kw < " << KW << "u; kw++) {\n";
432 os << " var iw: i32 = i32(ow) * " << stride << " + i32(kw) - " << pad << ";\n";
433 os << " if (iw < 0 || iw >= " << W << ") { continue; }\n";
434 os << " r += x[((n_ * " << C << "u + c) * " << H << "u + u32(ih)) * " << W << "u + u32(iw)] * w[((f * "
435 << C << "u + c) * " << KH << "u + kh) * " << KW << "u + kw];\n";
436 os << " }\n";
437 os << " }\n";
438 os << " }\n";
439 }
440 for (int u : grp.epilogue) ctx.emit(u);
441 const std::string finalExpr =
442 grp.epilogue.empty() ? std::string("r") : ctx.vars[static_cast<size_t>(grp.outputNode)];
443 os << " o[i_] = " << finalExpr << ";\n";
444 os << "}\n";
445 out.pass1.clear();
446 out.pass2 = os.str();
447 out.groupsX2 = groupsFor(g.node(grp.outputNode).size);
448 out.inputCount = 2 + int(hasConvBias) + int(hasBias);
449 return;
450}
451
452void genPool(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
453 const GraphNode &pn = g.node(grp.outputNode);
454 const GraphNode &X = g.node(pn.in0);
455 const int ksize = pn.i0, stride = pn.i1, pad = pn.i2;
456 const int C = X.dims[1], H = X.dims[2], W = X.dims[3];
457 const int OH = pn.dims[2], OW = pn.dims[3];
458 const bool maxPool = pn.type == OpType::MaxPool2d;
459 std::ostringstream os;
460 os << header(kLocalSize);
461 os << bufferDecl(0, "in_");
462 os << bufferDecl(1, "o");
463 os << pushConstant();
464 os << "@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
465 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) {\n";
466 os << " var i_: u32 = globalId.x;\n";
467 os << " if (i_ >= " << pn.size << "u) { return; }\n";
468 os << " var n_: u32 = i_ / (" << C << "u * " << OH << "u * " << OW << "u);\n";
469 os << " var rem: u32 = i_ % (" << C << "u * " << OH << "u * " << OW << "u);\n";
470 os << " var c: u32 = rem / (" << OH << "u * " << OW << "u);\n";
471 os << " var rem2: u32 = rem % (" << OH << "u * " << OW << "u);\n";
472 os << " var oh: u32 = rem2 / " << OW << "u;\n";
473 os << " var ow: u32 = rem2 % " << OW << "u;\n";
474 os << " var acc: f32 = " << (maxPool ? "-3.402823e38" : "0.0") << ";\n";
475 os << " var valid: i32 = 0;\n";
476 os << " for (var kh: i32 = 0; kh < " << ksize << "; kh++) {\n";
477 os << " var ih: i32 = i32(oh) * " << stride << " + kh - " << pad << ";\n";
478 os << " if (ih < 0 || ih >= " << H << ") { continue; }\n";
479 os << " for (var kw: i32 = 0; kw < " << ksize << "; kw++) {\n";
480 os << " var iw: i32 = i32(ow) * " << stride << " + kw - " << pad << ";\n";
481 os << " if (iw < 0 || iw >= " << W << ") { continue; }\n";
482 os << " var v: f32 = in_[((n_ * " << C << "u + c) * " << H << "u + u32(ih)) * " << W << "u + u32(iw)];\n";
483 if (maxPool) {
484 os << " acc = max(acc, v);\n";
485 } else {
486 os << " acc += v; valid++;\n";
487 }
488 os << " }\n";
489 os << " }\n";
490 if (!maxPool) os << " if (valid > 0) { acc /= f32(valid); }\n";
491 os << " o[i_] = acc;\n";
492 os << "}\n";
493 out.pass1.clear();
494 out.pass2 = os.str();
495 out.groupsX2 = groupsFor(pn.size);
496 out.inputCount = 1;
497 return;
498}
499
500} // namespace wgsl_detail
501
503 using namespace wgsl_detail;
504 out = KernelSpec{};
505 switch (group.kind) {
506 case GroupKind::Elementwise: return genElementwise(graph, group, out);
508 // naive variant first; the runtime autotunes between naive and tiled
509 return genMatMul(graph, group, false, out);
511 case GroupKind::Conv2d: return genConv(graph, group, out);
513 case GroupKind::AvgPool2d: return genPool(graph, group, out);
514 case GroupKind::Softmax: return genSoftmax(graph, group, out);
515 case GroupKind::LayerNorm: return genNorm(graph, group, false, out);
516 case GroupKind::RMSNorm: return genNorm(graph, group, true, out);
518 case GroupKind::ArgMax: return genReduceOrArgmax(graph, group, group.kind == GroupKind::ArgMax, out);
519 case GroupKind::Embedding: return genEmbedding(graph, group, out);
520 case GroupKind::Concat: return genConcat(graph, group, out);
521 case GroupKind::Slice: return genSlice(graph, group, out);
522 case GroupKind::Permute: return genPermute(graph, group, out);
523 case GroupKind::Resize2d: return genResize2d(graph, group, out);
524 case GroupKind::Sdpa: return genSdpa(graph, group, out);
525 case GroupKind::Alias: return; // no kernel; pure buffer alias
526 }
527 throw eve::Exception("Tensor WGSL: unsupported kernel variant or binding count");
528}
529
531 try {
532 KernelSpec out;
534 if (group.kind != GroupKind::MatMul)
536 DiagnosticCode::Unsupported, "Tiled WGSL requires a matrix product", "tensor.kernel"));
538 } else {
540 }
542 return Result<KernelSpec>::success(std::move(out));
543 } catch (const std::exception &error) {
545 Diagnostic::error(DiagnosticCode::Unsupported, error.what(), "tensor.kernel"));
546 }
547}
548
549} // namespace eve::tensor
float x
Definition AnimClip.cpp:738
const std::string & s
std::string variant
building::EdgeCurveGroup group
std::string nodeId
tensor::Graph g
Definition GpuGraph.cpp:7
float u
Definition Grass.cpp:233
glm::vec3 n
Definition Grass.cpp:63
float v
std::string name
std::string error
Definition Package.cpp:60
std::map< std::string, std::vector< std::string > > graph
Definition Package.cpp:59
eve::action::ActionVfxBinding binding
int idx
float d
float t
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
EVENGINE_API_FOUNDATION public API.
Definition Exception.h:13
Move-only operation result carrying either a value or Status.
Definition Result.h:155
static Result success(T value)
Construct a successful result owning value.
Definition Result.h:164
static Result failure(Status status)
Construct a failed result from a structured status.
Definition Result.h:175
EVENGINE_API_DOMAINS public API.
Definition Graph.h:111
const GraphNode & node(int id) const
Node.
Definition Graph.h:116
bool isQuantDType(DType dt)
True when quant d type.
Definition Quant.h:17
std::string scalarStr(float v)
Scalar str.
std::string emitQuantizedBVal(DType dt, int group, const char *bufName="b")
Emit quantized b val.
std::string bufferDeclUint(int binding, const char *name)
Buffer decl uint.
void genConv(const Graph &g, const FusedGroup &grp, KernelSpec &out)
void specializeInputBindings(const Graph &graph, const FusedGroup &group, KernelSpec &spec)
Specialize input bindings.
void genElementwise(const Graph &g, const FusedGroup &grp, KernelSpec &out)
std::string bufferDecl(int binding, const char *name)
Buffer decl.
std::string pushConstant()
Pushes constant.
int groupsFor(int count)
Groups for.
void genPool(const Graph &g, const FusedGroup &grp, KernelSpec &out)
void emitInputIndexExprs(std::ostringstream &os, const Graph &g, const FusedGroup &grp, std::vector< std::string > &indexExprs)
std::string header(int localX, int localY)
Header.
void genMatMul(const Graph &g, const FusedGroup &grp, bool tiled, KernelSpec &out)
DType
Tensor element types.
Definition Tensor.h:24
KernelVariant
Internal choice of generated matrix multiplication implementation.
void generateKernelWgslImpl(const Graph &graph, const FusedGroup &group, KernelSpec &out)
constexpr int kMaxKernelBindings
Definition KernelGen.h:69
Result< KernelSpec > generateWgslKernel(const Graph &graph, const FusedGroup &group, KernelVariant variant)
Lower an optimizer-produced group directly to owning WGSL source and dispatch metadata.
FusedGroup public API.
Definition Optimizer.h:46
std::vector< int > epilogue
Definition Optimizer.h:65
std::vector< int > nodes
Definition Optimizer.h:51
std::vector< int > inputs
Definition Optimizer.h:53
GraphNode public API.
Definition Graph.h:80
int dims[Tensor::kMaxRank]
Definition Graph.h:82
KernelSpec public API.
Definition KernelGen.h:27
const std::vector< std::string > & indexExprs
uint32_t pad[2]