use anyhow::Result;
use inferencelayer::GpuCtx;
fn pipeline(ctx: &GpuCtx, label: &str, src: &str) -> wgpu::ComputePipeline {
let m = ctx
.device
.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some(label),
source: wgpu::ShaderSource::Wgsl(src.into()),
});
ctx.device
.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some(label),
layout: None,
module: &m,
entry_point: Some("main"),
compilation_options: Default::default(),
cache: None,
})
}
const GEMV_F16S: &str = r#"enable f16;
fn q4_lo(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xU8(word & 0x0F0F0F0Fu)) - 8.0; }
fn q4_hi(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xU8((word >> 4u) & 0x0F0F0F0Fu)) - 8.0; }
@group(0) @binding(0) var<storage, read> scales: array<f16>;
@group(0) @binding(1) var<storage, read> quants: 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>;
@group(0) @binding(4) var<uniform> dims: vec4<u32>;
const NR: u32 = 16u;
const LANES: u32 = 16u;
const KC: u32 = 4u;
const TB: u32 = 16u;
const TV: u32 = 128u;
var<workgroup> xs: array<vec4<f32>, 512>;
var<workgroup> red: array<f32, 256>;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid3: vec3<u32>) {
let lid = lid3.x;
let row = wid.x * NR + lid / LANES;
let lane = lid % LANES;
let col0 = wid.y * KC;
let m = dims.x; let n = dims.y; let ncols = dims.w;
let nblk = n / 32u;
let xstride = n / 4u;
var acc = vec4<f32>(0.0);
let ntiles = (nblk + TB - 1u) / TB;
for (var t = 0u; t < ntiles; t = t + 1u) {
for (var j = 0u; j < 2u; j = j + 1u) {
let idx = lid * 2u + j;
let cc = idx / TV;
let e = t * TV + (idx % TV);
var v = vec4<f32>(0.0);
if (col0 + cc < ncols && e < xstride) { v = x[(col0 + cc) * xstride + e]; }
xs[idx] = v;
}
workgroupBarrier();
let b = t * TB + lane;
if (row < m && b < nblk) {
let d = f32(scales[row * nblk + b]);
let q = quants[row * nblk + b];
let l0 = q4_lo(q.x); let h0 = q4_hi(q.x);
let l1 = q4_lo(q.y); let h1 = q4_hi(q.y);
let l2 = q4_lo(q.z); let h2 = q4_hi(q.z);
let l3 = q4_lo(q.w); let h3 = q4_hi(q.w);
let xb = lane * 8u;
for (var cc = 0u; cc < KC; cc = cc + 1u) {
let base = cc * TV + xb;
var s = dot(l0, xs[base]) + dot(h0, xs[base + 4u]);
s = s + dot(l1, xs[base + 1u]) + dot(h1, xs[base + 5u]);
s = s + dot(l2, xs[base + 2u]) + dot(h2, xs[base + 6u]);
s = s + dot(l3, xs[base + 3u]) + dot(h3, xs[base + 7u]);
acc[cc] = acc[cc] + d * s;
}
}
workgroupBarrier();
}
for (var cc = 0u; cc < KC; cc = cc + 1u) {
red[lid] = acc[cc];
workgroupBarrier();
if (lane < 8u) { red[lid] = red[lid] + red[lid + 8u]; }
workgroupBarrier();
if (lane < 4u) { red[lid] = red[lid] + red[lid + 4u]; }
workgroupBarrier();
if (lane < 2u) { red[lid] = red[lid] + red[lid + 2u]; }
workgroupBarrier();
if (lane == 0u && row < m && col0 + cc < ncols) {
let v = red[lid] + red[lid + 1u];
let yo = (col0 + cc) * m + row;
y[yo] = v;
}
workgroupBarrier();
}
}
"#;
const MOE_PURE_READ: &str = r#"enable f16;
@group(0) @binding(0) var<storage, read> s1: array<f16>;
@group(0) @binding(1) var<storage, read> q1: array<vec4<u32>>;
@group(0) @binding(2) var<storage, read> s3: array<f16>;
@group(0) @binding(3) var<storage, read> q3: array<vec4<u32>>;
@group(0) @binding(4) var<storage, read> x: array<vec4<f32>>;
@group(0) @binding(5) var<storage, read> sel: array<u32>;
@group(0) @binding(6) var<storage, read_write> y: array<f32>;
@group(0) @binding(7) var<uniform> dims: vec4<u32>;
@group(0) @binding(8) var<uniform> epsm: vec4<f32>;
@compute @workgroup_size(32)
fn main(@builtin(workgroup_id) wid: vec3<u32>,
@builtin(subgroup_invocation_id) sid: u32) {
let mi = dims.x; let h = dims.y;
let col = wid.y / dims.w;
let slot = wid.y % dims.w;
let eid = sel[col * dims.w + slot];
let row0 = wid.x * 4u;
let nblk = h / 32u;
var qacc = 0u;
var sacc = 0.0;
for (var b = sid; b < nblk; b = b + 32u) {
for (var r = 0u; r < 4u; r = r + 1u) {
let row = eid * mi + min(row0 + r, mi - 1u);
let q = q1[row * nblk + b];
let p = q3[row * nblk + b];
qacc = qacc + q.x + q.y + q.z + q.w + p.x + p.y + p.z + p.w;
sacc = sacc + f32(s1[row * nblk + b]) + f32(s3[row * nblk + b]);
}
}
// wgpu derives the bind-group layout from what the shader USES, and the caller binds all nine
// buffers — so x and epsm must be referenced or the layout comes back with seven and the bind
// group fails validation. One BROADCAST read each (every thread hits the same address = a single
// cache line) keeps them live for ~free, and critically does NOT reintroduce the strided,
// per-thread x loads whose cost is the whole point of excluding them here.
let keep = x[0].x * epsm.x;
let t = subgroupAdd(sacc) + f32(subgroupAdd(qacc) & 1u) + keep * 0.0;
if (sid == 0u) { y[(col * dims.w + slot) * mi + row0] = t; }
}
"#;
fn moe_lab(ctx: &GpuCtx) -> Result<()> {
const H: usize = 2048; const MI: usize = 512; const TOPK: usize = 8; const E: usize = 64; const BLK: usize = 32;
let nblk = H / BLK;
let rows = E * MI;
let quants: Vec<u32> = (0..rows * nblk * 4)
.map(|i| (i as u32).wrapping_mul(2654435761))
.collect();
let scales: Vec<u16> = (0..rows * nblk)
.map(|i| half::f16::from_f32(0.01 + (i % 7) as f32 * 0.001).to_bits())
.collect();
let s1 = ctx.storage_bytes(bytemuck::cast_slice(&scales));
let q1 = ctx.storage(bytemuck::cast_slice(&quants));
let s3 = ctx.storage_bytes(bytemuck::cast_slice(&scales));
let q3 = ctx.storage(bytemuck::cast_slice(&quants));
let prod = inferencelayer::forward::moe_gate_q4_lcpp_src();
eprintln!("\n=== MoE expert gate+up GEMV (h={H} mi={MI} top_k={TOPK}, Q4_0) ===");
eprintln!(
"cols variant ms GB/s spread (traffic = cols × top_k × (w1+w3); spread = worst/best, high ⇒ busy machine ⇒ distrust)"
);
for ncols in [1usize, 2, 3, 5, 9] {
let sel: Vec<u32> = (0..ncols * TOPK).map(|i| (i % E) as u32).collect();
let x: Vec<f32> = (0..ncols * H).map(|i| (i % 97) as f32 * 0.01).collect();
let selb = ctx.storage(bytemuck::cast_slice(&sel));
let xb = ctx.storage(&x);
let yb = ctx.storage(&vec![0f32; ncols * TOPK * MI]);
let dims = wgpu::util::DeviceExt::create_buffer_init(
&ctx.device,
&wgpu::util::BufferInitDescriptor {
label: None,
contents: bytemuck::cast_slice(&[MI as u32, H as u32, 0u32, TOPK as u32]),
usage: wgpu::BufferUsages::UNIFORM,
},
);
let epsm = wgpu::util::DeviceExt::create_buffer_init(
&ctx.device,
&wgpu::util::BufferInitDescriptor {
label: None,
contents: bytemuck::cast_slice(&[1e-6f32, 0.0, 0.0, 0.0]),
usage: wgpu::BufferUsages::UNIFORM,
},
);
let bytes = ncols as f64 * TOPK as f64 * 2.0 * (MI * H) as f64 * (0.5 + 2.0 / BLK as f64);
let variants: Vec<(String, String)> = vec![
("prod32".into(), prod.clone()),
("sg128".into(), moe_gate_variant(&prod, 4, false)),
("xshare32".into(), moe_gate_variant(&prod, 1, true)),
("sg128x".into(), moe_gate_variant(&prod, 4, true)),
("sg256x".into(), moe_gate_variant(&prod, 8, true)),
("pureread".into(), MOE_PURE_READ.to_string()),
];
for (label, src) in &variants {
if label.contains('x') && label != "pureread" {
assert!(
src.contains("workgroupBarrier()") && src.contains("let v0 = xs[xb];"),
"{label}: x-staging rewrite did not apply — production source shape changed"
);
}
}
for (label, src) in &variants {
let nsg: u32 = match label.as_str() {
"sg128" | "sg128x" => 4,
"sg256x" => 8,
_ => 1, };
let pl = pipeline(ctx, label, src);
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pl.get_bind_group_layout(0),
entries: &[&s1, &q1, &s3, &q3, &xb, &selb, &yb, &dims, &epsm]
.iter()
.enumerate()
.map(|(i, b)| wgpu::BindGroupEntry {
binding: i as u32,
resource: b.as_entire_binding(),
})
.collect::<Vec<_>>(),
});
let gx = (MI as u32).div_ceil(4 * nsg);
let gy = (ncols * TOPK) as u32;
let mut times: Vec<f64> = Vec::new();
for rep in 0..6 {
let reps = if rep == 0 { 3 } else { 30 }; let t0 = std::time::Instant::now();
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut p = enc.begin_compute_pass(&Default::default());
for _ in 0..reps {
p.set_pipeline(&pl);
p.set_bind_group(0, &bg, &[]);
p.dispatch_workgroups(gx, gy, 1);
}
}
ctx.queue.submit([enc.finish()]);
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
if rep > 0 {
times.push(t0.elapsed().as_secs_f64() / reps as f64);
}
}
let best = times.iter().copied().fold(f64::MAX, f64::min);
let worst = times.iter().copied().fold(0.0, f64::max);
eprintln!(
"{ncols:>4} {label:<9} {:>6.3} {:>6.1} {:>4.0}%",
best * 1e3,
bytes / best / 1e9,
(worst / best - 1.0) * 100.0
);
}
}
Ok(())
}
fn moe_gate_variant(prod: &str, nsg: u32, stage_x: bool) -> String {
let mut s = if nsg == 1 {
prod.to_string()
} else {
moe_gate_multi_sg(prod, nsg)
};
if stage_x {
let (idx, stride) = if nsg == 1 {
("sid".to_string(), 32)
} else {
("lid.x".to_string(), nsg * 32)
};
s = s.replace(
"@compute @workgroup_size(",
"var<workgroup> xs: array<vec4<f32>, 512>;\n@compute @workgroup_size(",
);
s = s.replace(
" let xoff = col * (h / 4u);\n var ag = vec4<f32>(0.0);",
&format!(
" let xoff = col * (h / 4u);\n let nvec = h / 4u;\n for (var i = {idx}; i < nvec; i = i + {stride}u) {{ xs[i] = x[xoff + i]; }}\n workgroupBarrier();\n var ag = vec4<f32>(0.0);"
),
);
s = s.replace(
" let xb = xoff + b * 8u;\n let v0 = x[xb]; let v1 = x[xb + 1u];\n let v2 = x[xb + 2u]; let v3 = x[xb + 3u];\n let v4 = x[xb + 4u]; let v5 = x[xb + 5u];\n let v6 = x[xb + 6u]; let v7 = x[xb + 7u];",
" let xb = b * 8u;\n let v0 = xs[xb]; let v1 = xs[xb + 1u];\n let v2 = xs[xb + 2u]; let v3 = xs[xb + 3u];\n let v4 = xs[xb + 4u]; let v5 = xs[xb + 5u];\n let v6 = xs[xb + 6u]; let v7 = xs[xb + 7u];",
);
}
s
}
fn moe_gate_multi_sg(prod: &str, nsg: u32) -> String {
prod.replace(
"@compute @workgroup_size(32)\nfn main(@builtin(workgroup_id) wid: vec3<u32>,\n @builtin(subgroup_invocation_id) sid: u32) {",
&format!(
"@compute @workgroup_size({})\nfn main(@builtin(workgroup_id) wid: vec3<u32>,\n @builtin(local_invocation_id) lid: vec3<u32>,\n @builtin(subgroup_invocation_id) sid: u32) {{",
nsg * 32
),
)
.replace(
" let row0 = wid.x * 4u;",
&format!(" let row0 = (wid.x * {nsg}u + lid.x / 32u) * 4u;"),
)
}
fn main() -> Result<()> {
let ctx = GpuCtx::new()?;
eprintln!("adapter subgroups: {}", ctx.subgroups);
if std::env::args().any(|a| a == "--moe") {
return moe_lab(&ctx);
}
let gemv_u32s = GEMV_F16S
.replace("enable f16;\n", "")
.replace(
"@group(0) @binding(0) var<storage, read> scales: array<f16>;",
"@group(0) @binding(0) var<storage, read> scales: array<u32>;",
)
.replace(
"let d = f32(scales[row * nblk + b]);",
"let sw = scales[(row * nblk + b) >> 1u];\n let d = unpack2x16float(sw)[(row * nblk + b) & 1u];",
);
let gemv_f32s = GEMV_F16S
.replace("enable f16;\n", "")
.replace(
"@group(0) @binding(0) var<storage, read> scales: array<f16>;",
"@group(0) @binding(0) var<storage, read> scales: array<f32>;",
)
.replace(
"let d = f32(scales[row * nblk + b]);",
"let d = scales[row * nblk + b];",
);
let gemv_noscale = GEMV_F16S
.replace("enable f16;\n", "")
.replace(
"@group(0) @binding(0) var<storage, read> scales: array<f16>;",
"@group(0) @binding(0) var<storage, read> scales: array<u32>;",
)
.replace(
"let d = f32(scales[row * nblk + b]);",
"let d = 1.0 + f32(min(scales[0], 0u));",
);
let pure_read = r#"
@group(0) @binding(0) var<storage, read> scales: array<u32>;
@group(0) @binding(1) var<storage, read> quants: 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>;
@group(0) @binding(4) var<uniform> dims: vec4<u32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
let m = dims.x; let n = dims.y;
let nblk = n / 32u;
let row = gid.x / 16u;
let lane = gid.x % 16u;
if (row >= m) { return; }
var acc = vec4<u32>();
for (var b = lane; b < nblk; b = b + 16u) {
acc = acc + quants[row * nblk + b];
}
if ((acc.x | acc.y | acc.z | acc.w) == 123456789u + min(scales[0], 0u)) { y[row] = x[0].x; }
}
"#;
let nobar = r#"
fn q4_lo(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xU8(word & 0x0F0F0F0Fu)) - 8.0; }
fn q4_hi(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xU8((word >> 4u) & 0x0F0F0F0Fu)) - 8.0; }
@group(0) @binding(0) var<storage, read> scales: array<u32>;
@group(0) @binding(1) var<storage, read> quants: 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>;
@group(0) @binding(4) var<uniform> dims: vec4<u32>;
var<workgroup> red: array<f32, 256>;
const KC: u32 = 4u;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid3: vec3<u32>) {
let lid = lid3.x;
let m = dims.x; let n = dims.y; let ncols = dims.w;
let nblk = n / 32u;
let xstride = n / 4u;
let row = wid.x * 16u + lid / 16u;
let lane = lid % 16u;
let col0 = wid.y * KC;
var acc = vec4<f32>(0.0);
for (var b = lane; b < nblk; b = b + 16u) {
let sw = scales[(row * nblk + b) >> 1u];
let d = unpack2x16float(sw)[(row * nblk + b) & 1u];
let q = quants[row * nblk + b];
let l0 = q4_lo(q.x); let h0 = q4_hi(q.x);
let l1 = q4_lo(q.y); let h1 = q4_hi(q.y);
let l2 = q4_lo(q.z); let h2 = q4_hi(q.z);
let l3 = q4_lo(q.w); let h3 = q4_hi(q.w);
let xb = b * 8u;
for (var cc = 0u; cc < KC; cc = cc + 1u) {
if (col0 + cc < ncols) {
let base = (col0 + cc) * xstride + xb;
var s = dot(l0, x[base]) + dot(h0, x[base + 4u]);
s = s + dot(l1, x[base + 1u]) + dot(h1, x[base + 5u]);
s = s + dot(l2, x[base + 2u]) + dot(h2, x[base + 6u]);
s = s + dot(l3, x[base + 3u]) + dot(h3, x[base + 7u]);
acc[cc] = acc[cc] + d * s;
}
}
}
if (row >= m) { return; }
for (var cc = 0u; cc < KC; cc = cc + 1u) {
red[lid] = acc[cc];
workgroupBarrier();
if (lane < 8u) { red[lid] = red[lid] + red[lid + 8u]; }
workgroupBarrier();
if (lane < 4u) { red[lid] = red[lid] + red[lid + 4u]; }
workgroupBarrier();
if (lane < 2u) { red[lid] = red[lid] + red[lid + 2u]; }
workgroupBarrier();
if (lane == 0u && col0 + cc < ncols) {
y[(col0 + cc) * m + row] = red[lid] + red[lid + 1u];
}
workgroupBarrier();
}
}
"#;
let v3_kc4 = r#"
fn q4_lo(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xU8(word & 0x0F0F0F0Fu)) - 8.0; }
fn q4_hi(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xU8((word >> 4u) & 0x0F0F0F0Fu)) - 8.0; }
@group(0) @binding(0) var<storage, read> scales: array<f32>;
@group(0) @binding(1) var<storage, read> quants: 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>;
@group(0) @binding(4) var<uniform> dims: vec4<u32>;
const KC: u32 = 4u;
const XV: u32 = 512u; // vec4 per column (h/4)
var<workgroup> xs: array<vec4<f32>, 2048>; // KC*XV = 32 KiB
var<workgroup> red: array<f32, 256>;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid3: vec3<u32>) {
let lid = lid3.x;
let m = dims.x; let n = dims.y; let ncols = dims.w;
let nblk = n / 32u;
let xstride = n / 4u;
let row = wid.x * 16u + lid / 16u;
let lane = lid % 16u;
let col0 = wid.y * KC;
// Stage KC full columns of x ONCE (8 vec4/thread/col), one barrier total.
for (var cc = 0u; cc < KC; cc = cc + 1u) {
for (var e = lid; e < xstride; e = e + 256u) {
var v = vec4<f32>(0.0);
if (col0 + cc < ncols) { v = x[(col0 + cc) * xstride + e]; }
xs[cc * XV + e] = v;
}
}
workgroupBarrier();
var acc = vec4<f32>(0.0);
for (var b = lane; b < nblk; b = b + 16u) {
let d = scales[row * nblk + b];
let q = quants[row * nblk + b];
let l0 = q4_lo(q.x); let h0 = q4_hi(q.x);
let l1 = q4_lo(q.y); let h1 = q4_hi(q.y);
let l2 = q4_lo(q.z); let h2 = q4_hi(q.z);
let l3 = q4_lo(q.w); let h3 = q4_hi(q.w);
let xb = b * 8u;
for (var cc = 0u; cc < KC; cc = cc + 1u) {
let base = cc * XV + xb;
var s = dot(l0, xs[base]) + dot(h0, xs[base + 4u]);
s = s + dot(l1, xs[base + 1u]) + dot(h1, xs[base + 5u]);
s = s + dot(l2, xs[base + 2u]) + dot(h2, xs[base + 6u]);
s = s + dot(l3, xs[base + 3u]) + dot(h3, xs[base + 7u]);
acc[cc] = acc[cc] + d * s;
}
}
if (row >= m) { return; }
for (var cc = 0u; cc < KC; cc = cc + 1u) {
red[lid] = acc[cc];
workgroupBarrier();
if (lane < 8u) { red[lid] = red[lid] + red[lid + 8u]; }
workgroupBarrier();
if (lane < 4u) { red[lid] = red[lid] + red[lid + 4u]; }
workgroupBarrier();
if (lane < 2u) { red[lid] = red[lid] + red[lid + 2u]; }
workgroupBarrier();
if (lane == 0u && col0 + cc < ncols) {
y[(col0 + cc) * m + row] = red[lid] + red[lid + 1u];
}
workgroupBarrier();
}
}
"#;
let v3_kc2 = v3_kc4
.replace("const KC: u32 = 4u;", "const KC: u32 = 2u;")
.replace(
"var<workgroup> xs: array<vec4<f32>, 2048>; // KC*XV = 32 KiB",
"var<workgroup> xs: array<vec4<f32>, 1024>; // KC*XV = 16 KiB",
);
let lcpp2r = r#"enable f16;
fn q4_lo(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xU8(word & 0x0F0F0F0Fu)) - 8.0; }
fn q4_hi(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xU8((word >> 4u) & 0x0F0F0F0Fu)) - 8.0; }
@group(0) @binding(0) var<storage, read> scales: array<f16>;
@group(0) @binding(1) var<storage, read> quants: 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>;
@group(0) @binding(4) var<uniform> dims: vec4<u32>;
@compute @workgroup_size(32)
fn main(@builtin(workgroup_id) wid: vec3<u32>,
@builtin(subgroup_invocation_id) sid: u32) {
let m = dims.x; let n = dims.y;
let nblk = n / 32u;
let xstride = n / 4u;
let row0 = wid.x * 2u;
let col = wid.y;
let xoff = col * xstride;
var a0 = 0.0;
var a1 = 0.0;
for (var b = sid; b < nblk; b = b + 32u) {
let xb = xoff + b * 8u;
let v0 = x[xb]; let v1 = x[xb + 1u];
let v2 = x[xb + 2u]; let v3 = x[xb + 3u];
let v4 = x[xb + 4u]; let v5 = x[xb + 5u];
let v6 = x[xb + 6u]; let v7 = x[xb + 7u];
{
let d = f32(scales[row0 * nblk + b]);
let q = quants[row0 * nblk + b];
var s = dot(q4_lo(q.x), v0) + dot(q4_hi(q.x), v4);
s = s + dot(q4_lo(q.y), v1) + dot(q4_hi(q.y), v5);
s = s + dot(q4_lo(q.z), v2) + dot(q4_hi(q.z), v6);
s = s + dot(q4_lo(q.w), v3) + dot(q4_hi(q.w), v7);
a0 = a0 + d * s;
}
{
let d = f32(scales[(row0 + 1u) * nblk + b]);
let q = quants[(row0 + 1u) * nblk + b];
var s = dot(q4_lo(q.x), v0) + dot(q4_hi(q.x), v4);
s = s + dot(q4_lo(q.y), v1) + dot(q4_hi(q.y), v5);
s = s + dot(q4_lo(q.z), v2) + dot(q4_hi(q.z), v6);
s = s + dot(q4_lo(q.w), v3) + dot(q4_hi(q.w), v7);
a1 = a1 + d * s;
}
}
let t0 = subgroupAdd(a0);
let t1 = subgroupAdd(a1);
if (sid == 0u) {
y[col * m + row0] = t0;
if (row0 + 1u < m) { y[col * m + row0 + 1u] = t1; }
}
}
"#;
let lcpp4r = r#"enable f16;
fn q4_lo(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xU8(word & 0x0F0F0F0Fu)) - 8.0; }
fn q4_hi(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xU8((word >> 4u) & 0x0F0F0F0Fu)) - 8.0; }
@group(0) @binding(0) var<storage, read> scales: array<f16>;
@group(0) @binding(1) var<storage, read> quants: 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>;
@group(0) @binding(4) var<uniform> dims: vec4<u32>;
@compute @workgroup_size(32)
fn main(@builtin(workgroup_id) wid: vec3<u32>,
@builtin(subgroup_invocation_id) sid: u32) {
let m = dims.x; let n = dims.y;
let nblk = n / 32u;
let xstride = n / 4u;
let row0 = wid.x * 4u;
let col = wid.y;
let xoff = col * xstride;
var acc = vec4<f32>(0.0);
for (var b = sid; b < nblk; b = b + 32u) {
let xb = xoff + b * 8u;
let v0 = x[xb]; let v1 = x[xb + 1u];
let v2 = x[xb + 2u]; let v3 = x[xb + 3u];
let v4 = x[xb + 4u]; let v5 = x[xb + 5u];
let v6 = x[xb + 6u]; let v7 = x[xb + 7u];
for (var r = 0u; r < 4u; r = r + 1u) {
let row = row0 + r;
let d = f32(scales[row * nblk + b]);
let q = quants[row * nblk + b];
var s = dot(q4_lo(q.x), v0) + dot(q4_hi(q.x), v4);
s = s + dot(q4_lo(q.y), v1) + dot(q4_hi(q.y), v5);
s = s + dot(q4_lo(q.z), v2) + dot(q4_hi(q.z), v6);
s = s + dot(q4_lo(q.w), v3) + dot(q4_hi(q.w), v7);
acc[r] = acc[r] + d * s;
}
}
let tot = subgroupAdd(acc);
if (sid == 0u) {
for (var r = 0u; r < 4u; r = r + 1u) {
if (row0 + r < m) { y[col * m + row0 + r] = tot[r]; }
}
}
}
"#;
let cases: &[(&str, usize, usize)] = &[
("in_proj 12352x2048", 12352, 2048),
("head 248320x2048", 248320, 2048),
];
{
let chain = pipeline(
&ctx,
"chain",
"@group(0) @binding(0) var<storage, read_write> y: array<f32>;\n@compute @workgroup_size(32)\nfn main(@builtin(local_invocation_id) l: vec3<u32>) { if (l.x == 0u) { y[0] = y[0] + 1.0; } }",
);
let yb = ctx.storage(&[0f32; 4]);
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &chain.get_bind_group_layout(0),
entries: &[wgpu::BindGroupEntry {
binding: 0,
resource: yb.as_entire_binding(),
}],
});
for n in [50usize, 300, 1000] {
for rep in 0..2 {
let t0 = std::time::Instant::now();
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut p = enc.begin_compute_pass(&Default::default());
for _ in 0..n {
p.set_pipeline(&chain);
p.set_bind_group(0, &bg, &[]);
p.dispatch_workgroups(1, 1, 1);
}
}
ctx.queue.submit([enc.finish()]);
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
if rep == 1 {
let dt = t0.elapsed().as_secs_f64();
eprintln!(
"dependent-chain n={n}: {:.2} ms total = {:.1} µs/dispatch",
dt * 1e3,
dt * 1e6 / n as f64
);
}
}
}
}
{
let tiny = pipeline(
&ctx,
"tiny",
"@group(0) @binding(0) var<storage, read_write> y: array<f32>;\n@compute @workgroup_size(1)\nfn main() { y[0] = y[0] + 1.0; }",
);
let yb = ctx.storage(&[0f32; 4]);
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &tiny.get_bind_group_layout(0),
entries: &[wgpu::BindGroupEntry {
binding: 0,
resource: yb.as_entire_binding(),
}],
});
for mode in ["wait", "spin"] {
let t0 = std::time::Instant::now();
let reps = 50;
for _ in 0..reps {
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut p = enc.begin_compute_pass(&Default::default());
p.set_pipeline(&tiny);
p.set_bind_group(0, &bg, &[]);
p.dispatch_workgroups(1, 1, 1);
}
let _idx = ctx.queue.submit([enc.finish()]);
if mode == "wait" {
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
} else {
loop {
match ctx.device.poll(wgpu::PollType::Poll) {
Ok(s) if s.is_queue_empty() => break,
_ => std::hint::spin_loop(),
}
}
}
}
eprintln!(
"submit+{mode} turnaround: {:.0} µs each",
t0.elapsed().as_secs_f64() * 1e6 / reps as f64
);
}
}
for (name, m, n) in cases {
let (m, n) = (*m, *n);
let nblk = n / 32;
let scales_f16: Vec<u8> = (0..m * nblk * 2).map(|i| (i % 251) as u8).collect();
let quants: Vec<u32> = (0..m * nblk * 4)
.map(|i| (i as u32).wrapping_mul(2654435761))
.collect();
let scales_f32: Vec<f32> = (0..m * nblk)
.map(|i| {
half::f16::from_bits(u16::from_le_bytes([
scales_f16[2 * i],
scales_f16[2 * i + 1],
]))
.to_f32()
})
.collect();
for ncols in [1usize, 5, 16] {
let x: Vec<f32> = (0..ncols * n).map(|i| (i % 97) as f32 * 0.01).collect();
let sb_f16 = ctx.storage_bytes(&scales_f16);
let sb_f32 = ctx.storage(&scales_f32);
let qb = ctx.storage(bytemuck::cast_slice(&quants));
let xb = ctx.storage(&x);
let yb = ctx.storage(&vec![0f32; ncols * m]);
let dims = wgpu::util::DeviceExt::create_buffer_init(
&ctx.device,
&wgpu::util::BufferInitDescriptor {
label: None,
contents: bytemuck::cast_slice(&[m as u32, n as u32, 0u32, ncols as u32]),
usage: wgpu::BufferUsages::UNIFORM,
},
);
let variants: &[(&str, &str, &wgpu::Buffer, u32, u32)] = &[
("f16s", GEMV_F16S, &sb_f16, 4, 16),
("u32s", &gemv_u32s, &sb_f16, 4, 16),
("f32s", &gemv_f32s, &sb_f32, 4, 16),
("nosc", &gemv_noscale, &sb_f16, 4, 16),
("pure", pure_read, &sb_f16, 4, 16),
("nbar", nobar, &sb_f16, 4, 16),
("v3k4", v3_kc4, &sb_f32, 4, 16),
("v3k2", &v3_kc2, &sb_f32, 2, 16),
("l2r", lcpp2r, &sb_f16, 1, 2),
("l4r", lcpp4r, &sb_f16, 1, 4),
];
for (vn, src, sbuf, kc, rpw) in variants {
let pl = pipeline(&ctx, vn, src);
let entries: Vec<wgpu::BindGroupEntry> = [*sbuf, &qb, &xb, &yb, &dims]
.iter()
.enumerate()
.map(|(i, b)| wgpu::BindGroupEntry {
binding: i as u32,
resource: b.as_entire_binding(),
})
.collect();
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pl.get_bind_group_layout(0),
entries: &entries,
});
let gx = (m as u32).div_ceil(*rpw);
let gy = (ncols as u32).div_ceil(*kc);
for rep in 0..2 {
let reps = if rep == 0 { 3 } else { 30 };
let t0 = std::time::Instant::now();
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut p = enc.begin_compute_pass(&Default::default());
for _ in 0..reps {
p.set_pipeline(&pl);
p.set_bind_group(0, &bg, &[]);
p.dispatch_workgroups(gx, gy, 1);
}
}
ctx.queue.submit([enc.finish()]);
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
if rep == 1 {
let dt = t0.elapsed().as_secs_f64() / reps as f64;
let wbytes = (m * nblk * (16 + 2)) as f64 * gy as f64;
eprintln!(
"{name} ncols={ncols} {vn}: {:.0} µs/dispatch {:.0} GB/s weight-traffic",
dt * 1e6,
wbytes / dt / 1e9
);
}
}
}
}
}
Ok(())
}