| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| export const GRID_X = 16384; |
|
|
| const SCALE_ATTN = '0.0883883461356163'; |
| const NEG_MAX = '-3.4028234663852886e+38'; |
|
|
| |
| |
| |
| |
| const qHelpers = (bits) => ` |
| 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);`} |
| } |
| `; |
|
|
| |
| |
| export function rmsnorm(D, withWeight) { |
| const WG = 256; |
| const PER = D / WG; |
| return ` |
| 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;'} |
| } |
| } |
| `; |
| } |
|
|
| |
| |
| |
| |
| export function gemvQ4(bits = 4) { |
| const WG = 128; |
| return ` |
| 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; } |
| } |
| } |
| } |
| `; |
| } |
|
|
| |
| |
| |
| |
| |
| 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; |
| return ` |
| 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]; } |
| } |
| } |
| } |
| } |
| `; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| export function qkvPost() { |
| const WG = 128; |
| return ` |
| 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); |
| } |
| } |
| } |
| } |
| `; |
| } |
|
|
| |
| |
| export function attnScores() { |
| return ` |
| 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}; |
| } |
| } |
| `; |
| } |
|
|
| |
| |
| |
| export function attnSoftmaxGate() { |
| const WG = 256; |
| return ` |
| 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; |
| } |
| } |
| `; |
| } |
|
|
| |
| |
| export function attnPV() { |
| return ` |
| 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; |
| } |
| } |
| `; |
| } |
|
|
| |
| |
| |
| |
| export function attnDecode() { |
| return ` |
| 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; |
| } |
| `; |
| } |
|
|
| |
| |
| |
| |
| |
| export function attnDecodeMq() { |
| return ` |
| 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; |
| } |
| `; |
| } |
|
|
| |
| |
| export function mlpAct() { |
| return ` |
| 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]; |
| } |
| `; |
| } |
|
|
| |
| |
| |
| |
| |
| export function argmaxStage1() { |
| return ` |
| 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 ` |
| 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]; } |
| } |
| `; |
| } |
|
|
| |
| |
| |
| |
| |
| export function gatherEmbed() { |
| return ` |
| 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]); |
| } |
| } |
| `; |
| } |
|
|
| |
| |
| |
| export function gemvF16W() { |
| const WG = 128; |
| return ` |
| 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; } |
| } |
| } |
| } |
| `; |
| } |
|
|
| |
| |
| |
| export function gemmF16W() { |
| return ` |
| 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; |
| } |
| } |
| `; |
| } |
|
|
| |
| |
| export function relu2() { |
| return ` |
| 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; |
| } |
| `; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| export function gemmQ4Sg(bits = 4) { |
| const BM = 32, BN = 64, BK = 32; |
| const VPW = bits === 4 ? 8 : 4; |
| |
| 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 ` |
| 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; } |
| } |
| } |
| } |
| `; |
| } |
|
|