use crate::deepencoder::{get_rel_pos, DeepEncoderConfig, SamBlockWeights};
use crate::forward::{make_bg, pipeline, uni};
use crate::GpuCtx;
const LINEAR: &str = r#"
@group(0) @binding(0) var<storage, read> x: array<f32>;
@group(0) @binding(1) var<storage, read> w: array<f32>;
@group(0) @binding(2) var<storage, read> b: array<f32>;
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
@group(0) @binding(4) var<uniform> d: vec4<u32>; // (m, k, n, has_bias)
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {
let m = d.x; let k = d.y; let n = d.z;
let idx = gid.x + gid.y * nwg.x * 64u; if (idx >= m * n) { return; }
let i = idx / n; let j = idx % n;
var acc = 0.0; if (d.w == 1u) { acc = b[j]; }
for (var p = 0u; p < k; p = p + 1u) { acc = acc + x[i * k + p] * w[j * k + p]; }
y[idx] = acc;
}
"#;
pub(crate) const LAYERNORM: &str = r#"
@group(0) @binding(0) var<storage, read> x: array<f32>;
@group(0) @binding(1) var<storage, read> g: array<f32>;
@group(0) @binding(2) var<storage, read> b: array<f32>;
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
@group(0) @binding(4) var<uniform> d: vec4<u32>; // (rows, c, eps_bits, _)
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {
let rows = d.x; let c = d.y; let eps = bitcast<f32>(d.z);
let r = gid.x + gid.y * nwg.x * 64u; if (r >= rows) { return; }
var mean = 0.0;
for (var j = 0u; j < c; j = j + 1u) { mean = mean + x[r * c + j]; }
mean = mean / f32(c);
var vsum = 0.0;
for (var j = 0u; j < c; j = j + 1u) { let t = x[r * c + j] - mean; vsum = vsum + t * t; }
let inv = 1.0 / sqrt(vsum / f32(c) + eps);
for (var j = 0u; j < c; j = j + 1u) { y[r * c + j] = (x[r * c + j] - mean) * inv * g[j] + b[j]; }
}
"#;
#[allow(dead_code)]
const SAM_SCORES: &str = r#"
@group(0) @binding(0) var<storage, read> qkv: array<f32>; // [n, 3*width]
@group(0) @binding(1) var<storage, read> rh: array<f32>; // [gh, gh, hd]
@group(0) @binding(2) var<storage, read> rw: array<f32>; // [gw, gw, hd]
@group(0) @binding(3) var<storage, read_write> attn: array<f32>; // [heads, n, n]
@group(0) @binding(4) var<uniform> d: vec4<u32>; // (n, heads, hd, gw)
@group(0) @binding(5) var<uniform> e: vec4<u32>; // (width, gh, scale_bits, _)
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {
let n = d.x; let heads = d.y; let hd = d.z; let gw = d.w;
let width = e.x; let gh = e.y; let scale = bitcast<f32>(e.z);
let idx = gid.x + gid.y * nwg.x * 64u; if (idx >= heads * n * n) { return; }
let h = idx / (n * n); let rem = idx % (n * n); let i = rem / n; let j = rem % n;
let qbase = i * 3u * width + h * hd;
let kbase = j * 3u * width + width + h * hd;
var qk = 0.0;
for (var c = 0u; c < hd; c = c + 1u) { qk = qk + qkv[qbase + c] * qkv[kbase + c]; }
let qy = i / gw; let qx = i % gw; let ky = j / gw; let kx = j % gw;
var dh = 0.0;
for (var c = 0u; c < hd; c = c + 1u) { dh = dh + qkv[qbase + c] * rh[(qy * gh + ky) * hd + c]; }
var dw = 0.0;
for (var c = 0u; c < hd; c = c + 1u) { dw = dw + qkv[qbase + c] * rw[(qx * gw + kx) * hd + c]; }
attn[idx] = qk * scale + dh + dw;
}
"#;
const DHDW: &str = r#"
@group(0) @binding(0) var<storage, read> qkv: array<f32>;
@group(0) @binding(1) var<storage, read> rh: array<f32>;
@group(0) @binding(2) var<storage, read> rw: array<f32>;
@group(0) @binding(3) var<storage, read_write> dh: array<f32>; // [heads, n, gh]
@group(0) @binding(4) var<storage, read_write> dw: array<f32>; // [heads, n, gw]
@group(0) @binding(5) var<uniform> d: vec4<u32>; // (n, heads, hd, gw)
@group(0) @binding(6) var<uniform> e: vec4<u32>; // (width, gh, _, _)
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {
let n = d.x; let heads = d.y; let hd = d.z; let gw = d.w; let width = e.x; let gh = e.y;
let idx = gid.x + gid.y * nwg.x * 64u; if (idx >= heads * n) { return; }
let h = idx / n; let i = idx % n; let qy = i / gw; let qx = i % gw;
let qbase = i * 3u * width + h * hd;
let dhb = (h * n + i) * gh;
for (var ky = 0u; ky < gh; ky = ky + 1u) {
var s = 0.0; let rb = (qy * gh + ky) * hd;
for (var c = 0u; c < hd; c = c + 1u) { s = s + qkv[qbase + c] * rh[rb + c]; }
dh[dhb + ky] = s;
}
let dwb = (h * n + i) * gw;
for (var kx = 0u; kx < gw; kx = kx + 1u) {
var s = 0.0; let rb = (qx * gw + kx) * hd;
for (var c = 0u; c < hd; c = c + 1u) { s = s + qkv[qbase + c] * rw[rb + c]; }
dw[dwb + kx] = s;
}
}
"#;
pub(crate) fn flash_reg_src(hd: usize, rb: usize, bias: bool) -> String {
assert_eq!(hd, 64, "flash_reg bakes hd = workgroup width = 64");
let per = |f: &dyn Fn(usize) -> String| (0..rb).map(|i| f(i)).collect::<Vec<_>>().join("\n");
let m_decl = per(&|i| format!(" var m{i} = -3.0e38; var l{i} = 0.0; var a{i} = 0.0;"));
let d_decl = per(&|i| format!(" var dp{i} = 0.0;"));
let dots = per(&|i| format!(" dp{i} = dp{i} + qsh[{i}u * 64u + c] * kv;"));
let score = if bias {
per(&|i| format!(
" var s{i} = -3.0e38;\n if (j < n && rb0 + {i}u < n) {{ s{i} = dp{i} * scale + dh[(h * n + rb0 + {i}u) * gh + ky] + dw[(h * n + rb0 + {i}u) * gw + kx]; }}\n psh[{i}u * 64u + t] = s{i};"))
} else {
per(&|i| format!(
" var s{i} = -3.0e38;\n if (j < n && rb0 + {i}u < n) {{ s{i} = dp{i} * scale; }}\n psh[{i}u * 64u + t] = s{i};"))
};
let redmax_seed = per(&|i| format!(" red[{i}u * 64u + t] = s{i};"));
let redmax_step = per(&|i| format!(" if (t < st) {{ red[{i}u * 64u + t] = max(red[{i}u * 64u + t], red[{i}u * 64u + t + st]); }}"));
let newm = per(&|i| format!(" let nm{i} = max(m{i}, red[{i}u * 64u]);"));
let prob = per(&|i| format!(" var p{i} = 0.0; if (psh[{i}u * 64u + t] > -3.0e37) {{ p{i} = exp(psh[{i}u * 64u + t] - nm{i}); }} psh[{i}u * 64u + t] = p{i};"));
let redsum_seed = per(&|i| format!(" red[{i}u * 64u + t] = p{i};"));
let redsum_step = per(&|i| format!(" if (t < st) {{ red[{i}u * 64u + t] = red[{i}u * 64u + t] + red[{i}u * 64u + t + st]; }}"));
let lupd = per(&|i| format!(" let rs{i} = exp(m{i} - nm{i}); l{i} = l{i} * rs{i} + red[{i}u * 64u]; a{i} = a{i} * rs{i}; m{i} = nm{i};"));
let vacc = per(&|i| format!(" a{i} = a{i} + psh[{i}u * 64u + jj] * vv;"));
let store = per(&|i| format!(" if (rb0 + {i}u < n && l{i} > 0.0) {{ out[(rb0 + {i}u) * width + h * 64u + t] = a{i} / l{i}; }}"));
let sh = rb * 64;
let qsh = rb * 64;
let bias_bindings = if bias {
"@group(0) @binding(1) var<storage, read> dh: array<f32>; // [heads, n, gh]\n@group(0) @binding(2) var<storage, read> dw: array<f32>; // [heads, n, gw]"
} else {
""
};
let kykx = if bias { " let ky = j / gw; let kx = j % gw;" } else { "" };
format!(
r#"
@group(0) @binding(0) var<storage, read> qkv: array<f32>; // [n, 3*width]
{bias_bindings}
@group(0) @binding(3) var<storage, read_write> out: array<f32>; // [n, width]
@group(0) @binding(4) var<uniform> d: vec4<u32>; // (n, heads, hd, gw)
@group(0) @binding(5) var<uniform> e: vec4<u32>; // (width, gh, scale_bits, _)
const RB: u32 = {rb}u;
var<workgroup> qsh: array<f32, {qsh}>; // RB query rows × 64
var<workgroup> psh: array<f32, {sh}>; // scores/probs: RB × 64
var<workgroup> red: array<f32, {sh}>; // reduction scratch
@compute @workgroup_size(64)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>,
@builtin(num_workgroups) nwg: vec3<u32>) {{
let n = d.x; let heads = d.y; let gw = d.w;
let width = e.x; let gh = e.y; let scale = bitcast<f32>(e.z);
let t = lid.x;
let nblocks = (n + RB - 1u) / RB;
let wgid = wid.x + wid.y * nwg.x;
let h = wgid / nblocks; if (h >= heads) {{ return; }}
let rb0 = (wgid % nblocks) * RB;
// load qsh: RB rows × 64 dims of head h
for (var el = t; el < RB * 64u; el = el + 64u) {{
let ri = el / 64u; let c = el % 64u; let row = rb0 + ri;
qsh[el] = select(0.0, qkv[row * 3u * width + h * 64u + c], row < n);
}}
workgroupBarrier();
{m_decl}
let nchunks = (n + 63u) / 64u;
for (var ch = 0u; ch < nchunks; ch = ch + 1u) {{
let j = ch * 64u + t; // this thread's key (scoring)
{kykx}
{d_decl}
if (j < n) {{
for (var c = 0u; c < 64u; c = c + 1u) {{
let kv = qkv[j * 3u * width + width + h * 64u + c];
{dots}
}}
}}
{score}
workgroupBarrier();
// max-reduce over the 64 keys, per row
{redmax_seed}
workgroupBarrier();
for (var st = 32u; st > 0u; st = st >> 1u) {{
{redmax_step}
workgroupBarrier();
}}
{newm}
{prob}
workgroupBarrier();
// sum-reduce probs over the 64 keys, per row
{redsum_seed}
workgroupBarrier();
for (var st = 32u; st > 0u; st = st >> 1u) {{
{redsum_step}
workgroupBarrier();
}}
{lupd}
workgroupBarrier();
// value pass: thread t owns output dim t; accumulate over the 64 keys of this chunk
for (var jj = 0u; jj < 64u; jj = jj + 1u) {{
let kj = ch * 64u + jj;
if (kj < n) {{
let vv = qkv[kj * 3u * width + 2u * width + h * 64u + t];
{vacc}
}}
}}
workgroupBarrier();
}}
{store}
}}
"#
)
}
pub(crate) fn flash_reg_kv16_src(hd: usize, rb: usize, bias: bool) -> String {
let src = flash_reg_src(hd, rb, bias);
let a1 = "@group(0) @binding(0) var<storage, read> qkv: array<f32>;";
let a2 = "qsh[el] = select(0.0, qkv[row * 3u * width + h * 64u + c], row < n);";
let a3 = " let kv = qkv[j * 3u * width + width + h * 64u + c];";
let a4 = " let vv = qkv[kj * 3u * width + 2u * width + h * 64u + t];";
for a in [a1, a2, a3, a4] {
assert!(src.contains(a), "kv16 anchor drifted: {a:?}");
}
format!("enable f16;\n{src}")
.replace(a1, "@group(0) @binding(0) var<storage, read> qkv: array<f16>;")
.replace(a2, "qsh[el] = select(0.0, f32(qkv[row * 3u * width + h * 64u + c]), row < n);")
.replace(a3, " let kv = f32(qkv[j * 3u * width + width + h * 64u + c]);")
.replace(a4, " let vv = f32(qkv[kj * 3u * width + 2u * width + h * 64u + t]);")
}
pub(crate) fn flash_reg_f16a_src(hd: usize, rb: usize, bias: bool) -> String {
let src = flash_reg_src(hd, rb, bias);
let a1 = "var<workgroup> qsh: array<f32,";
let a2 = " qsh[el] = select(0.0, qkv[row * 3u * width + h * 64u + c], row < n);";
let a3 = " let kv = qkv[j * 3u * width + width + h * 64u + c];";
for a in [a1, a2, a3] {
assert!(src.contains(a), "flash f16a anchor drifted: {a:?}");
}
let mut out = format!("enable f16;\n{src}")
.replace(a1, "var<workgroup> qsh: array<f16,")
.replace(a2, " qsh[el] = f16(select(0.0, qkv[row * 3u * width + h * 64u + c], row < n));")
.replace(a3, " let kv = f16(qkv[j * 3u * width + width + h * 64u + c]);");
for i in 0..rb {
let d0 = format!(" var dp{i} = 0.0;");
let d1 = format!(" var dp{i} = f16(0.0);");
assert!(out.contains(&d0));
out = out.replace(&d0, &d1);
let s0 = format!("s{i} = dp{i} * scale;");
let s1 = format!("s{i} = f32(dp{i}) * scale;");
assert!(out.contains(&s0), "score anchor {i}");
out = out.replace(&s0, &s1);
if bias {
let b0 = format!("s{i} = dp{i} * scale +");
if out.contains(&b0) {
out = out.replace(&b0, &format!("s{i} = f32(dp{i}) * scale +"));
}
}
}
out
}
pub(crate) fn flash_reg_sg_src(hd: usize, rb: usize, bias: bool) -> String {
assert_eq!(hd, 64, "flash_reg bakes hd = workgroup width = 64");
let per = |f: &dyn Fn(usize) -> String| (0..rb).map(|i| f(i)).collect::<Vec<_>>().join("\n");
let m_decl = per(&|i| format!(" var m{i} = -3.0e38; var l{i} = 0.0; var a{i} = 0.0;"));
let d_decl = per(&|i| format!(" var dp{i} = 0.0;"));
let dots = per(&|i| format!(" dp{i} = dp{i} + qsh[{i}u * 64u + c] * kv;"));
let score = if bias {
per(&|i| format!(
" var s{i} = -3.0e38;\n if (j < n && rb0 + {i}u < n) {{ s{i} = dp{i} * scale + dh[(h * n + rb0 + {i}u) * gh + ky] + dw[(h * n + rb0 + {i}u) * gw + kx]; }}\n psh[{i}u * 64u + t] = s{i};"))
} else {
per(&|i| format!(
" var s{i} = -3.0e38;\n if (j < n && rb0 + {i}u < n) {{ s{i} = dp{i} * scale; }}\n psh[{i}u * 64u + t] = s{i};"))
};
let sgmax = per(&|i| format!(
" let sgm{i} = subgroupMax(s{i});\n if (sid == 0u) {{ red[{i}u * 8u + sg] = sgm{i}; }}"));
let newm = per(&|i| format!(
" var wm{i} = -3.0e38;\n for (var g = 0u; g < nsg; g++) {{ wm{i} = max(wm{i}, red[{i}u * 8u + g]); }}\n let nm{i} = max(m{i}, wm{i});"));
let prob = per(&|i| format!(" var p{i} = 0.0; if (psh[{i}u * 64u + t] > -3.0e37) {{ p{i} = exp(psh[{i}u * 64u + t] - nm{i}); }} psh[{i}u * 64u + t] = p{i};"));
let sgsum = per(&|i| format!(
" let sgs{i} = subgroupAdd(p{i});\n if (sid == 0u) {{ red[{i}u * 8u + sg] = sgs{i}; }}"));
let lupd = per(&|i| format!(
" var ws{i} = 0.0;\n for (var g = 0u; g < nsg; g++) {{ ws{i} = ws{i} + red[{i}u * 8u + g]; }}\n let rs{i} = exp(m{i} - nm{i}); l{i} = l{i} * rs{i} + ws{i}; a{i} = a{i} * rs{i}; m{i} = nm{i};"));
let vacc = per(&|i| format!(" a{i} = a{i} + psh[{i}u * 64u + jj] * vv;"));
let store = per(&|i| format!(" if (rb0 + {i}u < n && l{i} > 0.0) {{ out[(rb0 + {i}u) * width + h * 64u + t] = a{i} / l{i}; }}"));
let sh = rb * 64;
let qsh = rb * 64;
let red = rb * 8;
let bias_bindings = if bias {
"@group(0) @binding(1) var<storage, read> dh: array<f32>; // [heads, n, gh]\n@group(0) @binding(2) var<storage, read> dw: array<f32>; // [heads, n, gw]"
} else {
""
};
let kykx = if bias { " let ky = j / gw; let kx = j % gw;" } else { "" };
format!(
r#"
@group(0) @binding(0) var<storage, read> qkv: array<f32>; // [n, 3*width]
{bias_bindings}
@group(0) @binding(3) var<storage, read_write> out: array<f32>; // [n, width]
@group(0) @binding(4) var<uniform> d: vec4<u32>; // (n, heads, hd, gw)
@group(0) @binding(5) var<uniform> e: vec4<u32>; // (width, gh, scale_bits, _)
const RB: u32 = {rb}u;
var<workgroup> qsh: array<f32, {qsh}>; // RB query rows × 64
var<workgroup> psh: array<f32, {sh}>; // scores/probs: RB × 64
var<workgroup> red: array<f32, {red}>; // cross-subgroup partials: RB × ≤8 subgroups
@compute @workgroup_size(64)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>,
@builtin(num_workgroups) nwg: vec3<u32>,
@builtin(subgroup_invocation_id) sid: u32, @builtin(subgroup_size) sgsz: u32) {{
let n = d.x; let heads = d.y; let gw = d.w;
let width = e.x; let gh = e.y; let scale = bitcast<f32>(e.z);
let t = lid.x;
let sg = t / sgsz;
let nsg = min((64u + sgsz - 1u) / sgsz, 8u);
let nblocks = (n + RB - 1u) / RB;
let wgid = wid.x + wid.y * nwg.x;
let h = wgid / nblocks; if (h >= heads) {{ return; }}
let rb0 = (wgid % nblocks) * RB;
for (var el = t; el < RB * 64u; el = el + 64u) {{
let ri = el / 64u; let c = el % 64u; let row = rb0 + ri;
qsh[el] = select(0.0, qkv[row * 3u * width + h * 64u + c], row < n);
}}
workgroupBarrier();
{m_decl}
let nchunks = (n + 63u) / 64u;
for (var ch = 0u; ch < nchunks; ch = ch + 1u) {{
let j = ch * 64u + t; // this thread's key (scoring)
{kykx}
{d_decl}
if (j < n) {{
for (var c = 0u; c < 64u; c = c + 1u) {{
let kv = qkv[j * 3u * width + width + h * 64u + c];
{dots}
}}
}}
{score}
{sgmax}
workgroupBarrier();
{newm}
{prob}
{sgsum}
workgroupBarrier();
{lupd}
// value pass: thread t owns output dim t; accumulate over the 64 keys of this chunk
for (var jj = 0u; jj < 64u; jj = jj + 1u) {{
let kj = ch * 64u + jj;
if (kj < n) {{
let vv = qkv[kj * 3u * width + 2u * width + h * 64u + t];
{vacc}
}}
}}
workgroupBarrier();
}}
{store}
}}
"#
)
}
pub(crate) const ADDACT: &str = r#"
@group(0) @binding(0) var<storage, read> a: array<f32>;
@group(0) @binding(1) var<storage, read> b: array<f32>;
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
@group(0) @binding(3) var<uniform> d: vec4<u32>; // (len, act, add_b, _)
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {
let len = d.x; let i = gid.x + gid.y * nwg.x * 64u; if (i >= len) { return; }
var v = a[i]; if (d.z == 1u) { v = v + b[i]; }
if (d.y == 1u) {
// exact gelu via erf approx (Abramowitz-Stegun) — matches the CPU reference's gelu
let x = v; let s = sign(x); let ax = abs(x) * 0.7071067811865476;
let t = 1.0 / (1.0 + 0.3275911 * ax);
let er = 1.0 - (((((1.061405429*t - 1.453152027)*t) + 1.421413741)*t - 0.284496736)*t + 0.254829592)*t*exp(-ax*ax);
v = 0.5 * x * (1.0 + s * er);
} else if (d.y == 2u) {
v = v / (1.0 + exp(-1.702 * v)); // quick_gelu
}
y[i] = v;
}
"#;
#[allow(dead_code)]
const SOFTMAX: &str = r#"
@group(0) @binding(0) var<storage, read_write> a: array<f32>;
@group(0) @binding(1) var<uniform> d: vec4<u32>; // (rows, n, _, _)
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {
let rows = d.x; let n = d.y; let r = gid.x + gid.y * nwg.x * 64u; if (r >= rows) { return; }
let base = r * n;
var mx = a[base];
for (var j = 1u; j < n; j = j + 1u) { mx = max(mx, a[base + j]); }
var s = 0.0;
for (var j = 0u; j < n; j = j + 1u) { let e = exp(a[base + j] - mx); a[base + j] = e; s = s + e; }
for (var j = 0u; j < n; j = j + 1u) { a[base + j] = a[base + j] / s; }
}
"#;
pub(crate) const TILED_LINEAR: &str = r#"
@group(0) @binding(0) var<storage, read> x: array<f32>;
@group(0) @binding(1) var<storage, read> w: array<f32>;
@group(0) @binding(2) var<storage, read> b: array<f32>;
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
@group(0) @binding(4) var<uniform> d: vec4<u32>; // (m, k, n, has_bias)
var<workgroup> xs: array<f32, 256>; // 16×16 x-tile
var<workgroup> ws: array<f32, 256>; // 16×16 w-tile
@compute @workgroup_size(16, 16)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
let m = d.x; let k = d.y; let n = d.z;
let row = wid.x * 16u + lid.x; // i in [0,m)
let col = wid.y * 16u + lid.y; // j in [0,n)
var acc = 0.0;
let ntile = (k + 15u) / 16u;
for (var t = 0u; t < ntile; t = t + 1u) {
let kx = t * 16u + lid.y;
xs[lid.x * 16u + lid.y] = select(0.0, x[row * k + kx], row < m && kx < k);
let kw = t * 16u + lid.x;
ws[lid.y * 16u + lid.x] = select(0.0, w[col * k + kw], col < n && kw < k);
workgroupBarrier();
for (var p = 0u; p < 16u; p = p + 1u) { acc = acc + xs[lid.x * 16u + p] * ws[lid.y * 16u + p]; }
workgroupBarrier();
}
if (row < m && col < n) {
if (d.w == 1u) { acc = acc + b[col]; }
y[row * n + col] = acc;
}
}
"#;
#[allow(dead_code)]
const AV_MERGE: &str = r#"
@group(0) @binding(0) var<storage, read> attn: array<f32>; // [heads, n, n]
@group(0) @binding(1) var<storage, read> qkv: array<f32>; // [n, 3*width]
@group(0) @binding(2) var<storage, read_write> out: array<f32>; // [n, width]
@group(0) @binding(3) var<uniform> d: vec4<u32>; // (n, width, hd, _)
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {
let n = d.x; let width = d.y; let hd = d.z;
let idx = gid.x + gid.y * nwg.x * 64u; if (idx >= n * width) { return; }
let i = idx / width; let c = idx % width; let h = c / hd;
var acc = 0.0;
for (var j = 0u; j < n; j = j + 1u) {
acc = acc + attn[(h * n + i) * n + j] * qkv[j * 3u * width + 2u * width + c];
}
out[idx] = acc;
}
"#;
pub(crate) fn run(ctx: &GpuCtx, pl: &wgpu::ComputePipeline, bg: &wgpu::BindGroup, threads: usize) {
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
{
let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
p.set_pipeline(pl);
p.set_bind_group(0, bg, &[]);
let wg = (threads + 63) / 64;
let gx = wg.min(65535) as u32;
let gy = ((wg + 65534) / 65535) as u32; p.dispatch_workgroups(gx, gy, 1);
}
ctx.queue.submit(Some(enc.finish()));
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
}
pub(crate) fn u32x4(a: u32, b: u32, c: u32, d: u32) -> Vec<u8> {
bytemuck::cast_slice(&[a, b, c, d]).to_vec()
}
fn gpu_linear(
ctx: &GpuCtx,
x: &wgpu::Buffer,
w: &wgpu::Buffer,
b: Option<&wgpu::Buffer>,
m: usize,
k: usize,
n: usize,
) -> wgpu::Buffer {
let pl = pipeline(ctx, "de_tiled_linear", TILED_LINEAR);
let y = ctx.empty(m * n);
let zero = ctx.storage(&[0.0]);
let bias = b.unwrap_or(&zero);
let meta = uni(ctx, &u32x4(m as u32, k as u32, n as u32, if b.is_some() { 1 } else { 0 }));
let bg = make_bg(ctx, &pl, &[x, w, bias, &y], &meta);
let (gx, gy) = ((m as u32).div_ceil(16), (n as u32).div_ceil(16));
let mut enc = ctx.device.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
{
let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
p.set_pipeline(&pl);
p.set_bind_group(0, &bg, &[]);
p.dispatch_workgroups(gx, gy, 1);
}
ctx.queue.submit(Some(enc.finish()));
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
y
}
pub fn sam_attention_gpu(
ctx: &GpuCtx, x: &[f32], gh: usize, gw: usize, width: usize, heads: usize, w: &SamBlockWeights,
) -> Vec<f32> {
let n = gh * gw;
let hd = width / heads;
let xb = ctx.storage(x);
let qkv_w = ctx.storage(&w.qkv_w);
let qkv_b = ctx.storage(&w.qkv_b);
let proj_w = ctx.storage(&w.proj_w);
let proj_b = ctx.storage(&w.proj_b);
let rh = ctx.storage(&get_rel_pos(gh, gh, &w.rel_pos_h, hd));
let rw = ctx.storage(&get_rel_pos(gw, gw, &w.rel_pos_w, hd));
let out = sam_attn_buf(ctx, &xb, gh, gw, width, heads, &qkv_w, &qkv_b, &proj_w, &proj_b, &rh, &rw);
ctx.read(&out, n * width).expect("read sam attn out")
}
fn ln_buf(ctx: &GpuCtx, x: &wgpu::Buffer, rows: usize, c: usize, g: &wgpu::Buffer, b: &wgpu::Buffer, eps: f32) -> wgpu::Buffer {
let pl = pipeline(ctx, "de_ln", LAYERNORM);
let y = ctx.empty(rows * c);
let meta = uni(ctx, &u32x4(rows as u32, c as u32, eps.to_bits(), 0));
let bg = make_bg(ctx, &pl, &[x, g, b, &y], &meta);
run(ctx, &pl, &bg, rows);
y
}
fn addact_buf(ctx: &GpuCtx, a: &wgpu::Buffer, b: Option<&wgpu::Buffer>, len: usize, act: u32) -> wgpu::Buffer {
let pl = pipeline(ctx, "de_addact", ADDACT);
let y = ctx.empty(len);
let zero = ctx.storage(&[0.0]);
let bb = b.unwrap_or(&zero);
let meta = uni(ctx, &u32x4(len as u32, act, if b.is_some() { 1 } else { 0 }, 0));
let bg = make_bg(ctx, &pl, &[a, bb, &y], &meta);
run(ctx, &pl, &bg, len);
y
}
#[allow(clippy::too_many_arguments)]
fn sam_attn_buf(ctx: &GpuCtx, xb: &wgpu::Buffer, gh: usize, gw: usize, width: usize, heads: usize,
qkv_w: &wgpu::Buffer, qkv_b: &wgpu::Buffer, proj_w: &wgpu::Buffer, proj_b: &wgpu::Buffer,
rh: &wgpu::Buffer, rw: &wgpu::Buffer) -> wgpu::Buffer {
let n = gh * gw;
let hd = width / heads;
let scale = 1.0f32 / (hd as f32).sqrt();
let qkv = gpu_linear(ctx, xb, qkv_w, Some(qkv_b), n, width, 3 * width);
let d = uni(ctx, &u32x4(n as u32, heads as u32, hd as u32, gw as u32));
let e = uni(ctx, &u32x4(width as u32, gh as u32, scale.to_bits(), 0));
let e0 = uni(ctx, &u32x4(width as u32, gh as u32, 0, 0));
let dhb = ctx.empty(heads * n * gh);
let dwb = ctx.empty(heads * n * gw);
let dp = pipeline(ctx, "de_dhdw", DHDW);
let dbg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None, layout: &dp.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry { binding: 0, resource: qkv.as_entire_binding() },
wgpu::BindGroupEntry { binding: 1, resource: rh.as_entire_binding() },
wgpu::BindGroupEntry { binding: 2, resource: rw.as_entire_binding() },
wgpu::BindGroupEntry { binding: 3, resource: dhb.as_entire_binding() },
wgpu::BindGroupEntry { binding: 4, resource: dwb.as_entire_binding() },
wgpu::BindGroupEntry { binding: 5, resource: d.as_entire_binding() },
wgpu::BindGroupEntry { binding: 6, resource: e0.as_entire_binding() },
],
});
run(ctx, &dp, &dbg, heads * n);
let merged = ctx.empty(n * width);
let rb = 4usize;
let fp = pipeline(ctx, "de_flash_reg", &flash_reg_src(hd, rb, true));
let fbg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None, layout: &fp.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry { binding: 0, resource: qkv.as_entire_binding() },
wgpu::BindGroupEntry { binding: 1, resource: dhb.as_entire_binding() },
wgpu::BindGroupEntry { binding: 2, resource: dwb.as_entire_binding() },
wgpu::BindGroupEntry { binding: 3, resource: merged.as_entire_binding() },
wgpu::BindGroupEntry { binding: 4, resource: d.as_entire_binding() },
wgpu::BindGroupEntry { binding: 5, resource: e.as_entire_binding() },
],
});
let wgs = heads * n.div_ceil(rb);
let (gx, gy) = ((wgs.min(65535)) as u32, wgs.div_ceil(65535) as u32);
let mut enc = ctx.device.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
{
let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
p.set_pipeline(&fp);
p.set_bind_group(0, &fbg, &[]);
p.dispatch_workgroups(gx, gy, 1);
}
ctx.queue.submit(Some(enc.finish()));
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
gpu_linear(ctx, &merged, proj_w, Some(proj_b), n, width, width)
}
pub fn sam_block_global_gpu(ctx: &GpuCtx, x: &[f32], gh: usize, gw: usize, width: usize, heads: usize, hidden: usize, eps: f32, w: &SamBlockWeights) -> Vec<f32> {
let n = gh * gw;
let xb = ctx.storage(x);
let (n1w, n1b) = (ctx.storage(&w.norm1_w), ctx.storage(&w.norm1_b));
let (n2w, n2b) = (ctx.storage(&w.norm2_w), ctx.storage(&w.norm2_b));
let (qw, qb) = (ctx.storage(&w.qkv_w), ctx.storage(&w.qkv_b));
let (pw, pb) = (ctx.storage(&w.proj_w), ctx.storage(&w.proj_b));
let (f1w, f1b) = (ctx.storage(&w.mlp_fc1_w), ctx.storage(&w.mlp_fc1_b));
let (f2w, f2b) = (ctx.storage(&w.mlp_fc2_w), ctx.storage(&w.mlp_fc2_b));
let hd = width / heads;
let rh = ctx.storage(&get_rel_pos(gh, gh, &w.rel_pos_h, hd));
let rw = ctx.storage(&get_rel_pos(gw, gw, &w.rel_pos_w, hd));
let normed = ln_buf(ctx, &xb, n, width, &n1w, &n1b, eps);
let attn = sam_attn_buf(ctx, &normed, gh, gw, width, heads, &qw, &qb, &pw, &pb, &rh, &rw);
let y = addact_buf(ctx, &xb, Some(&attn), n * width, 0);
let normed2 = ln_buf(ctx, &y, n, width, &n2w, &n2b, eps);
let fc1 = gpu_linear(ctx, &normed2, &f1w, Some(&f1b), n, width, hidden);
let act = addact_buf(ctx, &fc1, None, n * hidden, 1); let fc2 = gpu_linear(ctx, &act, &f2w, Some(&f2b), n, hidden, width);
let out = addact_buf(ctx, &y, Some(&fc2), n * width, 0);
ctx.read(&out, n * width).expect("read sam block")
}
pub fn layernorm_gpu(ctx: &GpuCtx, x: &[f32], rows: usize, c: usize, g: &[f32], b: &[f32], eps: f32) -> Vec<f32> {
let pl = pipeline(ctx, "de_ln", LAYERNORM);
let xb = ctx.storage(x);
let gb = ctx.storage(g);
let bb = ctx.storage(b);
let y = ctx.empty(rows * c);
let meta = uni(ctx, &u32x4(rows as u32, c as u32, eps.to_bits(), 0));
let bg = make_bg(ctx, &pl, &[&xb, &gb, &bb, &y], &meta);
run(ctx, &pl, &bg, rows);
ctx.read(&y, rows * c).expect("read ln")
}
pub fn linear_gpu(ctx: &GpuCtx, x: &[f32], m: usize, k: usize, n: usize, w: &[f32], b: Option<&[f32]>) -> Vec<f32> {
let xb = ctx.storage(x);
let wb = ctx.storage(w);
let bb = b.map(|bb| ctx.storage(bb));
let out = gpu_linear(ctx, &xb, &wb, bb.as_ref(), m, k, n);
ctx.read(&out, m * n).expect("read linear")
}
pub fn bench_sam_attention(ctx: &GpuCtx, gh: usize, gw: usize, width: usize, heads: usize, iters: usize) -> (f64, f64) {
let n = gh * gw;
let hd = width / heads;
let fill = |len: usize, seed: u32| -> Vec<f32> {
let mut s = seed.wrapping_add(1);
(0..len).map(|_| { s = s.wrapping_mul(1664525).wrapping_add(1013904223); ((s >> 8) as f32 / 16_777_216.0 - 0.5) * 0.2 }).collect()
};
let qkv = ctx.storage(&fill(n * 3 * width, 42));
let rh = ctx.storage(&get_rel_pos(gh, gh, &fill((2 * gh - 1) * hd, 7), hd));
let rw = ctx.storage(&get_rel_pos(gw, gw, &fill((2 * gw - 1) * hd, 8), hd));
let nv = bench_naive_attn(ctx, &qkv, &rh, &rw, n, gh, gw, width, heads, iters);
let f1 = bench_flash_attn(ctx, &qkv, &rh, &rw, n, gh, gw, width, heads, iters);
let f2 = bench_flash_attn(ctx, &qkv, &rh, &rw, n, gh, gw, width, heads, iters);
(nv, (f1 + f2) / 2.0)
}
pub fn bench_flash_attn(ctx: &GpuCtx, qkv: &wgpu::Buffer, rh: &wgpu::Buffer, rw: &wgpu::Buffer,
n: usize, gh: usize, gw: usize, width: usize, heads: usize, iters: usize) -> f64 {
let hd = width / heads;
let scale = 1.0f32 / (hd as f32).sqrt();
let d = uni(ctx, &u32x4(n as u32, heads as u32, hd as u32, gw as u32));
let e = uni(ctx, &u32x4(width as u32, gh as u32, scale.to_bits(), 0));
let e0 = uni(ctx, &u32x4(width as u32, gh as u32, 0, 0));
let dp = pipeline(ctx, "de_dhdw", DHDW);
let fp = pipeline(ctx, "de_flash_reg", &flash_reg_src(hd, 4, true));
let go = |ctx: &GpuCtx| {
let dhb = ctx.empty(heads * n * gh);
let dwb = ctx.empty(heads * n * gw);
let dbg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor { label: None, layout: &dp.get_bind_group_layout(0), entries: &[
wgpu::BindGroupEntry { binding: 0, resource: qkv.as_entire_binding() },
wgpu::BindGroupEntry { binding: 1, resource: rh.as_entire_binding() },
wgpu::BindGroupEntry { binding: 2, resource: rw.as_entire_binding() },
wgpu::BindGroupEntry { binding: 3, resource: dhb.as_entire_binding() },
wgpu::BindGroupEntry { binding: 4, resource: dwb.as_entire_binding() },
wgpu::BindGroupEntry { binding: 5, resource: d.as_entire_binding() },
wgpu::BindGroupEntry { binding: 6, resource: e0.as_entire_binding() }] });
run(ctx, &dp, &dbg, heads * n);
let merged = ctx.empty(n * width);
let fbg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor { label: None, layout: &fp.get_bind_group_layout(0), entries: &[
wgpu::BindGroupEntry { binding: 0, resource: qkv.as_entire_binding() },
wgpu::BindGroupEntry { binding: 1, resource: dhb.as_entire_binding() },
wgpu::BindGroupEntry { binding: 2, resource: dwb.as_entire_binding() },
wgpu::BindGroupEntry { binding: 3, resource: merged.as_entire_binding() },
wgpu::BindGroupEntry { binding: 4, resource: d.as_entire_binding() },
wgpu::BindGroupEntry { binding: 5, resource: e.as_entire_binding() }] });
let wgs = heads * n.div_ceil(4);
let (gx, gy) = ((wgs.min(65535)) as u32, wgs.div_ceil(65535) as u32);
let mut enc = ctx.device.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
{ let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default()); p.set_pipeline(&fp); p.set_bind_group(0, &fbg, &[]); p.dispatch_workgroups(gx, gy, 1); }
ctx.queue.submit(Some(enc.finish()));
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
};
go(ctx); let t = std::time::Instant::now();
for _ in 0..iters { go(ctx); }
t.elapsed().as_secs_f64() * 1000.0 / iters as f64
}
pub fn bench_naive_attn(ctx: &GpuCtx, qkv: &wgpu::Buffer, rh: &wgpu::Buffer, rw: &wgpu::Buffer,
n: usize, gh: usize, gw: usize, width: usize, heads: usize, iters: usize) -> f64 {
let hd = width / heads;
let scale = 1.0f32 / (hd as f32).sqrt();
let d = uni(ctx, &u32x4(n as u32, heads as u32, hd as u32, gw as u32));
let e = uni(ctx, &u32x4(width as u32, gh as u32, scale.to_bits(), 0));
let sp = pipeline(ctx, "de_sam_scores", SAM_SCORES);
let smp = pipeline(ctx, "de_softmax", SOFTMAX);
let avp = pipeline(ctx, "de_av_merge", AV_MERGE);
let go = |ctx: &GpuCtx| {
let attn = ctx.empty(heads * n * n);
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor { label: None, layout: &sp.get_bind_group_layout(0), entries: &[
wgpu::BindGroupEntry { binding: 0, resource: qkv.as_entire_binding() },
wgpu::BindGroupEntry { binding: 1, resource: rh.as_entire_binding() },
wgpu::BindGroupEntry { binding: 2, resource: rw.as_entire_binding() },
wgpu::BindGroupEntry { binding: 3, resource: attn.as_entire_binding() },
wgpu::BindGroupEntry { binding: 4, resource: d.as_entire_binding() },
wgpu::BindGroupEntry { binding: 5, resource: e.as_entire_binding() }] });
run(ctx, &sp, &bg, heads * n * n);
let sd = uni(ctx, &u32x4((heads * n) as u32, n as u32, 0, 0));
let sbg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor { label: None, layout: &smp.get_bind_group_layout(0), entries: &[
wgpu::BindGroupEntry { binding: 0, resource: attn.as_entire_binding() },
wgpu::BindGroupEntry { binding: 1, resource: sd.as_entire_binding() }] });
run(ctx, &smp, &sbg, heads * n);
let merged = ctx.empty(n * width);
let ad = uni(ctx, &u32x4(n as u32, width as u32, hd as u32, 0));
let abg = make_bg(ctx, &avp, &[&attn, qkv, &merged], &ad);
run(ctx, &avp, &abg, n * width);
};
go(ctx);
let t = std::time::Instant::now();
for _ in 0..iters { go(ctx); }
t.elapsed().as_secs_f64() * 1000.0 / iters as f64
}
#[cfg(test)]
mod tests {
use super::*;
use crate::deepencoder::sam_attention;
fn maxabs(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b).map(|(x, y)| (x - y).abs()).fold(0.0, f32::max)
}
fn fill(n: usize, seed: u32) -> Vec<f32> {
let mut s = seed.wrapping_add(1);
(0..n)
.map(|_| {
s = s.wrapping_mul(1664525).wrapping_add(1013904223);
((s >> 8) as f32 / 16_777_216.0 - 0.5) * 0.2
})
.collect()
}
#[test]
fn sam_attention_cpu_gpu_iso() {
let Ok(ctx) = GpuCtx::new() else {
eprintln!("no wgpu adapter; skipping iso test");
return;
};
let (gh, gw, width, heads) = (8usize, 8usize, 64usize, 1usize);
let n = gh * gw;
let hd = width / heads;
let x = fill(n * width, 1);
let w = SamBlockWeights {
norm1_w: vec![1.0; width], norm1_b: vec![0.0; width],
qkv_w: fill(3 * width * width, 2), qkv_b: fill(3 * width, 3),
proj_w: fill(width * width, 4), proj_b: fill(width, 5),
norm2_w: vec![1.0; width], norm2_b: vec![0.0; width],
mlp_fc1_w: vec![0.0; width], mlp_fc1_b: vec![0.0; width],
mlp_fc2_w: vec![0.0; width], mlp_fc2_b: vec![0.0; width],
rel_pos_h: fill((2 * gh - 1) * hd, 6),
rel_pos_w: fill((2 * gw - 1) * hd, 7),
};
let cpu = sam_attention(&x, gh, gw, width, heads, &w);
let gpu = sam_attention_gpu(&ctx, &x, gh, gw, width, heads, &w);
let d = maxabs(&cpu, &gpu);
println!("sam_attention CPU vs GPU max_abs = {d:.3e}");
assert!(d < 1e-4, "sam_attention CPU/GPU not iso: max_abs {d:.3e}");
}
#[test]
fn linear_cpu_gpu_iso() {
let Ok(ctx) = GpuCtx::new() else { return };
let (m, k, n) = (37usize, 53usize, 41usize);
let x = fill(m * k, 11);
let w = fill(n * k, 12);
let b = fill(n, 13);
let cpu = crate::deepencoder::linear(&x, m, k, n, &w, Some(&b));
let gpu = linear_gpu(&ctx, &x, m, k, n, &w, Some(&b));
let d = maxabs(&cpu, &gpu);
println!("linear CPU vs GPU max_abs = {d:.3e}");
assert!(d < 1e-4, "linear not iso: {d:.3e}");
}
#[test]
fn layernorm_cpu_gpu_iso() {
let Ok(ctx) = GpuCtx::new() else { return };
let (rows, c) = (29usize, 64usize);
let x = fill(rows * c, 21);
let g = fill(c, 22).iter().map(|v| 1.0 + v).collect::<Vec<_>>();
let b = fill(c, 23);
let eps = 1e-6f32;
let cpu = crate::deepencoder::layernorm(&x, rows, c, &g, &b, eps);
let gpu = layernorm_gpu(&ctx, &x, rows, c, &g, &b, eps);
let d = maxabs(&cpu, &gpu);
println!("layernorm CPU vs GPU max_abs = {d:.3e}");
assert!(d < 1e-4, "layernorm not iso: {d:.3e}");
}
fn mk_block(width: usize, heads: usize, hidden: usize, gh: usize, gw: usize, seed: u32) -> SamBlockWeights {
let hd = width / heads;
SamBlockWeights {
norm1_w: fill(width, seed).iter().map(|v| 1.0 + v).collect(), norm1_b: fill(width, seed + 1),
qkv_w: fill(3 * width * width, seed + 2), qkv_b: fill(3 * width, seed + 3),
proj_w: fill(width * width, seed + 4), proj_b: fill(width, seed + 5),
norm2_w: fill(width, seed + 6).iter().map(|v| 1.0 + v).collect(), norm2_b: fill(width, seed + 7),
mlp_fc1_w: fill(hidden * width, seed + 8), mlp_fc1_b: fill(hidden, seed + 9),
mlp_fc2_w: fill(width * hidden, seed + 10), mlp_fc2_b: fill(width, seed + 11),
rel_pos_h: fill((2 * gh - 1) * hd, seed + 12), rel_pos_w: fill((2 * gw - 1) * hd, seed + 13),
}
}
#[test]
fn sam_block_global_cpu_gpu_iso() {
use crate::deepencoder::{sam_block, DeepEncoderConfig};
let Ok(ctx) = GpuCtx::new() else { return };
let (grid, width, heads, hidden) = (6usize, 64usize, 1usize, 256usize);
let mut cfg = DeepEncoderConfig::default();
cfg.sam_width = width; cfg.sam_heads = heads; cfg.sam_mlp_ratio = 4.0; cfg.eps = 1e-6;
let w = mk_block(width, heads, hidden, grid, grid, 100);
let x = fill(grid * grid * width, 99);
let cpu = sam_block(&x, grid, &cfg, false, &w); let gpu = sam_block_global_gpu(&ctx, &x, grid, grid, width, heads, hidden, cfg.eps, &w);
let d = maxabs(&cpu, &gpu);
println!("sam_block(global) CPU vs GPU max_abs = {d:.3e}");
assert!(d < 1e-3, "sam_block not iso: {d:.3e}");
}
#[test]
#[ignore] fn perf_flash_vs_naive() {
let Ok(ctx) = GpuCtx::new() else { eprintln!("no wgpu"); return };
let (gh, gw, width, heads) = (64usize, 64usize, 768usize, 12usize);
let n = gh * gw;
let hd = width / heads;
let qkv = ctx.storage(&fill(n * 3 * width, 42));
let rh = ctx.storage(&get_rel_pos(gh, gh, &fill((2 * gh - 1) * hd, 7), hd));
let rw = ctx.storage(&get_rel_pos(gw, gw, &fill((2 * gw - 1) * hd, 8), hd));
let iters = 3;
let f1 = bench_flash_attn(&ctx, &qkv, &rh, &rw, n, gh, gw, width, heads, iters);
let nv = bench_naive_attn(&ctx, &qkv, &rh, &rw, n, gh, gw, width, heads, iters);
let f2 = bench_flash_attn(&ctx, &qkv, &rh, &rw, n, gh, gw, width, heads, iters);
let flash = (f1 + f2) / 2.0;
println!("\n=== SAM global attention (64²,768,12h) — same run, throttle-invariant ratio ===");
println!(" naive-dense (805MB n² buffer, 3 passes): {nv:7.1} ms");
println!(" register-flash (no n² buffer, portable): {flash:7.1} ms");
println!(" SPEEDUP: {:.2}× (and flash uses ~0 extra VRAM vs naive's {:.0} MB)",
nv / flash, (heads * n * n * 4) as f64 / 1.0e6);
}
#[test]
#[ignore] fn perf_full_encoder() {
use std::time::Instant;
let Ok(ctx) = GpuCtx::new() else { eprintln!("no wgpu"); return };
let time = |f: &dyn Fn()| { f(); let t = Instant::now(); for _ in 0..3 { f(); } t.elapsed().as_secs_f64() * 1000.0 / 3.0 };
let wg = mk_block(768, 12, 3072, 64, 64, 200);
let xg = fill(64 * 64 * 768, 199);
let sam_global = time(&|| { let _ = sam_block_global_gpu(&ctx, &xg, 64, 64, 768, 12, 3072, 1e-6, &wg); });
let x4 = ctx.storage(&fill(4096 * 768, 1));
let w1 = ctx.storage(&fill(3072 * 768, 2));
let w2 = ctx.storage(&fill(768 * 3072, 3));
let mlp = time(&|| {
let h = gpu_linear(&ctx, &x4, &w1, None, 4096, 768, 3072);
let _ = gpu_linear(&ctx, &h, &w2, None, 4096, 3072, 768);
});
let wc = mk_block(1024, 16, 4096, 16, 16, 300);
let xc = fill(256 * 1024, 299);
let clip = time(&|| { let _ = sam_block_global_gpu(&ctx, &xc, 16, 16, 1024, 16, 4096, 1e-6, &wc); });
let sam_win = mlp + (sam_global - mlp).max(0.0) / 17.0;
let encoder = 4.0 * sam_global + 8.0 * sam_win + 24.0 * clip;
println!("\n=== DeepEncoder per-page cost (Metal, naive-dense attn + tiled GEMM) ===");
println!(" SAM global block (64²,768,12h) {sam_global:7.1} ms × 4 = {:7.1} ms", 4.0 * sam_global);
println!(" SAM MLP-only (the LN+2 GEMMs) {mlp:7.1} ms");
println!(" SAM windowed block (est. attn/17) {sam_win:7.1} ms × 8 = {:7.1} ms", 8.0 * sam_win);
println!(" CLIP block (256,1024,16h) {clip:7.1} ms × 24 = {:7.1} ms", 24.0 * clip);
println!(" ------------------------------------------------------------");
println!(" FULL ENCODER (per 1024² page) ~{encoder:7.0} ms (~{:.1} s)", encoder / 1000.0);
println!(" (attention is naive-dense; the flash/simdgroup path would cut the SAM term)");
}
#[test]
#[ignore] fn perf_sam_block_real_scale() {
use std::time::Instant;
let Ok(ctx) = GpuCtx::new() else { eprintln!("no wgpu"); return };
let (grid, width, heads, hidden) = (64usize, 768usize, 12usize, 3072usize);
let w = mk_block(width, heads, hidden, grid, grid, 200);
let x = fill(grid * grid * width, 199);
let _ = sam_block_global_gpu(&ctx, &x, grid, grid, width, heads, hidden, 1e-6, &w);
let iters = 3;
let t0 = Instant::now();
for _ in 0..iters {
let _ = sam_block_global_gpu(&ctx, &x, grid, grid, width, heads, hidden, 1e-6, &w);
}
let ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
println!("SAM global block (64x64, 768d, 12h) GPU: {ms:.1} ms/block (register-flash attn)");
println!(" → 4 global blocks ~= {:.0} ms; windowed+CLIP add a fraction; per-page encoder is O(that)", ms * 4.0);
}
}