11namespace wgsl_detail {
14std::string
header(
int localX,
int localY) {
15 return "const workgroupX = " + std::to_string(localX) +
"u;\nconst workgroupY = " + std::to_string(localY) +
"u;\n";
18 return "@group(0) @binding(" + std::to_string(
binding) +
") var<storage, read_write> " +
name +
": array<f32>;\n";
21 return "@group(0) @binding(" + std::to_string(
binding) +
") var<storage, read_write> " +
name +
": array<u32>;\n";
30 std::ostringstream os;
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"
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"
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"
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"
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"
94 if (
v ==
int(
v) && std::fabs(
v) < 1e9f)
return std::to_string(
int(
v)) +
".0";
95 std::ostringstream os;
110 std::ostringstream &
os;
118 const size_t idx =
static_cast<size_t>(it -
group.
inputs.begin());
128 const std::string
x =
152 expr =
"(0.5 * " +
x +
" * (1.0 + tanh(0.7978845608028654 * (" +
x +
" + 0.044715 * " +
x +
" * " +
x +
155 case OpType::Silu: expr =
"(" +
x +
" / (1.0 + exp(-" +
x +
")))";
break;
157 float lo = nd.
s0, hi = nd.
s1;
158 if (lo > hi) std::swap(lo, hi);
169 const std::string var =
"t" + std::to_string(
temp++);
170 os <<
" var " << var <<
": f32 = " << expr <<
";\n";
183 std::vector<std::string> &indexExprs) {
185 const int rank = on.
rank;
186 std::vector<int> S(
static_cast<size_t>(rank), 1);
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];
191 indexExprs.resize(grp.
inputs.size());
192 for (
size_t k = 0; k < grp.
inputs.size(); ++k) {
194 bool identical = inN.
rank == rank;
196 for (
int d = 0;
d < rank; ++
d)
203 indexExprs[k] =
"i_";
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) {
210 if (dim == 1)
continue;
212 for (
int t =
d + 1;
t < rank; ++
t) {
214 if (dt != 1)
s *= dt;
216 stride[
static_cast<size_t>(
d)] =
s;
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";
231 throw eve::Exception(
"Tensor WGSL: unsupported kernel variant or binding count");
233 std::ostringstream os;
235 for (
size_t k = 0; k < grp.
inputs.size(); ++k) os <<
bufferDecl(
int(k), (
"a" + std::to_string(k)).c_str());
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;
245 for (
int u : grp.
nodes) ctx.emit(
u);
246 os <<
" o[i_] = " << ctx.vars[
static_cast<size_t>(grp.
outputNode)] <<
";\n";
249 out.
pass2 = os.str();
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;
268 throw eve::Exception(
"Tensor WGSL: unsupported kernel variant or binding count");
270 const int scalesBinding = hasBias ? 3 : 2;
271 const int bindingOut = bQuant ? (hasBias ? 4 : 3) : (hasBias ? 3 : 2);
273 std::ostringstream os;
288 if (bQuant) os << emitQuantizedBVal(static_cast<DType>(B.dtype), B.qGroup);
290 std::vector<std::string> indexExprs(grp.
inputs.size(),
"i_");
292 ctx.biasIndexExpr =
"jj";
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>) "
298 os <<
" var i_: u32 = globalId.x;\n";
299 os <<
" if (i_ >= " << batch *
m *
n <<
"u) { return; }\n";
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";
306 os <<
" var ii: u32 = i_ / " <<
n <<
"u;\n";
307 os <<
" var jj: u32 = i_ % " <<
n <<
"u;\n";
309 os <<
" var r: f32 = 0.0;\n";
310 os <<
" for (var t: u32 = 0u; t < " << k <<
"u; t++) {\n";
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]")
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]")
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";
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>) "
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";
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";
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";
363 out.
pass2 = os.str();
365 out.
qDtype = bQuant ? B.dtype : 0;
369 if (tiled && batched)
371 "Tensor WGSL: unsupported kernel variant or binding count");
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);
385 std::ostringstream os;
389 if (hasConvBias) os <<
bufferDecl(2,
"convBias");
390 if (hasBias) os <<
bufferDecl(hasConvBias ? 3 : 2,
"bias");
393 std::vector<std::string> indexExprs(grp.
inputs.size(),
"i_");
395 ctx.biasIndexExpr = is1d ?
"ol" :
"ow";
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";
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
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";
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";
446 out.
pass2 = os.str();
448 out.
inputCount = 2 + int(hasConvBias) + int(hasBias);
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];
459 std::ostringstream os;
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";
484 os <<
" acc = max(acc, v);\n";
486 os <<
" acc += v; valid++;\n";
490 if (!maxPool) os <<
" if (valid > 0) { acc /= f32(valid); }\n";
491 os <<
" o[i_] = acc;\n";
494 out.
pass2 = os.str();
503 using namespace wgsl_detail;
505 switch (
group.kind) {
527 throw eve::Exception(
"Tensor WGSL: unsupported kernel variant or binding count");
543 }
catch (
const std::exception &
error) {
building::EdgeCurveGroup group
std::map< std::string, std::vector< std::string > > graph
eve::action::ActionVfxBinding binding
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.
EVENGINE_API_FOUNDATION public API.
Move-only operation result carrying either a value or Status.
static Result success(T value)
Construct a successful result owning value.
static Result failure(Status status)
Construct a failed result from a structured status.
EVENGINE_API_DOMAINS public API.
const GraphNode & node(int id) const
Node.
bool isQuantDType(DType dt)
True when quant d type.
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.
KernelVariant
Internal choice of generated matrix multiplication implementation.
void generateKernelWgslImpl(const Graph &graph, const FusedGroup &group, KernelSpec &out)
constexpr int kMaxKernelBindings
Result< KernelSpec > generateWgslKernel(const Graph &graph, const FusedGroup &group, KernelVariant variant)
Lower an optimizer-produced group directly to owning WGSL source and dispatch metadata.
std::vector< int > epilogue
std::vector< int > inputs
int dims[Tensor::kMaxRank]
std::string biasIndexExpr
std::string operand(int nodeId)
std::string emit(int nodeId)
std::vector< std::string > vars
const std::vector< std::string > & indexExprs