borkiss's picture
Upload folder using huggingface_hub
2268f8e verified
Raw
History Blame Contribute Delete
46.1 kB
// 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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> x : array<f32>;
@group(0) @binding(2) var<storage, read_write> y : array<f32>;
${withWeight ? '@group(0) @binding(3) var<storage, read> w : array<f32>;' : ''}
const D : u32 = ${D}u;
var<workgroup> sums : array<f32, ${WG / 4}>;
var<workgroup> inv_rms : f32;
@compute @workgroup_size(${WG})
fn main(@builtin(workgroup_id) wg_id : vec3<u32>,
@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<u32>-загрузки + subgroupAdd), FPQ4-деквант:
// вклад vec4 = scale * (sum(q*x) - zp * sum(x)). flags&1 = аккумуляция в y.
// bits=4: vec4<u32> = 32 значения; bits=8: vec4<u32> = 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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> wq : array<vec4<u32>>;
@group(0) @binding(2) var<storage, read> scales : array<u32>;
@group(0) @binding(3) var<storage, read> zeros : array<u32>;
@group(0) @binding(4) var<storage, read> x : array<vec4<f32>>;
@group(0) @binding(5) var<storage, read_write> y : array<f32>;
const WG : u32 = ${WG}u;
const GRID_X : u32 = ${GRID_X}u;
var<workgroup> sums : array<f32, ${WG / 4}>;
${qHelpers(bits)}
${bits === 4 ? `
fn q_dot(word : u32, x0 : vec4<f32>, x1 : vec4<f32>) -> f32 {
let q0 = vec4<f32>(f32(word & 0xFu), f32((word >> 4u) & 0xFu),
f32((word >> 8u) & 0xFu), f32((word >> 12u) & 0xFu));
let q1 = vec4<f32>(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>) -> f32 {
let q = vec4<f32>(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<u32>,
@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<u32>: ${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<f32>(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<f32>(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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> wq : array<u32>;
@group(0) @binding(2) var<storage, read> scales : array<u32>;
@group(0) @binding(3) var<storage, read> zeros : array<u32>;
@group(0) @binding(4) var<storage, read> xin : array<f32>;
@group(0) @binding(5) var<storage, read_write> y : array<f32>;
const BM : u32 = ${BM}u; const BN : u32 = ${BN}u; const BK : u32 = ${BK}u;
var<workgroup> xs : array<${XT}, ${BM * BK}>;
var<workgroup> ws : array<f16, ${BN * BK}>;
${qHelpers(bits)}
@compute @workgroup_size(16, 16)
fn main(@builtin(workgroup_id) wg_id : vec3<u32>,
@builtin(local_invocation_id) l : vec3<u32>,
@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<f32, 4>;
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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> qkv : array<f32>; // [S,4096]
@group(0) @binding(2) var<storage, read> rope1 : array<vec2<f32>>; // [maxT,32]
@group(0) @binding(3) var<storage, read> rope2 : array<vec2<f32>>; // [maxT,16,32]
@group(0) @binding(4) var<storage, read_write> qbuf : array<f32>; // [16,S,128]
@group(0) @binding(5) var<storage, read_write> kcache : array<f16>; // [8,maxT,128]
@group(0) @binding(6) var<storage, read_write> vcache : array<f16>;
@group(0) @binding(7) var<storage, read_write> krot16 : array<f16>; // [16,S,128]
@group(0) @binding(8) var<storage, read_write> v16 : array<f16>; // [16,S,128]
const EPS : f32 = 1.1920928955078125e-07;
var<workgroup> vals : array<f32, 128>;
var<workgroup> sums : array<f32, 4>;
var<workgroup> 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<u32>,
@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<f32>;
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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> qbuf : array<f32>; // [16,S,128]
@group(0) @binding(2) var<storage, read> krot16 : array<f16>; // [16,T,128]
@group(0) @binding(3) var<storage, read_write> scores : array<f32>; // [16,S,T]
var<workgroup> qt : array<f32, 2048>; // [16][128]
var<workgroup> kt : array<f16, 2048>; // [16][128]
@compute @workgroup_size(16, 16)
fn main(@builtin(workgroup_id) wg_id : vec3<u32>,
@builtin(local_invocation_id) l : vec3<u32>,
@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<uniform> params : Params;
@group(0) @binding(1) var<storage, read_write> scores : array<f32>; // [16,S,T] -> P
@group(0) @binding(2) var<storage, read> gate_bias : array<f32>; // [28*16]
@group(0) @binding(3) var<storage, read_write> lse_out : array<f32>; // [16,S]
@group(0) @binding(4) var<storage, read_write> gate_out : array<f32>; // [16,S]
const WG : u32 = ${WG}u;
const NEG : f32 = ${NEG_MAX};
var<workgroup> red : array<f32, ${WG / 4}>;
var<workgroup> m_sh : f32;
var<workgroup> 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<u32>,
@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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> p : array<f32>; // [16,S,T] (гейт вшит)
@group(0) @binding(2) var<storage, read> v16 : array<f16>; // [16,T,128]
@group(0) @binding(3) var<storage, read_write> o : array<f32>; // [S,2048]
var<workgroup> pt : array<f32, 512>; // [16][32]
var<workgroup> vt : array<f16, 512>; // [32][16]
@compute @workgroup_size(16, 16)
fn main(@builtin(workgroup_id) wg_id : vec3<u32>,
@builtin(local_invocation_id) l : vec3<u32>,
@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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> qbuf : array<f32>; // [16,1,128]
@group(0) @binding(2) var<storage, read> kcache : array<f16>; // [8,maxT,128]
@group(0) @binding(3) var<storage, read> vcache : array<f16>;
@group(0) @binding(4) var<storage, read> rope2 : array<vec2<f32>>; // [maxT,16,32]
@group(0) @binding(5) var<storage, read> gate_bias : array<f32>;
@group(0) @binding(6) var<storage, read_write> o : array<f32>; // [1,2048]
const NEG : f32 = ${NEG_MAX};
var<workgroup> q2 : array<f32, 256>; // [2][128]
var<workgroup> sc : array<f32, 8>; // [2 головы][4 ключа]
@compute @workgroup_size(128)
fn main(@builtin(workgroup_id) wg_id : vec3<u32>,
@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>(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<f32>(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<f32>(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<f32>(q2[qb], q2[qb + 1u], q2[qb + 2u], q2[qb + 3u]);
let q1v = vec4<f32>(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<f32>(sc[0], sc[1], sc[2], sc[3]);
let c1 = vec4<f32>(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<f32>(nm0));
let e1 = exp(c1 - vec4<f32>(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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> qbuf : array<f32>; // [16,s_len,128]
@group(0) @binding(2) var<storage, read> kcache : array<f16>; // [8,maxT,128]
@group(0) @binding(3) var<storage, read> vcache : array<f16>;
@group(0) @binding(4) var<storage, read> rope2 : array<vec2<f32>>; // [maxT,16,32]
@group(0) @binding(5) var<storage, read> gate_bias : array<f32>;
@group(0) @binding(6) var<storage, read_write> o : array<f32>; // [s_len,2048]
const NEG : f32 = ${NEG_MAX};
var<workgroup> q2 : array<f32, 256>; // [2][128]
var<workgroup> sc : array<f32, 8>; // [2 головы][4 ключа]
@compute @workgroup_size(128)
fn main(@builtin(workgroup_id) wg_id : vec3<u32>,
@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>(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<f32>(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<f32>(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<f32>(q2[qb], q2[qb + 1u], q2[qb + 2u], q2[qb + 3u]);
let q1v = vec4<f32>(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<f32>(sc[0], sc[1], sc[2], sc[3]);
let c1 = vec4<f32>(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<f32>(nm0));
let e1 = exp(c1 - vec4<f32>(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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> u : array<f32>;
@group(0) @binding(2) var<storage, read_write> y : array<f32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> src : array<f32>;
@group(0) @binding(2) var<storage, read_write> pval : array<f32>;
@group(0) @binding(3) var<storage, read_write> pidx : array<u32>;
var<workgroup> sv : array<f32, 256>;
var<workgroup> si : array<u32, 256>;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg_id : vec3<u32>,
@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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> pval : array<f32>;
@group(0) @binding(2) var<storage, read> pidx : array<u32>;
@group(0) @binding(3) var<storage, read_write> out : array<u32>;
var<workgroup> sv : array<f32, 256>;
var<workgroup> si : array<u32, 256>;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg_id : vec3<u32>,
@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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> ids : array<u32>;
@group(0) @binding(2) var<storage, read> embed : array<f16>;
@group(0) @binding(3) var<storage, read_write> x : array<f32>;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg_id : vec3<u32>,
@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<u32> = 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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> w : array<vec4<u32>>;
@group(0) @binding(2) var<storage, read> x : array<vec4<f32>>;
@group(0) @binding(3) var<storage, read_write> y : array<f32>;
const WG : u32 = ${WG}u;
const GRID_X : u32 = ${GRID_X}u;
var<workgroup> sums : array<f32, ${WG / 4}>;
@compute @workgroup_size(${WG})
fn main(@builtin(workgroup_id) wg_id : vec3<u32>,
@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<u32>
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<f32>(a, b), x0) + dot(vec4<f32>(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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> w : array<f16>;
@group(0) @binding(2) var<storage, read> xin : array<f32>;
@group(0) @binding(3) var<storage, read_write> y : array<f32>;
var<workgroup> at : array<f32, 1024>; // [16][64]
var<workgroup> wt : array<f16, 1024>; // [16][64]
@compute @workgroup_size(16, 16)
fn main(@builtin(workgroup_id) wg_id : vec3<u32>,
@builtin(local_invocation_id) l : vec3<u32>,
@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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> x : array<f32>;
@group(0) @binding(2) var<storage, read_write> y : array<f32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
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<f16, 8, 8>;\n`;
for (let i = 0; i < 2; i++)
loads += ` let a${i} = subgroupMatrixLoad<subgroup_matrix_left<f16, 8, 8>>(&tile_a, a_off + ${i * 8}u * BK, false, BK);\n`;
for (let j = 0; j < 4; j++)
loads += ` let b${j} = subgroupMatrixLoad<subgroup_matrix_right<f16, 8, 8>>(&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<uniform> params : Params;
@group(0) @binding(1) var<storage, read> wq : array<u32>;
@group(0) @binding(2) var<storage, read> scales : array<u32>;
@group(0) @binding(3) var<storage, read> zeros : array<u32>;
@group(0) @binding(4) var<storage, read> xin : array<f32>;
@group(0) @binding(5) var<storage, read_write> y : array<f32>;
const BM : u32 = ${BM}u; const BN : u32 = ${BN}u; const BK : u32 = ${BK}u;
var<workgroup> tile_a : array<f16, ${BM * BK}>;
var<workgroup> tile_b : array<f16, ${BN * BK}>;
var<workgroup> scratch : array<array<array<f16, 64>, 8>, 4>;
${qHelpers(bits)}
@compute @workgroup_size(128)
fn main(@builtin(workgroup_id) wg_id : vec3<u32>,
@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; }
}
}
}
`;
}