// wgsl.js — WGSL-кернелы ядра трансформера Falcon-Perception-0.6B. // // Конвенции: // - веса FPQ4: асимметричная q4, блок 128 вдоль K, packed u8 [N, K/128, 64] // (= сплошной поток u32 по строке, ниббл j слова = биты [4j..4j+3]), // scales f16 [N, K/128] (пары в u32, чётный индекс — младшие 16 бит), // zeros u8 [N, nb/2] (младший ниббл — чётный блок); w = (q - zp) * scale; // - активации f32; GEMM держит их в shared как f16 (семантика зеркалится // эталоном make_refdata.py); KV-кэш f16, 8 голов, канонический K // (1D-RoPE на dims 0-63, верх без поворота), 2D-поворот — при чтении // по таблице rope2 (текстовые строки identity); // - редукции — subgroupAdd + досведение через shared (фича subgroups // обязательна для движка; фоллбэки — вне задачи ядра). export const GRID_X = 16384; // 2D-адресация строк GEMV (N=65536 > лимита 65535) const SCALE_ATTN = '0.0883883461356163'; const NEG_MAX = '-3.4028234663852886e+38'; // Общий кусок: чтение скейла (f16-пары в u32) и zero-point. // nb = k/128 всегда чётно для наших матриц => строки scales выровнены на u32. // bits=4: zeros u8 [N, nb/2] — ниблы (младший — чётный блок); // bits=8: zeros u8 [N, nb] — байт на блок. const qHelpers = (bits) => /* wgsl */ ` fn scale_at(row : u32, blk : u32, nb : u32) -> f32 { let idx = row * nb + blk; let pair = unpack2x16float(scales[idx >> 1u]); return select(pair.x, pair.y, (idx & 1u) == 1u); } fn zp_at(row : u32, blk : u32, nb : u32) -> f32 { ${bits === 4 ? ` let zb = (nb + 1u) >> 1u; // байт на 2 блока let byte_idx = row * zb + (blk >> 1u); let word = zeros[byte_idx >> 2u]; let byte = (word >> ((byte_idx & 3u) * 8u)) & 0xFFu; return f32(select(byte & 0xFu, byte >> 4u, (blk & 1u) == 1u));` : ` let byte_idx = row * nb + blk; let word = zeros[byte_idx >> 2u]; return f32((word >> ((byte_idx & 3u) * 8u)) & 0xFFu);`} } `; // ---------------------------------------------------------------- RMSNorm // Один workgroup (256 потоков) на строку. y = x * rsqrt(mean(x^2) + eps) [* w]. export function rmsnorm(D, withWeight) { const WG = 256; const PER = D / WG; // 1024/256 = 4 return /* wgsl */ ` enable subgroups; struct Params { rows : u32, x_base : u32, eps : f32, p0 : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var x : array; @group(0) @binding(2) var y : array; ${withWeight ? '@group(0) @binding(3) var w : array;' : ''} const D : u32 = ${D}u; var sums : array; var inv_rms : f32; @compute @workgroup_size(${WG}) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_index) lid : u32, @builtin(subgroup_invocation_id) sg_lane : u32, @builtin(subgroup_size) sg_size : u32) { let row = wg_id.x; let base = params.x_base + row * D; var ss = 0.0; for (var i = lid; i < D; i = i + ${WG}u) { let v = x[base + i]; ss = ss + v * v; } let sg_sum = subgroupAdd(ss); if (sg_lane == 0u) { sums[lid / sg_size] = sg_sum; } workgroupBarrier(); if (lid == 0u) { let n_sg = (${WG}u + sg_size - 1u) / sg_size; var total = 0.0; for (var i = 0u; i < n_sg; i = i + 1u) { total = total + sums[i]; } inv_rms = inverseSqrt(total / f32(D) + params.eps); } workgroupBarrier(); let r = inv_rms; for (var i = lid; i < D; i = i + ${WG}u) { ${withWeight ? 'y[row * D + i] = x[base + i] * r * w[i];' : 'y[row * D + i] = x[base + i] * r;'} } } `; } // ---------------------------------------------------------------- GEMV q4/q8 // Вариант C из бенча (vec4-загрузки + subgroupAdd), FPQ4-деквант: // вклад vec4 = scale * (sum(q*x) - zp * sum(x)). flags&1 = аккумуляция в y. // bits=4: vec4 = 32 значения; bits=8: vec4 = 16 значений. export function gemvQ4(bits = 4) { const WG = 128; return /* wgsl */ ` enable subgroups; struct Params { n : u32, k : u32, x_base_v4 : u32, flags : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var wq : array>; @group(0) @binding(2) var scales : array; @group(0) @binding(3) var zeros : array; @group(0) @binding(4) var x : array>; @group(0) @binding(5) var y : array; const WG : u32 = ${WG}u; const GRID_X : u32 = ${GRID_X}u; var sums : array; ${qHelpers(bits)} ${bits === 4 ? ` fn q_dot(word : u32, x0 : vec4, x1 : vec4) -> f32 { let q0 = vec4(f32(word & 0xFu), f32((word >> 4u) & 0xFu), f32((word >> 8u) & 0xFu), f32((word >> 12u) & 0xFu)); let q1 = vec4(f32((word >> 16u) & 0xFu), f32((word >> 20u) & 0xFu), f32((word >> 24u) & 0xFu), f32((word >> 28u) & 0xFu)); return dot(q0, x0) + dot(q1, x1); }` : ` fn q_dot8(word : u32, x0 : vec4) -> f32 { let q = vec4(f32(word & 0xFFu), f32((word >> 8u) & 0xFFu), f32((word >> 16u) & 0xFFu), f32((word >> 24u) & 0xFFu)); return dot(q, x0); }`} @compute @workgroup_size(${WG}) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_index) lid : u32, @builtin(subgroup_invocation_id) sg_lane : u32, @builtin(subgroup_size) sg_size : u32) { let row_raw = wg_id.y * GRID_X + wg_id.x; let row = min(row_raw, params.n - 1u); let nb = params.k >> 7u; // блоков по 128 let vecs_per_row = params.k >> ${bits === 4 ? '5u' : '4u'}; // значений на vec4: ${bits === 4 ? 32 : 16} let vbase = row * vecs_per_row; var acc = 0.0; for (var vi = lid; vi < vecs_per_row; vi = vi + WG) { let wv = wq[vbase + vi]; let blk = vi >> ${bits === 4 ? '2u' : '3u'}; // vec4 на блок из 128: ${bits === 4 ? 4 : 8} let s = scale_at(row, blk, nb); let zp = zp_at(row, blk, nb); ${bits === 4 ? ` let xb = params.x_base_v4 + vi * 8u; let x0 = x[xb + 0u]; let x1 = x[xb + 1u]; let x2 = x[xb + 2u]; let x3 = x[xb + 3u]; let x4 = x[xb + 4u]; let x5 = x[xb + 5u]; let x6 = x[xb + 6u]; let x7 = x[xb + 7u]; var qd = q_dot(wv.x, x0, x1) + q_dot(wv.y, x2, x3) + q_dot(wv.z, x4, x5) + q_dot(wv.w, x6, x7); let ones = vec4(1.0); let xs = dot(x0 + x1 + x2 + x3 + x4 + x5 + x6 + x7, ones);` : ` let xb = params.x_base_v4 + vi * 4u; let x0 = x[xb + 0u]; let x1 = x[xb + 1u]; let x2 = x[xb + 2u]; let x3 = x[xb + 3u]; var qd = q_dot8(wv.x, x0) + q_dot8(wv.y, x1) + q_dot8(wv.z, x2) + q_dot8(wv.w, x3); let ones = vec4(1.0); let xs = dot(x0 + x1 + x2 + x3, ones);`} acc = acc + s * (qd - zp * xs); } let sg_sum = subgroupAdd(acc); if (sg_lane == 0u) { sums[lid / sg_size] = sg_sum; } workgroupBarrier(); if (lid == 0u) { let n_sg = (WG + sg_size - 1u) / sg_size; var total = 0.0; for (var i = 0u; i < n_sg; i = i + 1u) { total = total + sums[i]; } if (row_raw < params.n) { if ((params.flags & 1u) != 0u) { y[row] = y[row] + total; } else { y[row] = total; } } } } `; } // ---------------------------------------------------------------- GEMM q4/q8 // t32 из бенча (тайл 32x32, BK=32, 16x16 потоков, регистровый блокинг 2x2), // FPQ4-деквант в shared f16. Активации в shared: f16 (по умолчанию) или f32 — // xf16=false обязателен для w2: relu^2-активация переполняет f16 (>65504). // flags&1 = аккумуляция в y. export function gemmQ4(xf16 = true, bits = 4) { const BM = 32, BN = 32, BK = 32; const XT = xf16 ? 'f16' : 'f32'; const VPW = bits === 4 ? 8 : 4; // значений на u32-слово return /* wgsl */ ` enable f16; struct Params { m : u32, n : u32, k : u32, flags : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var wq : array; @group(0) @binding(2) var scales : array; @group(0) @binding(3) var zeros : array; @group(0) @binding(4) var xin : array; @group(0) @binding(5) var y : array; const BM : u32 = ${BM}u; const BN : u32 = ${BN}u; const BK : u32 = ${BK}u; var xs : array<${XT}, ${BM * BK}>; var ws : array; ${qHelpers(bits)} @compute @workgroup_size(16, 16) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_id) l : vec3, @builtin(local_invocation_index) lidx : u32) { let row0 = wg_id.y * BM; let col0 = wg_id.x * BN; let words_per_row = params.k / ${VPW}u; let nb = params.k >> 7u; var acc : array; for (var i = 0u; i < 4u; i = i + 1u) { acc[i] = 0.0; } let n_tiles = params.k / BK; for (var t = 0u; t < n_tiles; t = t + 1u) { let k0 = t * BK; for (var i = lidx; i < BM * BK; i = i + 256u) { let r = i / BK; let c = i % BK; let gr = row0 + r; var v = ${XT}(0.0); if (gr < params.m) { v = ${XT}(xin[gr * params.k + k0 + c]); } xs[i] = v; } for (var i = lidx; i < BN * (BK / ${VPW}u); i = i + 256u) { let r = i / (BK / ${VPW}u); let wslot = i % (BK / ${VPW}u); let gcol = col0 + r; let kk = k0 + wslot * ${VPW}u; var word = 0u; var s = 0.0; var zp = 0.0; if (gcol < params.n) { word = wq[gcol * words_per_row + (kk / ${VPW}u)]; let blk = kk >> 7u; s = scale_at(gcol, blk, nb); zp = zp_at(gcol, blk, nb); } let base = r * BK + wslot * ${VPW}u; for (var j = 0u; j < ${VPW}u; j = j + 1u) { let q = f32((word >> (${bits}u * j)) & ${bits === 4 ? '0xFu' : '0xFFu'}); ws[base + j] = f16((q - zp) * s); } } workgroupBarrier(); for (var kk = 0u; kk < BK; kk = kk + 1u) { let x0 = f32(xs[l.y * BK + kk]); let x1 = f32(xs[(l.y + 16u) * BK + kk]); let w0 = f32(ws[l.x * BK + kk]); let w1 = f32(ws[(l.x + 16u) * BK + kk]); acc[0] = acc[0] + x0 * w0; acc[1] = acc[1] + x0 * w1; acc[2] = acc[2] + x1 * w0; acc[3] = acc[3] + x1 * w1; } workgroupBarrier(); } for (var i = 0u; i < 2u; i = i + 1u) { let gr = row0 + l.y + 16u * i; for (var j = 0u; j < 2u; j = j + 1u) { let gc = col0 + l.x + 16u * j; if (gr < params.m && gc < params.n) { let idx = gr * params.n + gc; if ((params.flags & 1u) != 0u) { y[idx] = y[idx] + acc[i * 2u + j]; } else { y[idx] = acc[i * 2u + j]; } } } } } `; } // ---------------------------------------------------------------- QKV post // Workgroup на (токен s, юнит u): u<16 — Q-голова, u>=16 — KV-голова g=u-16. // Q: RMSNorm + полный RoPE (1D пары 0-31 + 2D пары 32-63 по таблице) -> qbuf. // K: RMSNorm + 1D-RoPE низа -> f16 в kcache (канонический вид); при префилле // дополнительно krot16 для обеих Q-голов пары (поворот f16-значений). // V: как есть -> f16 в vcache (+ v16 при префилле). export function qkvPost() { const WG = 128; return /* wgsl */ ` enable subgroups; enable f16; struct Params { s_len : u32, p0 : u32, max_t : u32, prefill : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var qkv : array; // [S,4096] @group(0) @binding(2) var rope1 : array>; // [maxT,32] @group(0) @binding(3) var rope2 : array>; // [maxT,16,32] @group(0) @binding(4) var qbuf : array; // [16,S,128] @group(0) @binding(5) var kcache : array; // [8,maxT,128] @group(0) @binding(6) var vcache : array; @group(0) @binding(7) var krot16 : array; // [16,S,128] @group(0) @binding(8) var v16 : array; // [16,S,128] const EPS : f32 = 1.1920928955078125e-07; var vals : array; var sums : array; var inv_rms : f32; fn reduce_inv_rms(v : f32, lid : u32, sg_lane : u32, sg_size : u32) { let sg_sum = subgroupAdd(v * v); if (sg_lane == 0u) { sums[lid / sg_size] = sg_sum; } workgroupBarrier(); if (lid == 0u) { let n_sg = (128u + sg_size - 1u) / sg_size; var total = 0.0; for (var i = 0u; i < n_sg; i = i + 1u) { total = total + sums[i]; } inv_rms = inverseSqrt(total / 128.0 + EPS); } workgroupBarrier(); } @compute @workgroup_size(128) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_index) lid : u32, @builtin(subgroup_invocation_id) sg_lane : u32, @builtin(subgroup_size) sg_size : u32) { let s = wg_id.x; let u = wg_id.y; let pos = params.p0 + s; let d = lid; // Общая нормализация: юнит целиком Q или целиком K, оффсет считаем заранее, // чтобы все барьеры остались на верхнем уровне (uniform control flow). let is_q = u < 16u; let head = select(u - 16u, u, is_q); let off = select(2048u + head * 128u, head * 128u, is_q); let raw = qkv[s * 4096u + off + d]; reduce_inv_rms(raw, lid, sg_lane, sg_size); vals[d] = raw * inv_rms; workgroupBarrier(); // пара (a,b) — соседние dims; каждый поток вычисляет своё значение сам let even = (d & 1u) == 0u; let a = vals[d & ~1u]; let b = vals[d | 1u]; if (is_q) { let f = (d & 63u) >> 1u; var cs : vec2; if (d < 64u) { cs = rope1[pos * 32u + f]; } else { cs = rope2[(pos * 16u + head) * 32u + f]; } let r = select(a * cs.y + b * cs.x, a * cs.x - b * cs.y, even); qbuf[(head * params.s_len + s) * 128u + d] = r; } else { let g = head; let vv = qkv[s * 4096u + 3072u + g * 128u + d]; // канонический K: 1D-поворот низа, верх сырой var kcan = vals[d]; if (d < 64u) { let cs = rope1[pos * 32u + (d >> 1u)]; kcan = select(a * cs.y + b * cs.x, a * cs.x - b * cs.y, even); } let k16v = f16(kcan); kcache[(g * params.max_t + pos) * 128u + d] = k16v; vcache[(g * params.max_t + pos) * 128u + d] = f16(vv); if (params.prefill != 0u) { // krot16: 2D-поворот f16-округлённых канонических значений (низ = k16v); // партнёрские значения пары восстанавливаем локально с тем же округлением let a16 = f32(f16(a)); let b16 = f32(f16(b)); for (var hh = 0u; hh < 2u; hh = hh + 1u) { let h = g * 2u + hh; var r = f32(k16v); if (d >= 64u) { let cs = rope2[(pos * 16u + h) * 32u + ((d - 64u) >> 1u)]; r = select(a16 * cs.y + b16 * cs.x, a16 * cs.x - b16 * cs.y, even); } krot16[(h * params.s_len + s) * 128u + d] = f16(r); v16[(h * params.s_len + s) * 128u + d] = f16(vv); } } } } `; } // ---------------------------------------------------------------- scores // scores[h,i,j] = SCALE * dot(qbuf[h,i], krot16[h,j]); тайл 16x16. export function attnScores() { return /* wgsl */ ` enable f16; struct Params { s_len : u32, t_len : u32, p2 : u32, p3 : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var qbuf : array; // [16,S,128] @group(0) @binding(2) var krot16 : array; // [16,T,128] @group(0) @binding(3) var scores : array; // [16,S,T] var qt : array; // [16][128] var kt : array; // [16][128] @compute @workgroup_size(16, 16) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_id) l : vec3, @builtin(local_invocation_index) lidx : u32) { let h = wg_id.z; let i0 = wg_id.y * 16u; let j0 = wg_id.x * 16u; for (var i = lidx; i < 2048u; i = i + 256u) { let r = i >> 7u; let c = i & 127u; var qv = 0.0; if (i0 + r < params.s_len) { qv = qbuf[(h * params.s_len + i0 + r) * 128u + c]; } qt[i] = qv; var kv = f16(0.0); if (j0 + r < params.t_len) { kv = krot16[(h * params.t_len + j0 + r) * 128u + c]; } kt[i] = kv; } workgroupBarrier(); let i = i0 + l.y; let j = j0 + l.x; var acc = 0.0; for (var c = 0u; c < 128u; c = c + 4u) { acc = acc + qt[l.y * 128u + c] * f32(kt[l.x * 128u + c]) + qt[l.y * 128u + c + 1u] * f32(kt[l.x * 128u + c + 1u]) + qt[l.y * 128u + c + 2u] * f32(kt[l.x * 128u + c + 2u]) + qt[l.y * 128u + c + 3u] * f32(kt[l.x * 128u + c + 3u]); } if (i < params.s_len && j < params.t_len) { scores[(h * params.s_len + i) * params.t_len + j] = acc * ${SCALE_ATTN}; } } `; } // ---------------------------------------------------------------- softmax+gate // Workgroup на строку (h, i): маска процедурная (bidir-префикс + каузальность), // P := softmax * gate (гейт вшит в P), lse/gate — в дебаг-буферы. export function attnSoftmaxGate() { const WG = 256; return /* wgsl */ ` enable subgroups; struct Params { s_len : u32, t_len : u32, bidir_end : u32, layer : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var scores : array; // [16,S,T] -> P @group(0) @binding(2) var gate_bias : array; // [28*16] @group(0) @binding(3) var lse_out : array; // [16,S] @group(0) @binding(4) var gate_out : array; // [16,S] const WG : u32 = ${WG}u; const NEG : f32 = ${NEG_MAX}; var red : array; var m_sh : f32; var l_sh : f32; fn reduce_max(v : f32, lid : u32, sg_lane : u32, sg_size : u32) -> f32 { let sg_m = subgroupMax(v); if (sg_lane == 0u) { red[lid / sg_size] = sg_m; } workgroupBarrier(); if (lid == 0u) { let n_sg = (WG + sg_size - 1u) / sg_size; var m = NEG; for (var i = 0u; i < n_sg; i = i + 1u) { m = max(m, red[i]); } m_sh = m; } workgroupBarrier(); return m_sh; } fn reduce_sum(v : f32, lid : u32, sg_lane : u32, sg_size : u32) -> f32 { let sg_s = subgroupAdd(v); if (sg_lane == 0u) { red[lid / sg_size] = sg_s; } workgroupBarrier(); if (lid == 0u) { let n_sg = (WG + sg_size - 1u) / sg_size; var s = 0.0; for (var i = 0u; i < n_sg; i = i + 1u) { s = s + red[i]; } l_sh = s; } workgroupBarrier(); return l_sh; } @compute @workgroup_size(${WG}) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_index) lid : u32, @builtin(subgroup_invocation_id) sg_lane : u32, @builtin(subgroup_size) sg_size : u32) { let i = wg_id.x; let h = wg_id.y; let row = (h * params.s_len + i) * params.t_len; let prefix = i < params.bidir_end; var m = NEG; for (var j = lid; j < params.t_len; j = j + WG) { let allow = (j <= i) || (prefix && j < params.bidir_end); if (allow) { m = max(m, scores[row + j]); } } m = reduce_max(m, lid, sg_lane, sg_size); var lsum = 0.0; for (var j = lid; j < params.t_len; j = j + WG) { let allow = (j <= i) || (prefix && j < params.bidir_end); if (allow) { lsum = lsum + exp(scores[row + j] - m); } } lsum = reduce_sum(lsum, lid, sg_lane, sg_size); let lse = m + log(lsum); let gate = 1.0 / (1.0 + exp(-(lse - gate_bias[params.layer * 16u + h]))); let scale = gate / lsum; for (var j = lid; j < params.t_len; j = j + WG) { let allow = (j <= i) || (prefix && j < params.bidir_end); var p = 0.0; if (allow) { p = exp(scores[row + j] - m) * scale; } scores[row + j] = p; } if (lid == 0u) { lse_out[h * params.s_len + i] = lse; gate_out[h * params.s_len + i] = gate; } } `; } // ---------------------------------------------------------------- P @ V // o[i, h*128+d] = sum_j P[h,i,j] * v16[h,j,d]; выход токен-мажорный [S,2048]. export function attnPV() { return /* wgsl */ ` enable f16; struct Params { s_len : u32, t_len : u32, p2 : u32, p3 : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var p : array; // [16,S,T] (гейт вшит) @group(0) @binding(2) var v16 : array; // [16,T,128] @group(0) @binding(3) var o : array; // [S,2048] var pt : array; // [16][32] var vt : array; // [32][16] @compute @workgroup_size(16, 16) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_id) l : vec3, @builtin(local_invocation_index) lidx : u32) { let h = wg_id.z; let i0 = wg_id.y * 16u; let d0 = wg_id.x * 16u; // 8 тайлов по dim var acc = 0.0; let n_tiles = (params.t_len + 31u) / 32u; for (var t = 0u; t < n_tiles; t = t + 1u) { let j0 = t * 32u; for (var i = lidx; i < 512u; i = i + 256u) { let r = i >> 5u; let c = i & 31u; // pt[r][c] var pv = 0.0; if (i0 + r < params.s_len && j0 + c < params.t_len) { pv = p[(h * params.s_len + i0 + r) * params.t_len + j0 + c]; } pt[i] = pv; let jr = i >> 4u; let dc = i & 15u; // vt[jr][dc] var vv = f16(0.0); if (j0 + jr < params.t_len) { vv = v16[(h * params.t_len + j0 + jr) * 128u + d0 + dc]; } vt[i] = vv; } workgroupBarrier(); for (var j = 0u; j < 32u; j = j + 1u) { acc = acc + pt[l.y * 32u + j] * f32(vt[j * 16u + l.x]); } workgroupBarrier(); } let i = i0 + l.y; if (i < params.s_len) { o[i * 2048u + h * 128u + d0 + l.x] = acc; } } `; } // ---------------------------------------------------------------- decode attention // Workgroup на KV-группу g (8 wg, 128 потоков = 4 сабгруппы по 32). // Обе Q-головы пары за проход; k из 8-голового кэша, 2D-поворот по таблице // (текст/декод — identity), онлайн-softmax, LSE-гейт. S=1. export function attnDecode() { return /* wgsl */ ` enable subgroups; enable f16; struct Params { t_len : u32, max_t : u32, layer : u32, p3 : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var qbuf : array; // [16,1,128] @group(0) @binding(2) var kcache : array; // [8,maxT,128] @group(0) @binding(3) var vcache : array; @group(0) @binding(4) var rope2 : array>; // [maxT,16,32] @group(0) @binding(5) var gate_bias : array; @group(0) @binding(6) var o : array; // [1,2048] const NEG : f32 = ${NEG_MAX}; var q2 : array; // [2][128] var sc : array; // [2 головы][4 ключа] @compute @workgroup_size(128) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_index) lid : u32, @builtin(subgroup_invocation_id) sg_lane : u32, @builtin(subgroup_size) sg_size : u32) { let g = wg_id.x; let h0 = g * 2u; q2[lid] = qbuf[(h0 * 1u) * 128u + lid]; q2[128u + lid] = qbuf[((h0 + 1u) * 1u) * 128u + lid]; workgroupBarrier(); var m0 = NEG; var l0 = 0.0; var acc0 = 0.0; var m1 = NEG; var l1 = 0.0; var acc1 = 0.0; let d = lid; // dim этого потока для PV let sg = lid / sg_size; // сабгруппа = ключ в чанке (sg_size=32) let lane = sg_lane; let n_chunks = (params.t_len + 3u) / 4u; for (var ch = 0u; ch < n_chunks; ch = ch + 1u) { let t = ch * 4u + sg; // скоры чанка: сабгруппа sg считает ключ t; lane покрывает dims lane*4..+3 var p0dot = 0.0; var p1dot = 0.0; if (t < params.t_len) { let kb = (g * params.max_t + t) * 128u + lane * 4u; let kv = vec4(f32(kcache[kb]), f32(kcache[kb + 1u]), f32(kcache[kb + 2u]), f32(kcache[kb + 3u])); var k0 = kv; var k1 = kv; if (lane >= 16u) { // dims 64..127: 2 пары на поток let fb = (lane * 4u - 64u) >> 1u; // первая пара let cs00 = rope2[(t * 16u + h0) * 32u + fb]; let cs01 = rope2[(t * 16u + h0) * 32u + fb + 1u]; k0 = vec4(kv.x * cs00.x - kv.y * cs00.y, kv.x * cs00.y + kv.y * cs00.x, kv.z * cs01.x - kv.w * cs01.y, kv.z * cs01.y + kv.w * cs01.x); let cs10 = rope2[(t * 16u + h0 + 1u) * 32u + fb]; let cs11 = rope2[(t * 16u + h0 + 1u) * 32u + fb + 1u]; k1 = vec4(kv.x * cs10.x - kv.y * cs10.y, kv.x * cs10.y + kv.y * cs10.x, kv.z * cs11.x - kv.w * cs11.y, kv.z * cs11.y + kv.w * cs11.x); } let qb = lane * 4u; let q0v = vec4(q2[qb], q2[qb + 1u], q2[qb + 2u], q2[qb + 3u]); let q1v = vec4(q2[128u + qb], q2[129u + qb], q2[130u + qb], q2[131u + qb]); p0dot = dot(k0, q0v); p1dot = dot(k1, q1v); } let s0 = subgroupAdd(p0dot); let s1 = subgroupAdd(p1dot); if (lane == 0u) { sc[sg] = select(NEG, s0 * ${SCALE_ATTN}, t < params.t_len); sc[4u + sg] = select(NEG, s1 * ${SCALE_ATTN}, t < params.t_len); } workgroupBarrier(); // онлайн-softmax (каждый поток дублирует скалярную математику — равномерно) let c0 = vec4(sc[0], sc[1], sc[2], sc[3]); let c1 = vec4(sc[4], sc[5], sc[6], sc[7]); workgroupBarrier(); // sc можно перезаписывать в след. чанке let nm0 = max(m0, max(max(c0.x, c0.y), max(c0.z, c0.w))); let nm1 = max(m1, max(max(c1.x, c1.y), max(c1.z, c1.w))); let r0 = exp(m0 - nm0); let r1 = exp(m1 - nm1); let e0 = exp(c0 - vec4(nm0)); let e1 = exp(c1 - vec4(nm1)); l0 = l0 * r0 + e0.x + e0.y + e0.z + e0.w; l1 = l1 * r1 + e1.x + e1.y + e1.z + e1.w; m0 = nm0; m1 = nm1; let vb = (g * params.max_t + ch * 4u) * 128u + d; let v0 = f32(vcache[vb]); let v1 = f32(vcache[vb + 128u]); let v2 = f32(vcache[vb + 256u]); let v3 = f32(vcache[vb + 384u]); acc0 = acc0 * r0 + e0.x * v0 + e0.y * v1 + e0.z * v2 + e0.w * v3; acc1 = acc1 * r1 + e1.x * v0 + e1.y * v1 + e1.z * v2 + e1.w * v3; } let lse0 = m0 + log(l0); let lse1 = m1 + log(l1); let g0 = 1.0 / (1.0 + exp(-(lse0 - gate_bias[params.layer * 16u + h0]))); let g1 = 1.0 / (1.0 + exp(-(lse1 - gate_bias[params.layer * 16u + h0 + 1u]))); o[h0 * 128u + d] = acc0 / l0 * g0; o[(h0 + 1u) * 128u + d] = acc1 / l1 * g1; } `; } // ---------------------------------------------------------------- decode attention, multi-query // Батчевый verify спекулятивного декода: k запросов (wg_id.y) поверх KV-кэша, // та же схема, что attnDecode (онлайн-softmax, GQA, 2D-поворот ключей при // чтении), но qbuf [16,k,128], каузальность t_len_i = t0 + qi + 1 // (kv-строки драфта уже записаны qkvPost'ом на p0=t0), выход o [k,2048]. export function attnDecodeMq() { return /* wgsl */ ` enable subgroups; enable f16; struct Params { t0 : u32, max_t : u32, layer : u32, s_len : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var qbuf : array; // [16,s_len,128] @group(0) @binding(2) var kcache : array; // [8,maxT,128] @group(0) @binding(3) var vcache : array; @group(0) @binding(4) var rope2 : array>; // [maxT,16,32] @group(0) @binding(5) var gate_bias : array; @group(0) @binding(6) var o : array; // [s_len,2048] const NEG : f32 = ${NEG_MAX}; var q2 : array; // [2][128] var sc : array; // [2 головы][4 ключа] @compute @workgroup_size(128) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_index) lid : u32, @builtin(subgroup_invocation_id) sg_lane : u32, @builtin(subgroup_size) sg_size : u32) { let g = wg_id.x; let qi = wg_id.y; let t_len = params.t0 + qi + 1u; // каузальность внутри драфта let h0 = g * 2u; q2[lid] = qbuf[(h0 * params.s_len + qi) * 128u + lid]; q2[128u + lid] = qbuf[((h0 + 1u) * params.s_len + qi) * 128u + lid]; workgroupBarrier(); var m0 = NEG; var l0 = 0.0; var acc0 = 0.0; var m1 = NEG; var l1 = 0.0; var acc1 = 0.0; let d = lid; let sg = lid / sg_size; let lane = sg_lane; let n_chunks = (t_len + 3u) / 4u; for (var ch = 0u; ch < n_chunks; ch = ch + 1u) { let t = ch * 4u + sg; var p0dot = 0.0; var p1dot = 0.0; if (t < t_len) { let kb = (g * params.max_t + t) * 128u + lane * 4u; let kv = vec4(f32(kcache[kb]), f32(kcache[kb + 1u]), f32(kcache[kb + 2u]), f32(kcache[kb + 3u])); var k0 = kv; var k1 = kv; if (lane >= 16u) { let fb = (lane * 4u - 64u) >> 1u; let cs00 = rope2[(t * 16u + h0) * 32u + fb]; let cs01 = rope2[(t * 16u + h0) * 32u + fb + 1u]; k0 = vec4(kv.x * cs00.x - kv.y * cs00.y, kv.x * cs00.y + kv.y * cs00.x, kv.z * cs01.x - kv.w * cs01.y, kv.z * cs01.y + kv.w * cs01.x); let cs10 = rope2[(t * 16u + h0 + 1u) * 32u + fb]; let cs11 = rope2[(t * 16u + h0 + 1u) * 32u + fb + 1u]; k1 = vec4(kv.x * cs10.x - kv.y * cs10.y, kv.x * cs10.y + kv.y * cs10.x, kv.z * cs11.x - kv.w * cs11.y, kv.z * cs11.y + kv.w * cs11.x); } let qb = lane * 4u; let q0v = vec4(q2[qb], q2[qb + 1u], q2[qb + 2u], q2[qb + 3u]); let q1v = vec4(q2[128u + qb], q2[129u + qb], q2[130u + qb], q2[131u + qb]); p0dot = dot(k0, q0v); p1dot = dot(k1, q1v); } let s0 = subgroupAdd(p0dot); let s1 = subgroupAdd(p1dot); if (lane == 0u) { sc[sg] = select(NEG, s0 * ${SCALE_ATTN}, t < t_len); sc[4u + sg] = select(NEG, s1 * ${SCALE_ATTN}, t < t_len); } workgroupBarrier(); let c0 = vec4(sc[0], sc[1], sc[2], sc[3]); let c1 = vec4(sc[4], sc[5], sc[6], sc[7]); workgroupBarrier(); let nm0 = max(m0, max(max(c0.x, c0.y), max(c0.z, c0.w))); let nm1 = max(m1, max(max(c1.x, c1.y), max(c1.z, c1.w))); let r0 = exp(m0 - nm0); let r1 = exp(m1 - nm1); let e0 = exp(c0 - vec4(nm0)); let e1 = exp(c1 - vec4(nm1)); l0 = l0 * r0 + e0.x + e0.y + e0.z + e0.w; l1 = l1 * r1 + e1.x + e1.y + e1.z + e1.w; m0 = nm0; m1 = nm1; let vb = (g * params.max_t + ch * 4u) * 128u + d; let v0 = f32(vcache[vb]); let v1 = f32(vcache[vb + 128u]); let v2 = f32(vcache[vb + 256u]); let v3 = f32(vcache[vb + 384u]); acc0 = acc0 * r0 + e0.x * v0 + e0.y * v1 + e0.z * v2 + e0.w * v3; acc1 = acc1 * r1 + e1.x * v0 + e1.y * v1 + e1.z * v2 + e1.w * v3; } let lse0 = m0 + log(l0); let lse1 = m1 + log(l1); let g0 = 1.0 / (1.0 + exp(-(lse0 - gate_bias[params.layer * 16u + h0]))); let g1 = 1.0 / (1.0 + exp(-(lse1 - gate_bias[params.layer * 16u + h0 + 1u]))); o[qi * 2048u + h0 * 128u + d] = acc0 / l0 * g0; o[qi * 2048u + (h0 + 1u) * 128u + d] = acc1 / l1 * g1; } `; } // ---------------------------------------------------------------- MLP-активация // y[s,i] = relu(u[s,i])^2 * u[s,3072+i]; u [S,6144] = [gate|up]. export function mlpAct() { return /* wgsl */ ` struct Params { total : u32, p1 : u32, p2 : u32, p3 : u32 } // total = S*3072 @group(0) @binding(0) var params : Params; @group(0) @binding(1) var u : array; @group(0) @binding(2) var y : array; @compute @workgroup_size(256) fn main(@builtin(global_invocation_id) gid : vec3) { let i = gid.x; if (i >= params.total) { return; } let s = i / 3072u; let c = i % 3072u; let gt = max(u[s * 6144u + c], 0.0); y[i] = gt * gt * u[s * 6144u + 3072u + c]; } `; } // ---------------------------------------------------------------- argmax // Двухстадийный argmax по f32-массиву (vocab 65536): стадия 1 — 256 wg по 256 // элементов -> частичные (val,idx); стадия 2 — 1 wg сводит 256 частичных, // пишет индекс u32 в out[row]. Ридбек 4 байта на строку. Строки — wg_id.y // (стадия 1) / wg_id.x (стадия 2); при dispatch (256,1)/(1) — прежний декод. export function argmaxStage1() { return /* wgsl */ ` struct Params { n : u32, p1 : u32, p2 : u32, p3 : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var src : array; @group(0) @binding(2) var pval : array; @group(0) @binding(3) var pidx : array; var sv : array; var si : array; @compute @workgroup_size(256) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_index) lid : u32) { let row = wg_id.y; let i = wg_id.x * 256u + lid; var v = -3.4e38; if (i < params.n) { v = src[row * params.n + i]; } sv[lid] = v; si[lid] = i; workgroupBarrier(); var stride = 128u; loop { if (stride == 0u) { break; } if (lid < stride) { if (sv[lid + stride] > sv[lid]) { sv[lid] = sv[lid + stride]; si[lid] = si[lid + stride]; } } workgroupBarrier(); stride = stride >> 1u; } if (lid == 0u) { pval[row * 256u + wg_id.x] = sv[0]; pidx[row * 256u + wg_id.x] = si[0]; } } `; } export function argmaxStage2() { return /* wgsl */ ` struct Params { n : u32, p1 : u32, p2 : u32, p3 : u32 } // n = число частичных @group(0) @binding(0) var params : Params; @group(0) @binding(1) var pval : array; @group(0) @binding(2) var pidx : array; @group(0) @binding(3) var out : array; var sv : array; var si : array; @compute @workgroup_size(256) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_index) lid : u32) { let row = wg_id.x; var v = -3.4e38; var ix = 0u; if (lid < params.n) { v = pval[row * 256u + lid]; ix = pidx[row * 256u + lid]; } sv[lid] = v; si[lid] = ix; workgroupBarrier(); var stride = 128u; loop { if (stride == 0u) { break; } if (lid < stride) { if (sv[lid + stride] > sv[lid]) { sv[lid] = sv[lid + stride]; si[lid] = si[lid + stride]; } } workgroupBarrier(); stride = stride >> 1u; } if (lid == 0u) { out[row] = si[0]; } } `; } // ---------------------------------------------------------------- gather embed // x[t, d] = f32(embed[ids[t] * 1024 + d]); wg на токен, 256 потоков x4. // ids — u32-буфер (для декода это выход argmax — CPU не участвует). // Сентинел 0xFFFFFFFF — строку не трогать (туда CPU уже записал coord/size- // эмбеддинг; нужно для спекулятивного verify одним submit'ом). export function gatherEmbed() { return /* wgsl */ ` enable f16; struct Params { count : u32, p1 : u32, p2 : u32, p3 : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var ids : array; @group(0) @binding(2) var embed : array; @group(0) @binding(3) var x : array; @compute @workgroup_size(256) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_index) lid : u32) { let t = wg_id.x; if (t >= params.count) { return; } let id = ids[t]; if (id == 0xFFFFFFFFu) { return; } for (var d = lid; d < 1024u; d = d + 256u) { x[t * 1024u + d] = f32(embed[id * 1024u + d]); } } `; } // ---------------------------------------------------------------- GEMV f16-веса // Каркас варианта C: vec4 = 8 f16-весов через unpack2x16float. // Для голов декода (coord/size_decoder W1/W2). flags&1 = аккумуляция. export function gemvF16W() { const WG = 128; return /* wgsl */ ` enable subgroups; struct Params { n : u32, k : u32, x_base_v4 : u32, flags : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var w : array>; @group(0) @binding(2) var x : array>; @group(0) @binding(3) var y : array; const WG : u32 = ${WG}u; const GRID_X : u32 = ${GRID_X}u; var sums : array; @compute @workgroup_size(${WG}) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_index) lid : u32, @builtin(subgroup_invocation_id) sg_lane : u32, @builtin(subgroup_size) sg_size : u32) { let row_raw = wg_id.y * GRID_X + wg_id.x; let row = min(row_raw, params.n - 1u); let vecs_per_row = params.k >> 3u; // 8 f16 на vec4 let vbase = row * vecs_per_row; var acc = 0.0; for (var vi = lid; vi < vecs_per_row; vi = vi + WG) { let wv = w[vbase + vi]; let xb = params.x_base_v4 + vi * 2u; let x0 = x[xb]; let x1 = x[xb + 1u]; let a = unpack2x16float(wv.x); let b = unpack2x16float(wv.y); let c = unpack2x16float(wv.z); let d = unpack2x16float(wv.w); acc = acc + dot(vec4(a, b), x0) + dot(vec4(c, d), x1); } let sg_sum = subgroupAdd(acc); if (sg_lane == 0u) { sums[lid / sg_size] = sg_sum; } workgroupBarrier(); if (lid == 0u) { let n_sg = (WG + sg_size - 1u) / sg_size; var total = 0.0; for (var i = 0u; i < n_sg; i = i + 1u) { total = total + sums[i]; } if (row_raw < params.n) { if ((params.flags & 1u) != 0u) { y[row] = y[row] + total; } else { y[row] = total; } } } } `; } // ---------------------------------------------------------------- GEMM f16-веса // Для img_projector: Y[m, n] = dot(X[m], W[n]); W f16 [N,K] row-major, // X f32 [M,K]; out_base — оффсет строк в выходном буфере (скаттер в x-буфер). export function gemmF16W() { return /* wgsl */ ` enable f16; struct Params { m : u32, n : u32, k : u32, out_base : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var w : array; @group(0) @binding(2) var xin : array; @group(0) @binding(3) var y : array; var at : array; // [16][64] var wt : array; // [16][64] @compute @workgroup_size(16, 16) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_id) l : vec3, @builtin(local_invocation_index) lidx : u32) { let m0 = wg_id.y * 16u; let n0 = wg_id.x * 16u; var acc = 0.0; let n_tiles = (params.k + 63u) / 64u; for (var t = 0u; t < n_tiles; t = t + 1u) { let k0 = t * 64u; for (var i = lidx; i < 1024u; i = i + 256u) { let r = i >> 6u; let c = i & 63u; var av = 0.0; if (m0 + r < params.m && k0 + c < params.k) { av = xin[(m0 + r) * params.k + k0 + c]; } at[i] = av; var wv = f16(0.0); if (n0 + r < params.n && k0 + c < params.k) { wv = w[(n0 + r) * params.k + k0 + c]; } wt[i] = wv; } workgroupBarrier(); for (var c = 0u; c < 64u; c = c + 1u) { acc = acc + at[l.y * 64u + c] * f32(wt[l.x * 64u + c]); } workgroupBarrier(); } let m = m0 + l.y; let n = n0 + l.x; if (m < params.m && n < params.n) { y[(params.out_base + m) * params.n + n] = acc; } } `; } // ---------------------------------------------------------------- relu^2 // y = relu(x)^2 — активация голов coord/size_decoder (без гейта). export function relu2() { return /* wgsl */ ` struct Params { total : u32, p1 : u32, p2 : u32, p3 : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var x : array; @group(0) @binding(2) var y : array; @compute @workgroup_size(256) fn main(@builtin(global_invocation_id) gid : vec3) { let i = gid.x; if (i >= params.total) { return; } let r = max(x[i], 0.0); y[i] = r * r; } `; } // ---------------------------------------------------------------- GEMM sgmat // Опциональный быстрый префилл-GEMM на chromium-experimental-subgroup-matrix // (схема ORT/llama.cpp, сверена в engine/bench: тайл 32x64, BK=32, 128 потоков // = 4 сабгруппы, каждая — субтайл 16x32 из 8x8 f16-матриц). FPQ4-деквант q4/q8. // ВНИМАНИЕ: аккумуляция f16 — НЕ использовать для w2 (relu^2 переполняет f16); // включается автодетектом + рантайм-smoke против t32 (model.js). export function gemmQ4Sg(bits = 4) { const BM = 32, BN = 64, BK = 32; const VPW = bits === 4 ? 8 : 4; // 4 сабгруппы: субтайл 16x32 = 2x4 матриц 8x8 let decls = '', loads = '', mads = '', stores = ''; for (let i = 0; i < 2; i++) for (let j = 0; j < 4; j++) decls += ` var c${i}${j} : subgroup_matrix_result;\n`; for (let i = 0; i < 2; i++) loads += ` let a${i} = subgroupMatrixLoad>(&tile_a, a_off + ${i * 8}u * BK, false, BK);\n`; for (let j = 0; j < 4; j++) loads += ` let b${j} = subgroupMatrixLoad>(&tile_b, b_off + ${j * 8}u * BK, true, BK);\n`; for (let i = 0; i < 2; i++) for (let j = 0; j < 4; j++) mads += ` c${i}${j} = subgroupMatrixMultiplyAccumulate(a${i}, b${j}, c${i}${j});\n`; for (let i = 0; i < 2; i++) for (let j = 0; j < 4; j++) stores += ` subgroupMatrixStore(&scratch[subtile_id][${i * 4 + j}], 0u, c${i}${j}, false, 8u);\n`; return /* wgsl */ ` enable subgroups; enable chromium_experimental_subgroup_matrix; enable f16; diagnostic (off, chromium.subgroup_matrix_uniformity); struct Params { m : u32, n : u32, k : u32, flags : u32 } @group(0) @binding(0) var params : Params; @group(0) @binding(1) var wq : array; @group(0) @binding(2) var scales : array; @group(0) @binding(3) var zeros : array; @group(0) @binding(4) var xin : array; @group(0) @binding(5) var y : array; const BM : u32 = ${BM}u; const BN : u32 = ${BN}u; const BK : u32 = ${BK}u; var tile_a : array; var tile_b : array; var scratch : array, 8>, 4>; ${qHelpers(bits)} @compute @workgroup_size(128) fn main(@builtin(workgroup_id) wg_id : vec3, @builtin(local_invocation_index) lid : u32, @builtin(subgroup_size) sg_size : u32) { let row0 = wg_id.y * BM; let col0 = wg_id.x * BN; let words_per_row = params.k / ${VPW}u; let nb = params.k >> 7u; let subtile_id = lid / sg_size; // предполагается sg_size == 32 let base_a = (subtile_id % 2u) * 16u; let base_b = (subtile_id / 2u) * 32u; ${decls} for (var kidx = 0u; kidx < params.k; kidx = kidx + BK) { for (var i = lid; i < BM * BK; i = i + 128u) { let r = i / BK; let c = i % BK; let gr = row0 + r; var v = f16(0.0); if (gr < params.m) { v = f16(xin[gr * params.k + kidx + c]); } tile_a[i] = v; } for (var i = lid; i < BN * (BK / ${VPW}u); i = i + 128u) { let r = i / (BK / ${VPW}u); let wslot = i % (BK / ${VPW}u); let gcol = col0 + r; let kk = kidx + wslot * ${VPW}u; var word = 0u; var s = 0.0; var zp = 0.0; if (gcol < params.n) { word = wq[gcol * words_per_row + (kk / ${VPW}u)]; let blk = kk >> 7u; s = scale_at(gcol, blk, nb); zp = zp_at(gcol, blk, nb); } let tbase = r * BK + wslot * ${VPW}u; for (var j = 0u; j < ${VPW}u; j = j + 1u) { let q = f32((word >> (${bits}u * j)) & ${bits === 4 ? '0xFu' : '0xFFu'}); tile_b[tbase + j] = f16((q - zp) * s); } } workgroupBarrier(); for (var step = 0u; step < BK; step = step + 8u) { let a_off = base_a * BK + step; let b_off = base_b * BK + step; ${loads} ${mads} } workgroupBarrier(); } ${stores} workgroupBarrier(); for (var i = lid; i < BM * BN; i = i + 128u) { let r = i / BN; let c = i % BN; let st = (r / 16u) + 2u * (c / 32u); let rr = r % 16u; let cc = c % 32u; let mi = (rr / 8u) * 4u + (cc / 8u); let gr = row0 + r; let gc = col0 + c; if (gr < params.m && gc < params.n) { let acc = f32(scratch[st][mi][(rr % 8u) * 8u + (cc % 8u)]); let idx = gr * params.n + gc; if ((params.flags & 1u) != 0u) { y[idx] = y[idx] + acc; } else { y[idx] = acc; } } } } `; }