10namespace glsl_detail {
13std::string
header(
int localX,
int localY) {
14 std::ostringstream os;
15 os <<
"#version 450\n";
16 os <<
"layout(local_size_x = " << localX;
17 if (localY > 1) os <<
", local_size_y = " << localY;
23 std::ostringstream os;
24 os <<
"layout(set = 0, binding = " <<
binding <<
") buffer B" <<
binding <<
" { float "
25 <<
name <<
"[]; };\n";
30 std::ostringstream os;
31 os <<
"layout(set = 0, binding = " <<
binding <<
") buffer B" <<
binding <<
" { uint "
32 <<
name <<
"[]; };\n";
42 std::ostringstream os;
46 os <<
"float bval(uint idx) {\n"
47 <<
" uint word = " << bufName <<
"[idx >> 2u];\n"
48 <<
" uint byte = (word >> ((idx & 3u) * 8u)) & 0xFFu;\n"
49 <<
" int v = int(byte);\n"
50 <<
" if (v >= 128) v -= 256;\n"
51 <<
" return float(v) * bs[idx / " <<
g <<
"u];\n"
55 os <<
"float bval(uint idx) {\n"
56 <<
" uint word = " << bufName <<
"[idx >> 3u];\n"
57 <<
" uint byte = (word >> (((idx >> 1u) & 3u) * 8u)) & 0xFFu;\n"
58 <<
" uint nib = (idx & 1u) == 0u ? (byte & 0xFu) : (byte >> 4u);\n"
59 <<
" int v = int(nib);\n"
60 <<
" if (v >= 8) v -= 16;\n"
61 <<
" return float(v) * bs[idx / " <<
g <<
"u];\n"
65 os <<
"float bval(uint idx) {\n"
66 <<
" uint word = " << bufName <<
"[idx >> 1u];\n"
67 <<
" uint hb = (word >> ((idx & 1u) * 16u)) & 0xFFFFu;\n"
68 <<
" return unpackHalf2x16(hb).x;\n"
72 os <<
"float bval(uint idx) {\n"
73 <<
" uint word = " << bufName <<
"[idx >> 2u];\n"
74 <<
" uint byte = (word >> ((idx & 3u) * 8u)) & 0xFFu;\n"
75 <<
" int s = (int(byte & 0x80u) != 0) ? -1 : 1;\n"
76 <<
" int e = int((byte >> 3u) & 0xFu);\n"
77 <<
" int m = int(byte & 0x7u);\n"
78 <<
" float v = (e == 0) ? exp2(-6.0) * float(m) / 8.0\n"
79 <<
" : exp2(float(e - 7)) * (1.0 + float(m) / 8.0);\n"
80 <<
" return float(s) * v * bs[idx / " <<
g <<
"u];\n"
84 os <<
"float bval(uint idx) {\n"
85 <<
" uint word = " << bufName <<
"[idx >> 3u];\n"
86 <<
" uint byte = (word >> (((idx >> 1u) & 3u) * 8u)) & 0xFFu;\n"
87 <<
" uint nib = (idx & 1u) == 0u ? (byte & 0xFu) : (byte >> 4u);\n"
88 <<
" int s = (int(nib & 8u) != 0) ? -1 : 1;\n"
89 <<
" int e = int((nib >> 1u) & 3u);\n"
90 <<
" int m = int(nib & 1u);\n"
91 <<
" float v = (e == 0) ? 0.5 * float(m)\n"
92 <<
" : exp2(float(e - 1)) * (1.0 + 0.5 * float(m));\n"
93 <<
" return float(s) * v * bs[idx / " <<
g <<
"u];\n"
102 return "layout(push_constant) uniform PC { float data[32]; } pc;\n";
108 if (
v ==
int(
v) && std::fabs(
v) < 1e9f)
return std::to_string(
int(
v)) +
".0";
109 std::ostringstream os;
124 std::ostringstream &
os;
133 const size_t idx =
static_cast<size_t>(it -
group.
inputs.begin());
167 expr =
"(0.5 * " +
x +
168 " * (1.0 + tanh(0.7978845608028654 * (" +
x +
169 " + 0.044715 * " +
x +
" * " +
x +
" * " +
x +
"))))";
171 case OpType::Silu: expr =
"(" +
x +
" / (1.0 + exp(-" +
x +
")))";
break;
173 float lo = nd.
s0, hi = nd.
s1;
174 if (lo > hi) std::swap(lo, hi);
186 const std::string var =
"t" + std::to_string(
temp++);
187 os <<
" float " << var <<
" = " << expr <<
";\n";
200 std::vector<std::string> &indexExprs) {
202 const int rank = on.
rank;
203 std::vector<int> S(
static_cast<size_t>(rank), 1);
205 S[
static_cast<size_t>(rank - 1)] = 1;
206 for (
int k = rank - 2; k >= 0; --k) S[
static_cast<size_t>(k)] =
207 S[
static_cast<size_t>(k + 1)] * on.
dims[k + 1];
209 indexExprs.resize(grp.
inputs.size());
210 for (
size_t k = 0; k < grp.
inputs.size(); ++k) {
212 bool identical = inN.
rank == rank;
214 for (
int d = 0;
d < rank; ++
d)
221 indexExprs[k] =
"i_";
224 const int pad = rank - inN.
rank;
225 std::vector<int> stride(
static_cast<size_t>(rank), 0);
226 for (
int d = 0;
d < rank; ++
d) {
228 if (dim == 1)
continue;
230 for (
int t =
d + 1;
t < rank; ++
t) {
232 if (dt != 1)
s *= dt;
234 stride[
static_cast<size_t>(
d)] =
s;
236 const std::string var =
"idx" + std::to_string(k);
237 os <<
" uint " << var <<
" = 0u;\n";
238 for (
int d = 0;
d < rank; ++
d) {
239 if (stride[
static_cast<size_t>(
d)] == 0)
continue;
240 os <<
" " << var <<
" += ((i_ / " << S[
static_cast<size_t>(
d)] <<
"u) % "
241 << on.
dims[
d] <<
"u) * " << stride[
static_cast<size_t>(
d)] <<
"u;\n";
250 std::ostringstream os;
252 for (
size_t k = 0; k < grp.
inputs.size(); ++k)
253 os <<
bufferDecl(
int(k), (
"a" + std::to_string(k)).c_str());
256 os <<
"void main() {\n";
257 os <<
" uint i_ = gl_GlobalInvocationID.x;\n";
258 os <<
" if (i_ >= " << on.
size <<
"u) return;\n";
259 std::vector<std::string> indexExprs;
262 for (
int u : grp.
nodes) ctx.emit(
u);
263 os <<
" o[i_] = " << ctx.vars[
static_cast<size_t>(grp.
outputNode)] <<
";\n";
266 out.
pass2 = os.str();
277 const bool batched = mm.
rank == 3;
278 const int batch = batched ? A.dims[0] : 1;
279 const int m = A.dims[batched ? 1 : 0];
280 const int k = A.dims[batched ? 2 : 1];
281 const int n = B.dims[batched ? 2 : 1];
282 const bool hasBias = grp.
biasNode >= 0;
284 if (bQuant && tiled)
return false;
285 const int scalesBinding = hasBias ? 3 : 2;
286 const int bindingOut = bQuant ? (hasBias ? 4 : 3) : (hasBias ? 3 : 2);
288 std::ostringstream os;
298 if (bQuant && B.dtype !=
static_cast<int>(
DType::Fp16))
302 if (bQuant) os << emitQuantizedBVal(static_cast<DType>(B.dtype), B.qGroup);
304 std::vector<std::string> indexExprs(grp.
inputs.size(),
"i_");
306 ctx.biasIndexExpr =
"jj";
309 os <<
"void main() {\n";
310 os <<
" uint i_ = gl_GlobalInvocationID.x;\n";
311 os <<
" if (i_ >= " << batch *
m *
n <<
"u) return;\n";
313 os <<
" uint bb = i_ / " <<
m *
n <<
"u;\n";
314 os <<
" uint rem = i_ % " <<
m *
n <<
"u;\n";
315 os <<
" uint ii = rem / " <<
n <<
"u;\n";
316 os <<
" uint jj = rem % " <<
n <<
"u;\n";
318 os <<
" uint ii = i_ / " <<
n <<
"u;\n";
319 os <<
" uint jj = i_ % " <<
n <<
"u;\n";
321 os <<
" float r = 0.0;\n";
322 os <<
" for (uint t = 0u; t < " << k <<
"u; ++t) {\n";
324 os <<
" r += a[bb * " <<
m * k <<
"u + ii * " << k <<
"u + t] * "
325 << (bQuant ?
"bval(bb * " + std::to_string(k *
n) +
"u + t * " +
326 std::to_string(
n) +
"u + jj)"
327 :
"b[bb * " + std::to_string(k *
n) +
"u + t * " +
328 std::to_string(
n) +
"u + jj]")
331 os <<
" r += a[ii * " << k <<
"u + t] * "
332 << (bQuant ?
"bval(t * " + std::to_string(
n) +
"u + jj)"
333 :
"b[t * " + std::to_string(
n) +
"u + jj]")
338 const std::string finalExpr =
339 grp.
epilogue.empty() ? std::string(
"r") : ctx.vars[
static_cast<size_t>(grp.
outputNode)];
340 os <<
" o[i_] = " << finalExpr <<
";\n";
344 ctx.biasIndexExpr =
"col";
345 const int gx = (
n + 15) / 16;
346 const int gy = (
m + 15) / 16;
347 os <<
"shared float As[16][17];\n";
348 os <<
"shared float Bs[16][17];\n";
349 os <<
"void main() {\n";
350 os <<
" uint tx = gl_LocalInvocationID.x;\n";
351 os <<
" uint ty = gl_LocalInvocationID.y;\n";
352 os <<
" uint row = gl_GlobalInvocationID.y;\n";
353 os <<
" uint col = gl_GlobalInvocationID.x;\n";
354 os <<
" float acc = 0.0;\n";
355 os <<
" for (uint tile = 0u; tile < " << (k + 15) / 16 <<
"u; ++tile) {\n";
356 os <<
" uint aCol = tile * 16u + tx;\n";
357 os <<
" uint bRow = tile * 16u + ty;\n";
358 os <<
" As[ty][tx] = (aCol < " << k <<
"u && row < " <<
m <<
"u) ? a[row * "
359 << k <<
"u + aCol] : 0.0;\n";
360 os <<
" Bs[ty][tx] = (bRow < " << k <<
"u && col < " <<
n <<
"u) ? b[bRow * "
361 <<
n <<
"u + col] : 0.0;\n";
362 os <<
" barrier();\n";
363 os <<
" for (uint t = 0u; t < 16u; ++t) acc += As[ty][t] * Bs[t][tx];\n";
364 os <<
" barrier();\n";
366 os <<
" float r = acc;\n";
368 const std::string finalExpr =
369 grp.
epilogue.empty() ? std::string(
"r") : ctx.vars[
static_cast<size_t>(grp.
outputNode)];
370 os <<
" if (row < " <<
m <<
"u && col < " <<
n <<
"u) o[row * " <<
n <<
"u + col] = "
371 << finalExpr <<
";\n";
377 out.
pass2 = os.str();
379 out.
qDtype = bQuant ? B.dtype : 0;
383 if (tiled && batched)
return false;
392 const int stride = cn.
i0,
pad = cn.
i1;
393 const bool hasBias = grp.
biasNode >= 0;
394 const bool hasConvBias = cn.
in2 >= 0;
395 const int bindingOut = 2 + int(hasConvBias) + int(hasBias);
397 std::ostringstream os;
401 if (hasConvBias) os <<
bufferDecl(2,
"convBias");
402 if (hasBias) os <<
bufferDecl(hasConvBias ? 3 : 2,
"bias");
405 std::vector<std::string> indexExprs(grp.
inputs.size(),
"i_");
407 ctx.biasIndexExpr = is1d ?
"ol" :
"ow";
409 os <<
"void main() {\n";
410 os <<
" uint i_ = gl_GlobalInvocationID.x;\n";
411 os <<
" if (i_ >= " <<
g.node(grp.
outputNode).size <<
"u) return;\n";
413 const int C = X.dims[1], L = X.dims[2];
414 const int F = Wt.
dims[0], K = Wt.
dims[2];
415 os <<
" uint n_ = i_ / (" << F <<
"u * " << cn.
dims[2] <<
"u);\n";
416 os <<
" uint rem = i_ % (" << F <<
"u * " << cn.
dims[2] <<
"u);\n";
417 os <<
" uint f = rem / " << cn.
dims[2] <<
"u;\n";
418 os <<
" uint ol = rem % " << cn.
dims[2] <<
"u;\n";
419 os <<
" float r = " << (hasConvBias ?
"convBias[f]" :
"0.0") <<
";\n";
420 os <<
" for (uint c = 0u; c < " << C <<
"u; ++c) {\n";
421 os <<
" for (uint kk = 0u; kk < " << K <<
"u; ++kk) {\n";
422 os <<
" int il = int(ol) * " << stride <<
" + int(kk) - " <<
pad <<
";\n";
423 os <<
" if (il < 0 || il >= " << L <<
") continue;\n";
424 os <<
" r += x[(n_ * " << C <<
"u + c) * " << L <<
"u + uint(il)] * w[(f * "
425 << C <<
"u + c) * " << K <<
"u + kk];\n";
429 const int C = X.dims[1], H = X.dims[2], W = X.dims[3];
430 const int F = Wt.
dims[0], KH = Wt.
dims[2], KW = Wt.
dims[3];
431 os <<
" uint n_ = i_ / (" << F <<
"u * " << cn.
dims[2] <<
"u * " << cn.
dims[3] <<
"u);\n";
432 os <<
" uint rem = i_ % (" << F <<
"u * " << cn.
dims[2] <<
"u * " << cn.
dims[3] <<
"u);\n";
433 os <<
" uint f = rem / (" << cn.
dims[2] <<
"u * " << cn.
dims[3] <<
"u);\n";
434 os <<
" uint rem2 = rem % (" << cn.
dims[2] <<
"u * " << cn.
dims[3] <<
"u);\n";
435 os <<
" uint oh = rem2 / " << cn.
dims[3] <<
"u;\n";
436 os <<
" uint ow = rem2 % " << cn.
dims[3] <<
"u;\n";
437 os <<
" float r = " << (hasConvBias ?
"convBias[f]" :
"0.0") <<
";\n";
438 os <<
" for (uint c = 0u; c < " << C <<
"u; ++c) {\n";
439 os <<
" for (uint kh = 0u; kh < " << KH <<
"u; ++kh) {\n";
440 os <<
" int ih = int(oh) * " << stride <<
" + int(kh) - " <<
pad <<
";\n";
441 os <<
" if (ih < 0 || ih >= " << H <<
") continue;\n";
442 os <<
" for (uint kw = 0u; kw < " << KW <<
"u; ++kw) {\n";
443 os <<
" int iw = int(ow) * " << stride <<
" + int(kw) - " <<
pad <<
";\n";
444 os <<
" if (iw < 0 || iw >= " << W <<
") continue;\n";
445 os <<
" r += x[((n_ * " << C <<
"u + c) * " << H <<
"u + uint(ih)) * " << W
446 <<
"u + uint(iw)] * w[((f * " << C <<
"u + c) * " << KH <<
"u + kh) * " << KW
453 const std::string finalExpr =
454 grp.
epilogue.empty() ? std::string(
"r") : ctx.vars[
static_cast<size_t>(grp.
outputNode)];
455 os <<
" o[i_] = " << finalExpr <<
";\n";
458 out.
pass2 = os.str();
460 out.
inputCount = 2 + int(hasConvBias) + int(hasBias);
467 const int ksize = pn.
i0, stride = pn.
i1,
pad = pn.
i2;
468 const int C = X.dims[1], H = X.dims[2], W = X.dims[3];
469 const int OH = pn.
dims[2], OW = pn.
dims[3];
471 std::ostringstream os;
476 os <<
"void main() {\n";
477 os <<
" uint i_ = gl_GlobalInvocationID.x;\n";
478 os <<
" if (i_ >= " << pn.
size <<
"u) return;\n";
479 os <<
" uint n_ = i_ / (" << C <<
"u * " << OH <<
"u * " << OW <<
"u);\n";
480 os <<
" uint rem = i_ % (" << C <<
"u * " << OH <<
"u * " << OW <<
"u);\n";
481 os <<
" uint c = rem / (" << OH <<
"u * " << OW <<
"u);\n";
482 os <<
" uint rem2 = rem % (" << OH <<
"u * " << OW <<
"u);\n";
483 os <<
" uint oh = rem2 / " << OW <<
"u;\n";
484 os <<
" uint ow = rem2 % " << OW <<
"u;\n";
485 os <<
" float acc = " << (maxPool ?
"-3.402823e38" :
"0.0") <<
";\n";
486 os <<
" int valid = 0;\n";
487 os <<
" for (int kh = 0; kh < " << ksize <<
"; ++kh) {\n";
488 os <<
" int ih = int(oh) * " << stride <<
" + kh - " <<
pad <<
";\n";
489 os <<
" if (ih < 0 || ih >= " << H <<
") continue;\n";
490 os <<
" for (int kw = 0; kw < " << ksize <<
"; ++kw) {\n";
491 os <<
" int iw = int(ow) * " << stride <<
" + kw - " <<
pad <<
";\n";
492 os <<
" if (iw < 0 || iw >= " << W <<
") continue;\n";
493 os <<
" float v = in_[((n_ * " << C <<
"u + c) * " << H <<
"u + uint(ih)) * " << W
494 <<
"u + uint(iw)];\n";
496 os <<
" acc = max(acc, v);\n";
498 os <<
" acc += v; ++valid;\n";
502 if (!maxPool) os <<
" acc = valid > 0 ? acc / float(valid) : 0.0;\n";
503 os <<
" o[i_] = acc;\n";
506 out.
pass2 = os.str();
515 const int axis = sn.
i0;
516 int outer = 1,
reduce = 1, inner = 1;
517 for (
int k = 0; k < axis; ++k) outer *= X.dims[k];
519 for (
int k = axis + 1; k < X.rank; ++k) inner *= X.dims[k];
520 const int rows = outer * inner;
521 const bool logMode = grp.
logMode;
523 std::ostringstream os1, os2;
529 os1 <<
"void main() {\n";
530 os1 <<
" uint i_ = gl_GlobalInvocationID.x;\n";
531 os1 <<
" if (i_ >= " <<
rows <<
"u) return;\n";
532 os1 <<
" uint o_ = i_ / " << inner <<
"u;\n";
533 os1 <<
" uint ii = i_ % " << inner <<
"u;\n";
534 os1 <<
" float m = -3.402823e38;\n";
535 os1 <<
" for (uint j = 0u; j < " <<
reduce <<
"u; ++j) {\n";
536 os1 <<
" m = max(m, in_[(o_ * " <<
reduce <<
"u + j) * " << inner <<
"u + ii]);\n";
538 os1 <<
" float s = 0.0;\n";
539 os1 <<
" for (uint j = 0u; j < " <<
reduce <<
"u; ++j) {\n";
540 os1 <<
" s += exp(in_[(o_ * " <<
reduce <<
"u + j) * " << inner <<
"u + ii] - m);\n";
542 os1 <<
" mx[i_] = m;\n";
543 os1 <<
" sm[i_] = s;\n";
552 os2 <<
"void main() {\n";
553 os2 <<
" uint i_ = gl_GlobalInvocationID.x;\n";
554 os2 <<
" if (i_ >= " << sn.
size <<
"u) return;\n";
555 os2 <<
" uint o_ = (i_ / (" <<
reduce * inner <<
"u)) * " << inner <<
"u + (i_ % "
557 os2 <<
" float m = mx[o_];\n";
558 os2 <<
" float s = sm[o_];\n";
559 os2 <<
" float x = in_[i_];\n";
561 os2 <<
" o[i_] = (x - m) - log(s);\n";
563 os2 <<
" o[i_] = exp(x - m) / s;\n";
566 out.
pass1 = os1.str();
567 out.
pass2 = os2.str();
581 const int cols = X.dims[X.rank - 1];
583 const float eps = nn.
s0;
585 const bool hasBias = !rms && grp.
hasBias;
586 const int inputCount = 1 + (hasScale ? 1 : 0) + (hasBias ? 1 : 0);
587 const int statsCount = rms ? 1 : 2;
588 const int outBinding = inputCount + statsCount;
590 std::ostringstream os1, os2;
594 if (statsCount > 1) os1 <<
bufferDecl(2,
"st1");
596 os1 <<
"void main() {\n";
597 os1 <<
" uint i_ = gl_GlobalInvocationID.x;\n";
598 os1 <<
" if (i_ >= " <<
rows <<
"u) return;\n";
599 os1 <<
" float s0 = 0.0;\n";
600 if (statsCount > 1) os1 <<
" float s1 = 0.0;\n";
601 os1 <<
" for (uint j = 0u; j < " <<
cols <<
"u; ++j) {\n";
602 os1 <<
" float v = in_[i_ * " <<
cols <<
"u + j];\n";
604 os1 <<
" s0 += v * v;\n";
606 os1 <<
" s0 += v; s1 += v * v;\n";
609 os1 <<
" st0[i_] = s0;\n";
610 if (statsCount > 1) os1 <<
" st1[i_] = s1;\n";
614 for (
int k = 0; k < inputCount; ++k)
615 os2 <<
bufferDecl(k, (
"a" + std::to_string(k)).c_str());
617 if (statsCount > 1) os2 <<
bufferDecl(inputCount + 1,
"st1");
620 os2 <<
"void main() {\n";
621 os2 <<
" uint i_ = gl_GlobalInvocationID.x;\n";
622 os2 <<
" if (i_ >= " << nn.
size <<
"u) return;\n";
623 os2 <<
" uint r = i_ / " <<
cols <<
"u;\n";
624 os2 <<
" uint c = i_ % " <<
cols <<
"u;\n";
626 os2 <<
" float inv = 1.0 / sqrt(st0[r] / " <<
cols <<
".0 + " <<
scalarStr(eps)
628 os2 <<
" float y = a0[i_] * inv;\n";
629 if (hasScale) os2 <<
" y *= a1[c];\n";
631 os2 <<
" float mean = st0[r] / " <<
cols <<
".0;\n";
632 os2 <<
" float var = st1[r] / " <<
cols <<
".0 - mean * mean;\n";
633 os2 <<
" var = max(var, 0.0);\n";
634 os2 <<
" float inv = 1.0 / sqrt(var + " <<
scalarStr(eps) <<
");\n";
635 os2 <<
" float y = (a0[i_] - mean) * inv;\n";
636 if (hasScale) os2 <<
" y *= a1[c];\n";
637 if (hasBias) os2 <<
" y += a" << (hasScale ? 2 : 1) <<
"[c];\n";
639 os2 <<
" o[i_] = y;\n";
641 out.
pass1 = os1.str();
642 out.
pass2 = os2.str();
656 const int axis = rn.
i0;
657 int outer = 1,
reduce = 1, inner = 1;
658 for (
int k = 0; k < axis; ++k) outer *= X.dims[k];
660 for (
int k = axis + 1; k < X.rank; ++k) inner *= X.dims[k];
661 const int outSize = outer * inner;
662 std::ostringstream os;
667 os <<
"void main() {\n";
668 os <<
" uint i_ = gl_GlobalInvocationID.x;\n";
669 os <<
" if (i_ >= " << outSize <<
"u) return;\n";
670 os <<
" uint o_ = i_ / " << inner <<
"u;\n";
671 os <<
" uint ii = i_ % " << inner <<
"u;\n";
673 os <<
" float best = -3.402823e38;\n";
674 os <<
" float bestJ = 0.0;\n";
675 os <<
" for (uint j = 0u; j < " <<
reduce <<
"u; ++j) {\n";
676 os <<
" float v = in_[(o_ * " <<
reduce <<
"u + j) * " << inner <<
"u + ii];\n";
677 os <<
" if (v > best) { best = v; bestJ = float(j); }\n";
679 os <<
" o[i_] = bestJ;\n";
684 os <<
" float acc = 0.0;\n";
685 os <<
" for (uint j = 0u; j < " <<
reduce <<
"u; ++j) acc += in_[(o_ * "
686 <<
reduce <<
"u + j) * " << inner <<
"u + ii];\n";
691 os <<
" float acc = 3.402823e38;\n";
692 os <<
" for (uint j = 0u; j < " <<
reduce <<
"u; ++j) acc = min(acc, in_[(o_ * "
693 <<
reduce <<
"u + j) * " << inner <<
"u + ii]);\n";
696 os <<
" float acc = -3.402823e38;\n";
697 os <<
" for (uint j = 0u; j < " <<
reduce <<
"u; ++j) acc = max(acc, in_[(o_ * "
698 <<
reduce <<
"u + j) * " << inner <<
"u + ii]);\n";
700 default:
return false;
702 os <<
" o[i_] = acc;\n";
706 out.
pass2 = os.str();
717 const int vocab = T.
dims[0], dim = T.
dims[1];
718 std::ostringstream os;
723 if (tQuant && tInt) os <<
bufferDecl(2,
"bs");
726 if (tQuant) os << emitQuantizedBVal(static_cast<DType>(T.
dtype), T.
qGroup,
"table");
727 os <<
"void main() {\n";
728 os <<
" uint i_ = gl_GlobalInvocationID.x;\n";
729 os <<
" if (i_ >= " << en.
size <<
"u) return;\n";
730 os <<
" uint r = i_ / " << dim <<
"u;\n";
731 os <<
" uint d = i_ % " << dim <<
"u;\n";
732 os <<
" int ii = int(idx[r]);\n";
733 os <<
" ii = clamp(ii, 0, " << (vocab - 1) <<
");\n";
734 os <<
" o[i_] = " << (tQuant ?
"bval(uint(ii) * " + std::to_string(dim) +
"u + d)"
735 :
"table[uint(ii) * " + std::to_string(dim) +
"u + d]")
739 out.
pass2 = os.str();
751 const int axis = cn.
i0;
752 const int n = int(grp.
inputs.size());
753 if (n < 2 || n > 4)
return false;
756 for (
int k = 0; k <
n; ++k) {
757 starts[k] = axisTotal;
758 axisTotal +=
g.node(grp.
inputs[
static_cast<size_t>(k)]).dims[axis];
761 for (
int k = axis + 1; k < cn.
rank; ++k) inner *= cn.
dims[k];
762 std::ostringstream os;
764 for (
int k = 0; k <
n; ++k) os <<
bufferDecl(k, (
"a" + std::to_string(k)).c_str());
767 os <<
"void main() {\n";
768 os <<
" uint i_ = gl_GlobalInvocationID.x;\n";
769 os <<
" if (i_ >= " << cn.
size <<
"u) return;\n";
770 os <<
" uint ax = (i_ / " << inner <<
"u) % " << axisTotal <<
"u;\n";
771 os <<
" uint op = i_ / (" << axisTotal <<
"u * " << inner <<
"u);\n";
772 os <<
" uint ip = i_ % " << inner <<
"u;\n";
773 os <<
" float v = 0.0;\n";
774 for (
int k = 0; k <
n; ++k) {
775 const int sz =
g.node(grp.
inputs[
static_cast<size_t>(k)]).dims[axis];
776 const char *cond = k == 0 ?
"if" :
"else if";
777 os <<
" " << cond <<
" (ax >= " << starts[k] <<
"u && ax < " << starts[k] +
sz
779 os <<
" v = a" << k <<
"[op * " <<
sz <<
"u * " << inner <<
"u + (ax - " << starts[k]
780 <<
"u) * " << inner <<
"u + ip];\n";
783 os <<
" o[i_] = v;\n";
786 out.
pass2 = os.str();
798 for (
int k = axis + 1; k < sn.
rank; ++k) inner *= sn.
dims[k];
799 std::ostringstream os;
804 os <<
"void main() {\n";
805 os <<
" uint i_ = gl_GlobalInvocationID.x;\n";
806 os <<
" if (i_ >= " << sn.
size <<
"u) return;\n";
807 os <<
" uint ax = (i_ / " << inner <<
"u) % " << axisSize <<
"u;\n";
808 os <<
" uint op = i_ / (" << axisSize <<
"u * " << inner <<
"u);\n";
809 os <<
" uint ip = i_ % " << inner <<
"u;\n";
810 os <<
" o[i_] = in_[op * " << X.dims[axis] <<
"u * " << inner <<
"u + (ax + " <<
begin
811 <<
"u) * " << inner <<
"u + ip];\n";
814 out.
pass2 = os.str();
823 const int rank = pn.
rank;
826 for (
int k = rank - 2; k >= 0; --k) S[k] = S[k + 1] * pn.
dims[k + 1];
828 inStride[rank - 1] = 1;
829 for (
int k = rank - 2; k >= 0; --k) inStride[k] = inStride[k + 1] * X.dims[k + 1];
830 std::ostringstream os;
835 os <<
"void main() {\n";
836 os <<
" uint i_ = gl_GlobalInvocationID.x;\n";
837 os <<
" if (i_ >= " << pn.
size <<
"u) return;\n";
838 os <<
" uint idx = 0u;\n";
839 for (
int k = 0; k < rank; ++k) {
840 const int inAxis = pn.
perm[k];
841 os <<
" idx += ((i_ / " << S[k] <<
"u) % " << pn.
dims[k] <<
"u) * "
842 << inStride[inAxis] <<
"u;\n";
844 os <<
" o[i_] = in_[idx];\n";
847 out.
pass2 = os.str();
858 const int S = K.
dims[2];
859 if (S > 2048 || D > 512)
return false;
861 const bool masked = grp.
masked;
862 const int bindingOut = masked ? 4 : 3;
863 std::ostringstream os;
870 os <<
"shared float scores[" << S <<
"];\n";
871 os <<
"shared float maxv;\n";
872 os <<
"shared float sumv;\n";
874 os <<
"void main() {\n";
875 os <<
" uint tid = gl_LocalInvocationID.x;\n";
876 os <<
" uint bh = gl_WorkGroupID.x;\n";
877 os <<
" uint t = gl_WorkGroupID.y;\n";
878 os <<
" uint b = bh / " << H <<
"u;\n";
879 os <<
" uint h = bh % " << H <<
"u;\n";
880 os <<
" uint qbase = (b * " << H <<
"u + h) * " << T <<
"u * " << D <<
"u + t * " << D
882 os <<
" uint kbase = (b * " << H <<
"u + h) * " << S <<
"u * " << D <<
"u;\n";
883 os <<
" uint vbase = kbase;\n";
884 os <<
" for (uint s = tid; s < " << S <<
"u; s += 128u) {\n";
885 os <<
" float acc = 0.0;\n";
886 os <<
" for (uint d = 0u; d < " << D <<
"u; ++d) acc += q[qbase + d] * k[kbase + s * "
890 os <<
" acc += mask[(b * " << H <<
"u + h) * " << T <<
"u * " << S <<
"u + t * " << S
893 os <<
" scores[s] = acc;\n";
895 os <<
" barrier();\n";
896 os <<
" if (tid == 0u) {\n";
897 os <<
" float m = -3.402823e38;\n";
898 os <<
" for (uint s = 0u; s < " << S <<
"u; ++s) m = max(m, scores[s]);\n";
899 os <<
" float sm = 0.0;\n";
900 os <<
" for (uint s = 0u; s < " << S <<
"u; ++s) sm += exp(scores[s] - m);\n";
901 os <<
" maxv = m; sumv = sm;\n";
903 os <<
" barrier();\n";
904 os <<
" for (uint d = tid; d < " << D <<
"u; d += 128u) {\n";
905 os <<
" float acc = 0.0;\n";
906 os <<
" for (uint s = 0u; s < " << S <<
"u; ++s) acc += exp(scores[s] - maxv) * v[vbase + s * "
908 os <<
" o[qbase + d] = acc / sumv;\n";
912 out.
pass2 = os.str();
922 using namespace glsl_detail;
924 switch (
group.kind) {
952 using namespace glsl_detail;
building::EdgeCurveGroup group
std::array< float, 3 > scale
std::map< std::string, std::vector< std::string > > graph
eve::action::ActionVfxBinding binding
EVENGINE_API_DOMAINS public API.
const GraphNode & node(int id) const
Node.
static constexpr int kMaxRank
bool genConv(const Graph &g, const FusedGroup &grp, KernelSpec &out)
std::string emitQuantizedBVal(DType dt, int group, const char *bufName="b")
bool genSdpa(const Graph &g, const FusedGroup &grp, KernelSpec &out)
std::string header(int localX, int localY)
Header.
std::string bufferDecl(int binding, const char *name)
Buffer decl.
bool genEmbedding(const Graph &g, const FusedGroup &grp, KernelSpec &out)
bool genSoftmax(const Graph &g, const FusedGroup &grp, KernelSpec &out)
bool genConcat(const Graph &g, const FusedGroup &grp, KernelSpec &out)
int groupsFor(int count)
Groups for.
bool genNorm(const Graph &g, const FusedGroup &grp, bool rms, KernelSpec &out)
std::string bufferDeclUint(int binding, const char *name)
std::string scalarStr(float v)
Scalar str.
bool genSlice(const Graph &g, const FusedGroup &grp, KernelSpec &out)
bool genReduceOrArgmax(const Graph &g, const FusedGroup &grp, bool argmax, KernelSpec &out)
std::string pushConstant()
Pushes constant.
bool genMatMul(const Graph &g, const FusedGroup &grp, bool tiled, KernelSpec &out)
bool genPermute(const Graph &g, const FusedGroup &grp, KernelSpec &out)
bool 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)
bool genElementwise(const Graph &g, const FusedGroup &grp, KernelSpec &out)
bool isQuantDType(DType dt)
True when quant d type.
DType
Tensor element types.
bool generateKernel(const Graph &graph, const FusedGroup &group, KernelSpec &out)
Generate kernel.
constexpr int kMaxKernelBindings
bool generateMatMulVariant(const Graph &graph, const FusedGroup &group, bool tiled, KernelSpec &out)
Generate mat mul variant.
std::vector< int > epilogue
std::vector< int > inputs
int perm[Tensor::kMaxRank]
int dims[Tensor::kMaxRank]
std::vector< uint8_t > constBytes
std::string biasIndexExpr
std::string emit(int nodeId)
std::vector< std::string > vars
const std::vector< std::string > & indexExprs
std::string operand(int nodeId)
gpgpu::ComputeShader * reduce