| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| {{ENABLE_IMMEDIATE}} |
| {{ENABLE_SG}} |
| {{ENABLE_F16}} |
|
|
| struct Params { |
| B: u32, |
| t: u32, |
| S: u32, |
| _pad: u32, |
| } |
|
|
| {{PARAM_BINDING}}var<{{PARAM_ADDRESS}}> params: Params; |
| @group(0) @binding(1) var<storage, read> W: array<vec4<{{T}}>>; |
| @group(0) @binding(2) var<storage, read> ring: array<u32>; |
| @group(0) @binding(3) var<storage, read_write> Kc: array<vec4<{{T}}>>; |
| @group(0) @binding(4) var<storage, read_write> Vc: array<vec4<{{T}}>>; |
| @group(0) @binding(5) var<storage, read> CKV: array<vec4<{{T}}>>; |
| @group(0) @binding(6) var<storage, read> lens: array<u32>; |
| @group(0) @binding(7) var<storage, read_write> X: array<vec4<{{T}}>>; |
|
|
| const H: u32 = {{H}}u; |
| const D: u32 = {{D}}u; |
| const D4: u32 = D / 4u; |
| const HD4: u32 = H * D4; |
| const QKV4: u32 = 3u * HD4; |
| const FFN4: u32 = {{FFN4}}u; |
| const KQ_FFN: u32 = FFN4; |
| const LMAX: u32 = {{LMAX}}u; |
| const SCORES_MAX: u32 = {{SCORES_CAP}}u; |
| const ATTN_SCALE: f32 = {{ATTN_SCALE}}; |
| const WG: u32 = {{WG}}u; |
| const JT: u32 = WG / D4; |
| const TMP4: u32 = max(FFN4, HD4 + WG); |
|
|
| var<workgroup> xs4: array<vec4<f32>, HD4>; |
| var<workgroup> tmp4: array<vec4<f32>, TMP4>; |
| var<workgroup> out4: array<vec4<f32>, HD4>; |
| var<workgroup> scores: array<f32, SCORES_MAX>; |
| {{IF_NOSG}} |
| var<workgroup> red: array<f32, WG>; |
| {{/IF_NOSG}} |
| {{IF_SG}} |
| |
| |
| |
| |
| |
| |
| |
| const NSG_CAP: u32 = WG / 4u; |
| var<workgroup> red: array<f32, 2u * NSG_CAP>; |
| var<private> sgId: u32; |
| var<private> nSg: u32; |
| var<private> redSlot: u32 = 0u; |
| {{/IF_SG}} |
|
|
| fn wgMax(tid: u32, v: f32) -> f32 { |
| {{IF_SG}} |
| let s1 = subgroupMax(v); |
| let base = redSlot * NSG_CAP; |
| if (subgroupElect()) { red[base + sgId] = s1; } |
| workgroupBarrier(); |
| var r = red[base]; |
| for (var i = 1u; i < nSg; i = i + 1u) { r = max(r, red[base + i]); } |
| redSlot = 1u - redSlot; |
| return r; |
| {{/IF_SG}} |
| {{IF_NOSG}} |
| red[tid] = v; |
| workgroupBarrier(); |
| for (var s = WG / 2u; s > 0u; s = s >> 1u) { |
| if (tid < s) { red[tid] = max(red[tid], red[tid + s]); } |
| workgroupBarrier(); |
| } |
| let r = red[0]; |
| workgroupBarrier(); |
| return r; |
| {{/IF_NOSG}} |
| } |
|
|
| fn wgSum(tid: u32, v: f32) -> f32 { |
| {{IF_SG}} |
| let s1 = subgroupAdd(v); |
| let base = redSlot * NSG_CAP; |
| if (subgroupElect()) { red[base + sgId] = s1; } |
| workgroupBarrier(); |
| var r = red[base]; |
| for (var i = 1u; i < nSg; i = i + 1u) { r = r + red[base + i]; } |
| redSlot = 1u - redSlot; |
| return r; |
| {{/IF_SG}} |
| {{IF_NOSG}} |
| red[tid] = v; |
| workgroupBarrier(); |
| for (var s = WG / 2u; s > 0u; s = s >> 1u) { |
| if (tid < s) { red[tid] = red[tid] + red[tid + s]; } |
| workgroupBarrier(); |
| } |
| let r = red[0]; |
| workgroupBarrier(); |
| return r; |
| {{/IF_NOSG}} |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| fn gemvQuad(wOff: u32, bOff: u32, n4: u32, kq: u32, nq: u32, srcSel: u32) -> vec4<f32> { |
| var a0 = vec4<f32>(0.0); |
| var a1 = vec4<f32>(0.0); |
| var a2 = vec4<f32>(0.0); |
| var a3 = vec4<f32>(0.0); |
| for (var k4 = 0u; k4 < kq; k4 = k4 + 1u) { |
| var xq: vec4<f32>; |
| if (srcSel == 0u) { xq = xs4[k4]; } |
| else if (srcSel == 1u) { xq = out4[k4]; } |
| else { xq = tmp4[k4]; } |
| let kBase = wOff + (k4 << 2u) * nq + n4; |
| a0 = fma(vec4<f32>(xq.x), vec4<f32>(W[kBase]), a0); |
| a1 = fma(vec4<f32>(xq.y), vec4<f32>(W[kBase + nq]), a1); |
| a2 = fma(vec4<f32>(xq.z), vec4<f32>(W[kBase + 2u * nq]), a2); |
| a3 = fma(vec4<f32>(xq.w), vec4<f32>(W[kBase + 3u * nq]), a3); |
| } |
| return (a0 + a1) + (a2 + a3) + vec4<f32>(W[bOff + n4]); |
| } |
|
|
| @compute @workgroup_size({{WG}}) |
| fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>{{IF_SG}}, @builtin(subgroup_size) sgSize: u32{{/IF_SG}}) { |
| |
| if (wid.x >= params.B) { return; } |
| let b = wid.x; |
| let tid = lid.x; |
| let t = params.t; |
| {{IF_SG}} |
| sgId = tid / sgSize; |
| nSg = (WG + sgSize - 1u) / sgSize; |
| {{/IF_SG}} |
|
|
| |
| {{IF_EMBED}} |
| |
| |
| var id: u32 = {{DECODER_START}}u; |
| if (t != 0u) { id = ring[(t - 1u) * params.B + b]; } |
| for (var i = tid; i < HD4; i = i + WG) { |
| let e = vec4<f32>(W[{{TABLE4}}u + id * HD4 + i]) * {{EMBED_SCALE}} |
| + vec4<f32>(W[{{POS4}}u + t * HD4 + i]); |
| xs4[i] = vec4<f32>(vec4<{{T}}>(e)); |
| } |
| {{/IF_EMBED}} |
| {{IF_NOEMBED}} |
| |
| |
| |
| _ = ring[0]; |
| for (var i = tid; i < HD4; i = i + WG) { |
| xs4[i] = vec4<f32>(X[b * HD4 + i]); |
| } |
| {{/IF_NOEMBED}} |
| workgroupBarrier(); |
|
|
| |
| let kvBase = (b * LMAX + t) * HD4; |
| for (var n4 = tid; n4 < QKV4; n4 = n4 + WG) { |
| let g = vec4<{{T}}>(gemvQuad({{QKVW4}}u, {{QKVB4}}u, n4, HD4, QKV4, 0u)); |
| if (n4 < HD4) { tmp4[n4] = vec4<f32>(g); } |
| else if (n4 < 2u * HD4) { Kc[kvBase + n4 - HD4] = g; } |
| else { Vc[kvBase + n4 - 2u * HD4] = g; } |
| } |
| workgroupBarrier(); |
|
|
| |
| { |
| let len = min(t + 1u, LMAX); |
| for (var h = 0u; h < H; h = h + 1u) { |
| let hq = h * D4; |
| var lm: f32 = -1e30; |
| for (var j = tid; j < len; j = j + WG) { |
| let koff = (b * LMAX + j) * HD4 + hq; |
| var dot4 = vec4<f32>(0.0); |
| for (var i = 0u; i < D4; i = i + 1u) { |
| dot4 = dot4 + tmp4[hq + i] * vec4<f32>(Kc[koff + i]); |
| } |
| let sc = (dot4.x + dot4.y + dot4.z + dot4.w) * ATTN_SCALE; |
| scores[j] = sc; |
| lm = max(lm, sc); |
| } |
| let rowMax = wgMax(tid, lm); |
| var ls: f32 = 0.0; |
| for (var j = tid; j < len; j = j + WG) { |
| let e = exp(scores[j] - rowMax); |
| scores[j] = e; |
| ls = ls + e; |
| } |
| let denom = wgSum(tid, ls); |
| let dq = tid % D4; |
| let jg = tid / D4; |
| var acc = vec4<f32>(0.0); |
| if (jg < JT) { |
| for (var j = jg; j < len; j = j + JT) { |
| acc = acc + scores[j] * vec4<f32>(Vc[(b * LMAX + j) * HD4 + hq + dq]); |
| } |
| } |
| tmp4[HD4 + tid] = acc; |
| workgroupBarrier(); |
| if (tid < D4) { |
| var o = vec4<f32>(0.0); |
| for (var g = 0u; g < JT; g = g + 1u) { o = o + tmp4[HD4 + g * D4 + tid]; } |
| out4[hq + tid] = vec4<f32>(vec4<{{T}}>(o / denom)); |
| } |
| workgroupBarrier(); |
| } |
| } |
|
|
| |
| for (var n4 = tid; n4 < HD4; n4 = n4 + WG) { |
| let g = vec4<{{T}}>(gemvQuad({{OUTW4}}u, {{OUTB4}}u, n4, HD4, HD4, 1u)); |
| tmp4[n4] = vec4<f32>(g) + xs4[n4]; |
| } |
| workgroupBarrier(); |
| { |
| var s: f32 = 0.0; |
| for (var i = tid; i < HD4; i = i + WG) { |
| let v = tmp4[i]; |
| s = s + v.x + v.y + v.z + v.w; |
| } |
| let mu = wgSum(tid, s) / f32(H * D); |
| var sq: f32 = 0.0; |
| for (var i = tid; i < HD4; i = i + WG) { |
| let dv = tmp4[i] - vec4<f32>(mu); |
| sq = sq + dot(dv, dv); |
| } |
| let inv = inverseSqrt(wgSum(tid, sq) / f32(H * D) + {{EPS}}); |
| for (var i = tid; i < HD4; i = i + WG) { |
| let o = vec4<f32>(W[{{LN1G4}}u + i]) * (tmp4[i] - vec4<f32>(mu)) * inv |
| + vec4<f32>(W[{{LN1B4}}u + i]); |
| xs4[i] = vec4<f32>(vec4<{{T}}>(o)); |
| } |
| } |
| workgroupBarrier(); |
|
|
| |
| for (var n4 = tid; n4 < HD4; n4 = n4 + WG) { |
| out4[n4] = vec4<f32>(vec4<{{T}}>(gemvQuad({{CQW4}}u, {{CQB4}}u, n4, HD4, HD4, 0u))); |
| } |
| workgroupBarrier(); |
|
|
| |
| { |
| let len = min(lens[b], SCORES_MAX); |
| for (var h = 0u; h < H; h = h + 1u) { |
| let hq = h * D4; |
| var lm: f32 = -1e30; |
| for (var j = tid; j < len; j = j + WG) { |
| let koff = (b * params.S + j) * 2u * HD4 + hq; |
| var dot4 = vec4<f32>(0.0); |
| for (var i = 0u; i < D4; i = i + 1u) { |
| dot4 = dot4 + out4[hq + i] * vec4<f32>(CKV[koff + i]); |
| } |
| let sc = (dot4.x + dot4.y + dot4.z + dot4.w) * ATTN_SCALE; |
| scores[j] = sc; |
| lm = max(lm, sc); |
| } |
| let rowMax = wgMax(tid, lm); |
| var ls: f32 = 0.0; |
| for (var j = tid; j < len; j = j + WG) { |
| let e = exp(scores[j] - rowMax); |
| scores[j] = e; |
| ls = ls + e; |
| } |
| let denom = wgSum(tid, ls); |
| let dq = tid % D4; |
| let jg = tid / D4; |
| var acc = vec4<f32>(0.0); |
| if (jg < JT) { |
| for (var j = jg; j < len; j = j + JT) { |
| |
| acc = acc + scores[j] * vec4<f32>(CKV[(b * params.S + j) * 2u * HD4 + HD4 + hq + dq]); |
| } |
| } |
| tmp4[HD4 + tid] = acc; |
| workgroupBarrier(); |
| if (tid < D4) { |
| var o = vec4<f32>(0.0); |
| for (var g = 0u; g < JT; g = g + 1u) { o = o + tmp4[HD4 + g * D4 + tid]; } |
| tmp4[hq + tid] = vec4<f32>(vec4<{{T}}>(o / denom)); |
| } |
| workgroupBarrier(); |
| } |
| } |
|
|
| |
| for (var n4 = tid; n4 < HD4; n4 = n4 + WG) { |
| let g = vec4<{{T}}>(gemvQuad({{COW4}}u, {{COB4}}u, n4, HD4, HD4, 2u)); |
| out4[n4] = vec4<f32>(g) + xs4[n4]; |
| } |
| workgroupBarrier(); |
| { |
| var s: f32 = 0.0; |
| for (var i = tid; i < HD4; i = i + WG) { |
| let v = out4[i]; |
| s = s + v.x + v.y + v.z + v.w; |
| } |
| let mu = wgSum(tid, s) / f32(H * D); |
| var sq: f32 = 0.0; |
| for (var i = tid; i < HD4; i = i + WG) { |
| let dv = out4[i] - vec4<f32>(mu); |
| sq = sq + dot(dv, dv); |
| } |
| let inv = inverseSqrt(wgSum(tid, sq) / f32(H * D) + {{EPS}}); |
| for (var i = tid; i < HD4; i = i + WG) { |
| let o = vec4<f32>(W[{{LN2G4}}u + i]) * (out4[i] - vec4<f32>(mu)) * inv |
| + vec4<f32>(W[{{LN2B4}}u + i]); |
| xs4[i] = vec4<f32>(vec4<{{T}}>(o)); |
| } |
| } |
| workgroupBarrier(); |
|
|
| |
| for (var n4 = tid; n4 < FFN4; n4 = n4 + WG) { |
| var v = gemvQuad({{FC1W4}}u, {{FC1B4}}u, n4, HD4, FFN4, 0u); |
| v = v / (vec4<f32>(1.0) + exp(-v)); |
| tmp4[n4] = vec4<f32>(vec4<{{T}}>(v)); |
| } |
| workgroupBarrier(); |
|
|
| |
| for (var n4 = tid; n4 < HD4; n4 = n4 + WG) { |
| let g = vec4<{{T}}>(gemvQuad({{FC2W4}}u, {{FC2B4}}u, n4, KQ_FFN, HD4, 2u)); |
| out4[n4] = vec4<f32>(g) + xs4[n4]; |
| } |
| workgroupBarrier(); |
| { |
| var s: f32 = 0.0; |
| for (var i = tid; i < HD4; i = i + WG) { |
| let v = out4[i]; |
| s = s + v.x + v.y + v.z + v.w; |
| } |
| let mu = wgSum(tid, s) / f32(H * D); |
| var sq: f32 = 0.0; |
| for (var i = tid; i < HD4; i = i + WG) { |
| let dv = out4[i] - vec4<f32>(mu); |
| sq = sq + dot(dv, dv); |
| } |
| let inv = inverseSqrt(wgSum(tid, sq) / f32(H * D) + {{EPS}}); |
| for (var i = tid; i < HD4; i = i + WG) { |
| let o = vec4<f32>(W[{{LN3G4}}u + i]) * (out4[i] - vec4<f32>(mu)) * inv |
| + vec4<f32>(W[{{LN3B4}}u + i]); |
| X[b * HD4 + i] = vec4<{{T}}>(o); |
| } |
| } |
| } |
|
|