#![allow(clippy::needless_range_loop)]
#![allow(clippy::too_many_arguments)]
#![allow(clippy::type_complexity)]
#![allow(clippy::manual_clamp)]
#![allow(clippy::manual_memcpy)]
#![allow(clippy::if_same_then_else)]
#![allow(clippy::large_enum_variant)]
#![allow(clippy::unnecessary_unwrap)]
#![allow(dead_code)]
#![allow(unused_assignments)]
#![allow(private_interfaces)]
use crate::forward::ShaderModuleTuned as _;
use anyhow::{Result, anyhow};
use wgpu::util::DeviceExt;
#[cfg(not(target_arch = "wasm32"))]
pub mod bench_guard;
pub mod deltanet;
#[cfg(feature = "encoder-cpu")]
pub mod conv2d;
#[cfg(feature = "encoder-cpu")]
pub mod cpu_gemm;
#[cfg(feature = "encoder-cpu")]
pub mod deepencoder;
#[cfg(feature = "encoder-cpu")]
pub mod deepencoder_gpu;
#[cfg(feature = "encoder-cpu")]
pub mod embed_engine;
pub mod encoder;
#[cfg(feature = "encoder-cpu")]
pub mod encoder_cpu;
pub mod encoder_weights;
pub mod forward;
pub mod gguf;
pub mod iq_tables;
#[cfg(not(target_arch = "wasm32"))]
pub mod gguf_write;
#[cfg(feature = "encoder-cpu")]
pub mod siglip_vision;
pub mod kld;
pub mod simd;
#[cfg(feature = "encoder-cpu")]
pub mod simd_math;
#[cfg(feature = "encoder-cpu")]
pub mod gliner;
#[cfg(feature = "encoder-cpu")]
mod gliner_gpu;
#[cfg(feature = "encoder-cpu")]
pub mod gliner2;
#[cfg(feature = "encoder-cpu")]
pub mod glinker;
#[cfg(feature = "encoder-cpu")]
pub mod glinerrelex;
#[cfg(feature = "encoder-cpu")]
pub mod parakeet;
#[cfg(feature = "encoder-cpu")]
pub mod rtdetr;
#[cfg(feature = "encoder-cpu")]
pub mod cv_resize;
#[cfg(feature = "encoder-cpu")]
pub mod tableformer;
#[cfg(feature = "encoder-cpu")]
pub mod rtdetr_gpu;
#[cfg(feature = "encoder-cpu")]
pub mod parakeet_gpu;
#[cfg(feature = "encoder-cpu")]
pub mod whisper;
#[cfg(feature = "encoder-cpu")]
pub mod ark_asr;
#[cfg(feature = "encoder-cpu")]
pub mod whisper_gpu;
#[cfg(feature = "encoder-cpu")]
pub mod diarize;
#[cfg(feature = "cudarc")]
pub mod cuda_tower;
#[cfg(feature = "encoder-cpu")]
pub mod vision_glm;
#[cfg(feature = "encoder-cpu")]
pub mod vision_glm_gpu;
#[cfg(feature = "encoder-cpu")]
pub mod ast;
#[cfg(feature = "encoder-cpu")]
pub mod speecht5;
#[cfg(feature = "encoder-cpu")]
pub mod vits;
#[cfg(feature = "encoder-cpu")]
pub mod kokoro;
#[cfg(feature = "encoder-cpu")]
pub mod pocket_tts;
#[cfg(feature = "encoder-cpu")]
pub mod pocket_tts_gpu;
#[cfg(feature = "encoder-cpu")]
pub mod mimi;
#[cfg(feature = "encoder-cpu")]
pub mod mimi_gpu;
pub mod realtime;
#[cfg(feature = "encoder-cpu")]
pub mod moshi_lm;
#[cfg(feature = "encoder-cpu")]
pub mod moshi_gpu;
#[cfg(feature = "encoder-cpu")]
pub mod chatterbox;
#[cfg(feature = "encoder-cpu")]
pub mod qwen3tts;
#[cfg(feature = "encoder-cpu")]
pub mod qwen3tts_codec;
#[cfg(feature = "encoder-cpu")]
pub mod qwen3tts_spk;
#[cfg(feature = "encoder-cpu")]
pub mod qwen3tts_gpu;
#[cfg(feature = "webrtc-media")]
pub mod opus_rtc;
#[cfg(feature = "webrtc-media")]
pub mod webrtc_rtc;
#[cfg(feature = "encoder-cpu")]
pub mod vision;
#[cfg(feature = "encoder-cpu")]
pub mod vision_gpu;
#[cfg(feature = "encoder-cpu")]
pub mod diffusion_gemma;
#[cfg(feature = "encoder-cpu")]
pub mod diffusion_gemma_gpu;
pub mod grammar;
pub mod pooling;
pub mod reference;
pub mod reference_qwen35;
pub mod sampling;
pub mod server;
#[cfg(feature = "server")]
pub mod serve;
#[cfg(feature = "net")]
pub mod replica;
pub mod turboquant;
pub mod vocab;
pub mod weights;
#[cfg(feature = "net")]
pub mod shard;
#[cfg(feature = "net")]
pub mod shard_serve;
#[cfg(feature = "webrtc")]
pub mod webrtc_native;
#[cfg(feature = "encoder-cpu")]
pub use embed_engine::EmbedEngine;
pub use encoder::{EncKernels, EncoderGpu};
#[cfg(feature = "encoder-cpu")]
pub use encoder_cpu::CpuEncoder;
pub use encoder_weights::{
Act, EncArch, EncBatch, EncoderConfig, MaskKind, MlpKind, NormKind, PosKind,
};
pub use forward::{
BatchCol, BatchPlan, EngineOpts, Lfm2Gpu, MtpEngine, SpecGrammar, SpecPlan, StageBatchOut,
StageOut,
};
pub use pooling::{EmbedOut, Pooling};
#[cfg(feature = "net")]
pub use replica::{FleetPlan, Replica, ReplicaPolicy, max_replicas, plan_replicas};
pub use sampling::prompt_lookup_drafts;
pub use server::{
Emission, FinishReason, KvPool, RadixIndex, RequestParams, SamplingParams, Scheduler,
ServeStats, generate_once,
};
#[cfg(feature = "net")]
pub use shard::{
Fleet, FleetRegistry, LoadAck, ShardClient, ShardedPipeline, StageInfo, WorkerOptions,
accept_fleet, decode_greedy_mtp, decode_greedy_mtp_duo, decode_greedy_mtp_multi, plan_split,
run_worker, validate_chain, webrtc_chain, webrtc_pair,
};
#[cfg(feature = "net")]
pub use shard_serve::{PipelineServe, ReplicatedServe, StageStat};
pub use vocab::{ByteVocab, build_vocab_blob};
#[cfg(feature = "webrtc")]
pub use webrtc_native::{
IceServer, LocalTurnServer, NativeWebrtcPeer, WebrtcConfig, WebrtcPipe, loopback_pipes,
spawn_local_turn_server,
};
pub use weights::{Layer, Lfm2Config, ModelKind, Op, Weights, detect_model_kind};
#[cfg(feature = "cli")]
pub fn load_encoder_tokenizer(
dir: &std::path::Path,
max_seq: usize,
) -> Result<tokenizers::Tokenizer> {
let mut tok = tokenizers::Tokenizer::from_file(dir.join("tokenizer.json"))
.map_err(|e| anyhow!("tokenizer: {e}"))?;
tok.with_padding(None);
tok.with_truncation(Some(tokenizers::TruncationParams {
max_length: max_seq,
..Default::default()
}))
.map_err(|e| anyhow!("truncation: {e}"))?;
Ok(tok)
}
pub fn test_kernel_kc4() -> String {
gemv_q4_k_sg_src()
}
pub fn test_kernel_kc16() -> String {
gemv_q4_k_sg16_src()
}
pub struct GpuCtx {
pub device: wgpu::Device,
pub queue: wgpu::Queue,
pub backend: String,
pub subgroups32: bool,
sg32_probe: std::sync::OnceLock<bool>,
pub spin_poll: bool,
pub subgroups: bool,
pub coop_matrix: bool,
pub f16: bool,
pub timestamps: bool,
pub ts_period: f32,
}
impl GpuCtx {
pub fn new() -> Result<Self> {
pollster::block_on(Self::new_async(None))
}
pub fn new_at(adapter_index: usize) -> Result<Self> {
pollster::block_on(Self::new_async(Some(adapter_index)))
}
pub fn share(&self) -> Self {
Self {
device: self.device.clone(),
queue: self.queue.clone(),
backend: self.backend.clone(),
subgroups32: self.subgroups32,
sg32_probe: self.sg32_probe.clone(),
spin_poll: self.spin_poll,
subgroups: self.subgroups,
coop_matrix: self.coop_matrix,
f16: self.f16,
timestamps: self.timestamps,
ts_period: self.ts_period,
}
}
pub async fn new_async(force_idx: Option<usize>) -> Result<Self> {
let instance = wgpu::Instance::default();
let sel_env = std::env::var("OSFKB_WGPU_ADAPTER").ok();
let sel_opt = force_idx.map(|i| i.to_string()).or(sel_env);
let adapter = if let Some(sel) = sel_opt {
let adapters: Vec<wgpu::Adapter> = instance
.enumerate_adapters(wgpu::Backends::all())
.await
.into_iter()
.filter(|a| a.get_info().device_type != wgpu::DeviceType::Cpu)
.collect();
for (i, a) in adapters.iter().enumerate() {
let info = a.get_info();
eprintln!(
"adapter[{i}]: {:?}/{} ({:?})",
info.backend, info.name, info.device_type
);
}
let picked = if let Ok(idx) = sel.parse::<usize>() {
adapters.into_iter().nth(idx)
} else {
let needle = sel.to_lowercase();
adapters
.into_iter()
.find(|a| a.get_info().name.to_lowercase().contains(&needle))
};
picked.ok_or_else(|| anyhow!("OSFKB_WGPU_ADAPTER={sel} matched no adapter"))?
} else {
instance
.request_adapter(&wgpu::RequestAdapterOptions {
power_preference: wgpu::PowerPreference::HighPerformance,
..Default::default()
})
.await
.map_err(|e| anyhow!("no wgpu adapter available on any backend: {e:?}"))?
};
let info = adapter.get_info();
let force_portable = std::env::var("OSFKB_NO_SUBGROUPS").ok().as_deref() == Some("1");
let mut want = (wgpu::Features::SHADER_F16
| wgpu::Features::SUBGROUP
| wgpu::Features::TIMESTAMP_QUERY
| wgpu::Features::EXPERIMENTAL_COOPERATIVE_MATRIX)
& adapter.features();
if force_portable {
want.remove(wgpu::Features::SUBGROUP | wgpu::Features::EXPERIMENTAL_COOPERATIVE_MATRIX);
}
let (device, queue) = adapter
.request_device(&wgpu::DeviceDescriptor {
label: Some("inferencelayer"),
required_features: want,
required_limits: wgpu::Limits {
max_storage_buffer_binding_size: adapter
.limits()
.max_storage_buffer_binding_size,
max_buffer_size: adapter.limits().max_buffer_size,
max_storage_buffers_per_shader_stage: adapter
.limits()
.max_storage_buffers_per_shader_stage,
max_compute_workgroup_storage_size: adapter
.limits()
.max_compute_workgroup_storage_size,
max_compute_invocations_per_workgroup: adapter
.limits()
.max_compute_invocations_per_workgroup,
max_compute_workgroup_size_x: adapter.limits().max_compute_workgroup_size_x,
..wgpu::Limits::default()
},
memory_hints: wgpu::MemoryHints::Performance,
experimental_features: unsafe { wgpu::ExperimentalFeatures::enabled() },
trace: wgpu::Trace::Off,
})
.await
.map_err(|e| anyhow!("request_device: {e:?}"))?;
let ts_period = queue.get_timestamp_period();
Ok(Self {
backend: format!("{:?}/{}", info.backend, info.name),
subgroups32: !force_portable
&& info.subgroup_min_size == 32
&& info.subgroup_max_size == 32,
sg32_probe: std::sync::OnceLock::new(),
spin_poll: std::env::var("OSFKB_SPIN_POLL").ok().as_deref() == Some("1"),
subgroups: want.contains(wgpu::Features::SUBGROUP),
coop_matrix: want.contains(wgpu::Features::EXPERIMENTAL_COOPERATIVE_MATRIX),
f16: want.contains(wgpu::Features::SHADER_F16),
timestamps: want.contains(wgpu::Features::TIMESTAMP_QUERY),
ts_period,
device,
queue,
})
}
pub fn storage(&self, data: &[f32]) -> wgpu::Buffer {
self.storage_bytes(bytemuck::cast_slice(data))
}
pub fn storage_bytes(&self, data: &[u8]) -> wgpu::Buffer {
let buf = self.device.create_buffer(&wgpu::BufferDescriptor {
label: None,
size: data.len() as u64,
usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_SRC
| wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
const CHUNK: usize = 64 << 20;
for (i, chunk) in data.chunks(CHUNK).enumerate() {
self.queue.write_buffer(&buf, (i * CHUNK) as u64, chunk);
if data.len() > CHUNK {
self.queue.submit(std::iter::empty());
let _ = self.device.poll(wgpu::PollType::wait_indefinitely());
}
}
buf
}
pub fn empty(&self, len: usize) -> wgpu::Buffer {
self.device.create_buffer(&wgpu::BufferDescriptor {
label: None,
size: (len * 4) as u64,
usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_SRC
| wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
})
}
pub fn empty_f16(&self, len: usize) -> wgpu::Buffer {
self.device.create_buffer(&wgpu::BufferDescriptor {
label: None,
size: (len * 2) as u64,
usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_SRC
| wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
})
}
pub fn read(&self, buf: &wgpu::Buffer, len: usize) -> Result<Vec<f32>> {
let size = (len * 4) as u64;
let staging = self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("staging"),
size,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let mut enc = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
enc.copy_buffer_to_buffer(buf, 0, &staging, 0, size);
self.queue.submit([enc.finish()]);
let slice = staging.slice(..);
let (tx, rx) = std::sync::mpsc::channel();
slice.map_async(wgpu::MapMode::Read, move |r| {
let _ = tx.send(r);
});
if self.spin_poll {
loop {
let r = self.device.poll(wgpu::PollType::Poll);
if let Ok(status) = &r
&& status.is_queue_empty()
{
break;
}
if rx.try_recv().is_ok() {
let data = slice.get_mapped_range().expect("mapped range");
let out: Vec<f32> = bytemuck::cast_slice(&data).to_vec();
drop(data);
staging.unmap();
return Ok(out);
}
std::hint::spin_loop();
}
} else {
let _ = self.device.poll(wgpu::PollType::wait_indefinitely());
}
rx.recv()
.unwrap()
.map_err(|e| anyhow!("map_async: {e:?}"))?;
let data = slice.get_mapped_range().expect("mapped range");
let out: Vec<f32> = bytemuck::cast_slice(&data).to_vec();
drop(data);
staging.unmap();
Ok(out)
}
pub fn read_u32(&self, buf: &wgpu::Buffer, len: usize) -> Result<Vec<u32>> {
let f = self.read(buf, len)?;
Ok(f.iter().map(|x| x.to_bits()).collect())
}
pub fn read_f16(&self, buf: &wgpu::Buffer, len: usize) -> Result<Vec<f32>> {
let words = self.read_u32(buf, len.div_ceil(2))?;
let mut out = Vec::with_capacity(len);
for (i, w) in words.iter().enumerate() {
out.push(half::f16::from_bits(*w as u16).to_f32());
if 2 * i + 1 < len {
out.push(half::f16::from_bits((*w >> 16) as u16).to_f32());
}
}
out.truncate(len);
Ok(out)
}
pub fn subgroups32_effective(&self) -> bool {
if !self.subgroups {
return false;
}
if self.subgroups32 {
return true;
}
*self.sg32_probe.get_or_init(|| {
if std::env::var("OSFKB_SG32_PROBE").ok().as_deref() == Some("0") {
return false;
}
let ok = self.probe_sg32_lcpp().unwrap_or(false);
eprintln!(
"lcpp subgroup probe on {}: {}",
self.backend,
if ok {
"WG=32 subgroupAdd semantics VALIDATED — fast GEMV/head family eligible"
} else {
"validation failed — tree family retained"
}
);
ok
})
}
fn probe_sg32_lcpp(&self) -> Result<bool> {
const M: usize = 8; const N: usize = 1536; const NCOLS: usize = 2;
let nblk = N / 32;
let scales: Vec<u16> = vec![0x3C00; M * nblk];
let quants: Vec<u32> = (0..M * nblk * 4)
.map(|i| (i as u32).wrapping_mul(2654435761).rotate_left(7))
.collect();
let x: Vec<f32> = (0..NCOLS * N).map(|i| ((i % 8) as f32) - 4.0).collect();
let mut expect = vec![0f32; NCOLS * M];
for c in 0..NCOLS {
for row in 0..M {
let mut sum = 0f64;
for b in 0..nblk {
let q = &quants[(row * nblk + b) * 4..(row * nblk + b) * 4 + 4];
let mut s = 0f64;
for (w, word) in q.iter().enumerate() {
for j in 0..4 {
let byte = (word >> (8 * j)) & 0xFF;
let lo = (byte & 0xF) as f64 - 8.0;
let hi = (byte >> 4) as f64 - 8.0;
s += lo * x[c * N + b * 32 + 4 * w + j] as f64;
s += hi * x[c * N + b * 32 + 16 + 4 * w + j] as f64;
}
}
sum += s; }
expect[c * M + row] = sum as f32;
}
}
let scope = self.device.push_error_scope(wgpu::ErrorFilter::Validation);
let module = self
.device
.shader_module_tuned(wgpu::ShaderModuleDescriptor {
label: Some("sg32_probe"),
source: wgpu::ShaderSource::Wgsl(gemv_q4_k_lcpp_src().into()),
});
let pl = self
.device
.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some("sg32_probe"),
layout: None,
module: &module,
entry_point: Some("main"),
compilation_options: wgpu::PipelineCompilationOptions::default(),
cache: None,
});
if pollster::block_on(scope.pop()).is_some() {
return Ok(false);
}
let sb = self.storage_bytes(bytemuck::cast_slice(&scales));
let qb = self.storage_bytes(bytemuck::cast_slice(&quants));
let xb = self.storage(&x);
let yb = self.storage(&[0f32; NCOLS * M]);
let dims = wgpu::util::DeviceExt::create_buffer_init(
&self.device,
&wgpu::util::BufferInitDescriptor {
label: Some("sg32_probe_dims"),
contents: bytemuck::cast_slice(&[M as u32, N as u32, 0, NCOLS as u32]),
usage: wgpu::BufferUsages::UNIFORM,
},
);
let entries: Vec<wgpu::BindGroupEntry> = [&sb, &qb, &xb, &yb, &dims]
.iter()
.enumerate()
.map(|(i, b)| wgpu::BindGroupEntry {
binding: i as u32,
resource: b.as_entire_binding(),
})
.collect();
let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("sg32_probe"),
layout: &pl.get_bind_group_layout(0),
entries: &entries,
});
let mut enc = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pl);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups((M as u32).div_ceil(4), NCOLS as u32, 1);
}
self.queue.submit([enc.finish()]);
let got = self.read(&yb, NCOLS * M)?;
Ok(got == expect)
}
}
pub(crate) const Q4_FN: &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 fma(vec4<f32>(unpack4xU8(word & 0xF0F0F0F0u)), vec4<f32>(0.0625), vec4<f32>(-8.0)); }
"#;
pub fn gemm_q1_src() -> String {
gemm_q1_src_impl(1)
}
pub fn gemm_q1_src_chunk(chunk: u32) -> String {
gemm_q1_src_impl(chunk)
}
fn gemm_q1_src_impl(chunk: u32) -> String {
assert!(
matches!(chunk, 1 | 2 | 4),
"chunk must divide a 128-superblock"
);
let ch = chunk as usize;
let tile = 128 * ch;
let lanes = 4 * ch; let mut b = format!(
r##"fn q1s(word: u32, sh: u32) -> vec4<f32> {{
let bits = (vec4<u32>(word) >> vec4<u32>(sh, sh + 1u, sh + 2u, sh + 3u)) & vec4<u32>(1u);
return select(vec4<f32>(-1.0), vec4<f32>(1.0), bits == vec4<u32>(1u));
}}
@group(0) @binding(0) var<storage, read> scales: array<f16>; // [m, n/128]
@group(0) @binding(1) var<storage, read> quants: array<u32>; // [m, n/32] sign words
@group(0) @binding(2) var<storage, read> x: array<vec4<f32>>; // [ncols, N/4]
@group(0) @binding(3) var<storage, read_write> y: array<f32>; // [ncols, M]
@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)
const BM: u32 = 32u;
const BN: u32 = 16u;
var<workgroup> xsv: array<vec4<f32>, {tile}>;
var<workgroup> xtmp: array<vec4<f32>, {tile}>;
var<workgroup> wsc: array<f32, 32>;
var<workgroup> wq1: array<u32, {wq}>;
@compute @workgroup_size(32)
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 nsc = n / 128u;
let xstride = n / 4u;
let ty = lid / 4u;
let tx = lid % 4u;
let row0 = wid.x * BM + ty * 4u;
let col0 = wid.y * BN;
var acc0 = vec4<f32>(0.0);
var acc1 = vec4<f32>(0.0);
var acc2 = vec4<f32>(0.0);
var acc3 = vec4<f32>(0.0);
let nrounds = nblk / {chunk}u;
for (var rd = 0u; rd < nrounds; rd = rd + 1u) {{
let b0 = rd * {chunk}u;
for (var j = 0u; j < {lanes}u; j = j + 1u) {{
let idx = lid * {lanes}u + j;
let cn = idx / {w8}u;
let e = idx % {w8}u;
var v = vec4<f32>(0.0);
if (col0 + cn < ncols) {{ v = x[(col0 + cn) * xstride + b0 * 8u + e]; }}
xtmp[idx] = v;
}}
workgroupBarrier();
for (var j = 0u; j < {lanes}u; j = j + 1u) {{
let idx = lid * {lanes}u + j;
let kk = idx / 4u;
let c4 = idx % 4u;
let e = kk / 4u;
let comp = kk % 4u;
xsv[idx] = vec4<f32>(
xtmp[(c4 * 4u) * {w8}u + e][comp],
xtmp[(c4 * 4u + 1u) * {w8}u + e][comp],
xtmp[(c4 * 4u + 2u) * {w8}u + e][comp],
xtmp[(c4 * 4u + 3u) * {w8}u + e][comp],
);
}}
{{
let row = wid.x * BM + lid;
if (row < m) {{
wsc[lid] = f32(scales[row * nsc + (b0 >> 2u)]);
"##,
tile = tile,
chunk = chunk,
lanes = lanes,
w8 = 8 * ch,
wq = 32 * ch,
);
for c in 0..ch {
b.push_str(&format!(
" wq1[lid * {ch}u + {c}u] = quants[row * nblk + b0 + {c}u];\n"
));
}
b.push_str(
r##" } else {
wsc[lid] = 0.0;
"##,
);
for c in 0..ch {
b.push_str(&format!(" wq1[lid * {ch}u + {c}u] = 0u;\n"));
}
b.push_str(
r##" }
}
workgroupBarrier();
{
"##,
);
for r in 0..4 {
b.push_str(&format!(
" let d{r} = wsc[ty * 4u + {r}u];\n var s{r} = vec4<f32>(0.0);\n"
));
for c in 0..ch {
b.push_str(&format!(
" let w{r}_{c} = wq1[(ty * 4u + {r}u) * {ch}u + {c}u];\n"
));
for g in 0..8 {
let sh = 4 * g;
let e0 = (c * 32 + sh) * 4;
let (i0, i1, i2, i3) = (e0, e0 + 4, e0 + 8, e0 + 12);
b.push_str(&format!(
" let g{r}_{c}_{g} = q1s(w{r}_{c}, {sh}u);\n s{r} = s{r} + g{r}_{c}_{g}.x * xsv[{i0}u + tx] + g{r}_{c}_{g}.y * xsv[{i1}u + tx] + g{r}_{c}_{g}.z * xsv[{i2}u + tx] + g{r}_{c}_{g}.w * xsv[{i3}u + tx];\n"
));
}
}
b.push_str(&format!(" acc{r} = acc{r} + d{r} * s{r};\n"));
}
b.push_str(
r##" }
workgroupBarrier();
}
"##,
);
for r in 0..4 {
b.push_str(&format!(
r##" {{
let row = row0 + {r}u;
if (row < m) {{
for (var c = 0u; c < 4u; c = c + 1u) {{
let col = col0 + tx * 4u + c;
if (col < ncols) {{
let o = col * m + row;
var v = acc{r}.x;
if (c == 1u) {{ v = acc{r}.y; }}
if (c == 2u) {{ v = acc{r}.z; }}
if (c == 3u) {{ v = acc{r}.w; }}
if (dims.z == 1u) {{ y[o] = y[o] + v; }} else {{ y[o] = v; }}
}}
}}
}}
}}
"##
));
}
b.push_str("}\n");
format!("enable f16;\n{b}")
}
pub fn gemm_q1_xt_src() -> String {
let mut src = String::from(
r##"enable f16;
fn q1s(word: u32, sh: u32) -> vec4<f32> {
let bits = (vec4<u32>(word) >> vec4<u32>(sh, sh + 1u, sh + 2u, sh + 3u)) & vec4<u32>(1u);
return select(vec4<f32>(-1.0), vec4<f32>(1.0), bits == vec4<u32>(1u));
}
@group(0) @binding(0) var<storage, read> scales: array<f16>;
@group(0) @binding(1) var<storage, read> quants: array<u32>;
@group(0) @binding(2) var<storage, read> x: array<vec4<f32>>; // [n][ncols/4] TRANSPOSED
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)
const BM: u32 = 32u;
const BN: u32 = 16u;
var<workgroup> xsv: array<vec4<f32>, 128>;
var<workgroup> wsc: array<f32, 32>;
var<workgroup> wq1: array<u32, 32>;
@compute @workgroup_size(32)
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 nsc = n / 128u;
let cstride = (ncols + 3u) / 4u;
let ty = lid / 4u; let tx = lid % 4u;
let row0 = wid.x * BM + ty * 4u;
let col0 = wid.y * BN;
let cv0 = col0 / 4u;
var acc0 = vec4<f32>(0.0); var acc1 = vec4<f32>(0.0);
var acc2 = vec4<f32>(0.0); var acc3 = vec4<f32>(0.0);
for (var b = 0u; b < nblk; b = b + 1u) {
for (var j = 0u; j < 4u; j = j + 1u) {
let idx = lid * 4u + j;
let e = idx / 4u; let c4 = idx % 4u;
xsv[idx] = x[(b * 32u + e) * cstride + cv0 + c4];
}
{
let row = wid.x * BM + lid;
if (row < m) {
wsc[lid] = f32(scales[row * nsc + (b >> 2u)]);
wq1[lid] = quants[row * nblk + b];
} else { wsc[lid] = 0.0; wq1[lid] = 0u; }
}
workgroupBarrier();
{
"##,
);
for r in 0..4 {
src.push_str(&format!(
" let d{r} = wsc[ty * 4u + {r}u];\n let w{r} = wq1[ty * 4u + {r}u];\n"
));
for g in 0..8 {
let sh = 4 * g;
let (i0, i1, i2, i3) = (sh * 4, (sh + 1) * 4, (sh + 2) * 4, (sh + 3) * 4);
src.push_str(&format!(
" let g{r}_{g} = d{r} * q1s(w{r}, {sh}u);\n acc{r} = acc{r} + g{r}_{g}.x * xsv[{i0}u + tx] + g{r}_{g}.y * xsv[{i1}u + tx] + g{r}_{g}.z * xsv[{i2}u + tx] + g{r}_{g}.w * xsv[{i3}u + tx];\n"
));
}
}
src.push_str(" }\n workgroupBarrier();\n }\n");
for r in 0..4 {
src.push_str(&format!(
r##" {{
let row = row0 + {r}u;
if (row < m) {{
for (var c = 0u; c < 4u; c = c + 1u) {{
let col = col0 + tx * 4u + c;
if (col < ncols) {{
let o = col * m + row;
var v = acc{r}.x;
if (c == 1u) {{ v = acc{r}.y; }}
if (c == 2u) {{ v = acc{r}.z; }}
if (c == 3u) {{ v = acc{r}.w; }}
if (dims.z == 1u) {{ y[o] = y[o] + v; }} else {{ y[o] = v; }}
}}
}}
}}
}}
"##
));
}
src.push_str("}\n");
src
}
pub fn transpose_cols_src() -> String {
r##"@group(0) @binding(0) var<storage, read> x: array<f32>; // [ncols, n]
@group(0) @binding(1) var<storage, read_write> xt: array<vec4<f32>>; // [n, ncols/4]
@group(0) @binding(2) var<uniform> d: vec4<u32>; // (n, ncols, _, _)
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
let n = d.x; let ncols = d.y;
let cstride = (ncols + 3u) / 4u;
let tid = gid.x;
if (tid >= n * cstride) { return; }
let c4 = tid / n;
let e = tid % n;
// No dynamic component writes (Naga lowers them to scratch): unrolled.
let c0 = c4 * 4u;
var v = vec4<f32>(0.0);
if (c0 < ncols) { v.x = x[c0 * n + e]; }
if (c0 + 1u < ncols) { v.y = x[(c0 + 1u) * n + e]; }
if (c0 + 2u < ncols) { v.z = x[(c0 + 2u) * n + e]; }
if (c0 + 3u < ncols) { v.w = x[(c0 + 3u) * n + e]; }
xt[e * cstride + c4] = v;
}
"##
.to_string()
}
pub fn pack_f16_src() -> String {
r##"@group(0) @binding(0) var<storage, read> x: array<f32>;
@group(0) @binding(1) var<storage, read_write> y: array<u32>;
@group(0) @binding(2) var<uniform> d: vec4<u32>; // (n_f32, _, _, _)
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
let i = gid.x;
if (i * 2u < d.x) {
y[i] = pack2x16float(vec2<f32>(x[i * 2u], x[i * 2u + 1u]));
}
}
"##
.to_string()
}
pub fn gemv_q1_f16x_src() -> String {
r#"enable f16;
fn q1s(word: u32, sh: u32) -> vec4<f32> {
let bits = (vec4<u32>(word) >> vec4<u32>(sh, sh + 1u, sh + 2u, sh + 3u)) & vec4<u32>(1u);
return select(vec4<f32>(-1.0), vec4<f32>(1.0), bits == vec4<u32>(1u));
}
@group(0) @binding(0) var<storage, read> scales: array<f16>; // [m, n/128]
@group(0) @binding(1) var<storage, read> bits: array<u32>; // [m, n/32] sign words
@group(0) @binding(2) var<storage, read> x: array<vec4<u32>>; // packed f16 pairs
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)
@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 nsc = n / 128u;
let xstride = n / 8u;
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 * 4u;
let q0 = x[xb]; let q1 = x[xb + 1u];
let q2 = x[xb + 2u]; let q3 = x[xb + 3u];
let v0 = vec4<f32>(unpack2x16float(q0.x), unpack2x16float(q0.y));
let v1 = vec4<f32>(unpack2x16float(q0.z), unpack2x16float(q0.w));
let v2 = vec4<f32>(unpack2x16float(q1.x), unpack2x16float(q1.y));
let v3 = vec4<f32>(unpack2x16float(q1.z), unpack2x16float(q1.w));
let v4 = vec4<f32>(unpack2x16float(q2.x), unpack2x16float(q2.y));
let v5 = vec4<f32>(unpack2x16float(q2.z), unpack2x16float(q2.w));
let v6 = vec4<f32>(unpack2x16float(q3.x), unpack2x16float(q3.y));
let v7 = vec4<f32>(unpack2x16float(q3.z), unpack2x16float(q3.w));
for (var r = 0u; r < 4u; r = r + 1u) {
let row = min(row0 + r, m - 1u);
let d = f32(scales[row * nsc + (b >> 2u)]);
let w = bits[row * nblk + b];
var s = dot(q1s(w, 0u), v0) + dot(q1s(w, 4u), v1);
s = s + dot(q1s(w, 8u), v2) + dot(q1s(w, 12u), v3);
s = s + dot(q1s(w, 16u), v4) + dot(q1s(w, 20u), v5);
s = s + dot(q1s(w, 24u), v6) + dot(q1s(w, 28u), 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) {
let yo = col * m + row0 + r;
if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }
}
}
}
}
"#
.to_string()
}
pub fn q1_rp4_repack_src() -> String {
r##"enable f16;
@group(0) @binding(0) var<storage, read> scales: array<f16>; // [m, n/128]
@group(0) @binding(1) var<storage, read> bits: array<u32>; // [m, n/32]
@group(0) @binding(2) var<storage, read_write> scales4: array<vec4<f32>>; // [m/4, n/128]
@group(0) @binding(3) var<storage, read_write> bits4: array<vec4<u32>>; // [m/4, n/32]
@group(0) @binding(4) var<uniform> d: vec4<u32>; // (m, nblk, nsc, _)
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
let m = d.x; let nblk = d.y; let nsc = d.z;
let g4 = (m + 3u) / 4u;
let t = gid.x;
if (t < g4 * nblk) {
let g = t / nblk; let b = t % nblk;
var w = vec4<u32>(0u);
for (var r = 0u; r < 4u; r = r + 1u) {
let row = g * 4u + r;
if (row < m) { w[r] = bits[row * nblk + b]; }
}
bits4[t] = w;
}
if (t < g4 * nsc) {
let g = t / nsc; let j = t % nsc;
var sc = vec4<f32>(0.0);
for (var r = 0u; r < 4u; r = r + 1u) {
let row = g * 4u + r;
if (row < m) { sc[r] = f32(scales[row * nsc + j]); }
}
scales4[t] = sc;
}
}
"##
.to_string()
}
pub fn gemv_q1_rp4_f16x_src() -> String {
r#"enable f16;
@group(0) @binding(0) var<storage, read> scales4: array<vec4<f32>>; // [m/4, n/128]
@group(0) @binding(1) var<storage, read> bits4: array<vec4<u32>>; // [m/4, n/32]
@group(0) @binding(2) var<storage, read> x: array<vec4<u32>>; // packed f16 pairs
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
@group(0) @binding(4) var<storage, read_write> y16: array<u32>; // packed-f16 twin of y
@group(0) @binding(5) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)
@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 ncols = dims.w;
let nblk = n / 32u;
let nsc = n / 128u;
let xstride = n / 8u;
let g = wid.x;
let col = wid.y;
let xoff = col * xstride;
var accx = 0.0; var accy = 0.0; var accz = 0.0; var accw = 0.0;
var b = sid;
let xb0 = xoff + b * 4u;
var pp0 = x[xb0]; var pp1 = x[xb0 + 1u];
var pp2 = x[xb0 + 2u]; var pp3 = x[xb0 + 3u];
var wc = bits4[g * nblk + b];
var dc = scales4[g * nsc + (b >> 2u)];
loop {
let bn = b + 32u;
let has_next = bn < nblk;
var np0 = vec4<u32>(0u); var np1 = vec4<u32>(0u);
var np2 = vec4<u32>(0u); var np3 = vec4<u32>(0u);
var nw = vec4<u32>(0u); var nd = vec4<f32>(0.0);
if (has_next) {
let nxb = xoff + bn * 4u;
np0 = x[nxb]; np1 = x[nxb + 1u];
np2 = x[nxb + 2u]; np3 = x[nxb + 3u];
nw = bits4[g * nblk + bn];
nd = scales4[g * nsc + (bn >> 2u)];
}
let c0 = vec4<f32>(unpack2x16float(pp0.x), unpack2x16float(pp0.y));
let c1 = vec4<f32>(unpack2x16float(pp0.z), unpack2x16float(pp0.w));
let c2 = vec4<f32>(unpack2x16float(pp1.x), unpack2x16float(pp1.y));
let c3 = vec4<f32>(unpack2x16float(pp1.z), unpack2x16float(pp1.w));
let c4 = vec4<f32>(unpack2x16float(pp2.x), unpack2x16float(pp2.y));
let c5 = vec4<f32>(unpack2x16float(pp2.z), unpack2x16float(pp2.w));
let c6 = vec4<f32>(unpack2x16float(pp3.x), unpack2x16float(pp3.y));
let c7 = vec4<f32>(unpack2x16float(pp3.z), unpack2x16float(pp3.w));
{
let w = wc.x;
let s0 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(0u, 1u, 2u, 3u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s1 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(4u, 5u, 6u, 7u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s2 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(8u, 9u, 10u, 11u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s3 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(12u, 13u, 14u, 15u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s4 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(16u, 17u, 18u, 19u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s5 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(20u, 21u, 22u, 23u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s6 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(24u, 25u, 26u, 27u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s7 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(28u, 29u, 30u, 31u)) & vec4<u32>(1u)) == vec4<u32>(1u));
var t = dot(s0, c0) + dot(s1, c1) + dot(s2, c2) + dot(s3, c3);
t = t + dot(s4, c4) + dot(s5, c5) + dot(s6, c6) + dot(s7, c7);
accx = accx + dc.x * t;
}
{
let w = wc.y;
let s0 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(0u, 1u, 2u, 3u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s1 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(4u, 5u, 6u, 7u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s2 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(8u, 9u, 10u, 11u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s3 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(12u, 13u, 14u, 15u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s4 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(16u, 17u, 18u, 19u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s5 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(20u, 21u, 22u, 23u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s6 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(24u, 25u, 26u, 27u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s7 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(28u, 29u, 30u, 31u)) & vec4<u32>(1u)) == vec4<u32>(1u));
var t = dot(s0, c0) + dot(s1, c1) + dot(s2, c2) + dot(s3, c3);
t = t + dot(s4, c4) + dot(s5, c5) + dot(s6, c6) + dot(s7, c7);
accy = accy + dc.y * t;
}
{
let w = wc.z;
let s0 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(0u, 1u, 2u, 3u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s1 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(4u, 5u, 6u, 7u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s2 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(8u, 9u, 10u, 11u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s3 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(12u, 13u, 14u, 15u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s4 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(16u, 17u, 18u, 19u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s5 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(20u, 21u, 22u, 23u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s6 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(24u, 25u, 26u, 27u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s7 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(28u, 29u, 30u, 31u)) & vec4<u32>(1u)) == vec4<u32>(1u));
var t = dot(s0, c0) + dot(s1, c1) + dot(s2, c2) + dot(s3, c3);
t = t + dot(s4, c4) + dot(s5, c5) + dot(s6, c6) + dot(s7, c7);
accz = accz + dc.z * t;
}
{
let w = wc.w;
let s0 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(0u, 1u, 2u, 3u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s1 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(4u, 5u, 6u, 7u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s2 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(8u, 9u, 10u, 11u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s3 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(12u, 13u, 14u, 15u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s4 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(16u, 17u, 18u, 19u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s5 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(20u, 21u, 22u, 23u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s6 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(24u, 25u, 26u, 27u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s7 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(28u, 29u, 30u, 31u)) & vec4<u32>(1u)) == vec4<u32>(1u));
var t = dot(s0, c0) + dot(s1, c1) + dot(s2, c2) + dot(s3, c3);
t = t + dot(s4, c4) + dot(s5, c5) + dot(s6, c6) + dot(s7, c7);
accw = accw + dc.w * t;
}
if (!has_next) { break; }
b = bn;
pp0 = np0; pp1 = np1; pp2 = np2; pp3 = np3;
wc = nw; dc = nd;
}
let tx = subgroupAdd(accx); let ty = subgroupAdd(accy);
let tz = subgroupAdd(accz); let tw = subgroupAdd(accw);
if (sid == 0u) {
let r0 = g * 4u;
var f0 = tx; var f1 = ty; var f2 = tz; var f3 = tw;
if (r0 < m && col < ncols) { let o = col * m + r0; if (dims.z == 1u) { f0 = y[o] + tx; } y[o] = f0; }
if (r0 + 1u < m && col < ncols) { let o = col * m + r0 + 1u; if (dims.z == 1u) { f1 = y[o] + ty; } y[o] = f1; }
if (r0 + 2u < m && col < ncols) { let o = col * m + r0 + 2u; if (dims.z == 1u) { f2 = y[o] + tz; } y[o] = f2; }
if (r0 + 3u < m && col < ncols) { let o = col * m + r0 + 3u; if (dims.z == 1u) { f3 = y[o] + tw; } y[o] = f3; }
// packed-f16 twin (rows are 4-aligned per group; m is even at every site)
if (r0 + 1u < m && col < ncols) { y16[(col * m + r0) / 2u] = pack2x16float(vec2<f32>(f0, f1)); }
if (r0 + 3u < m && col < ncols) { y16[(col * m + r0 + 2u) / 2u] = pack2x16float(vec2<f32>(f2, f3)); }
}
}
"#
.to_string()
}
pub fn mlp_gate_q1_rp4_f16x_src() -> String {
r#"enable f16;
@group(0) @binding(0) var<storage, read> s14: array<vec4<f32>>; // gate scales rp4
@group(0) @binding(1) var<storage, read> b14: array<vec4<u32>>; // gate signs rp4
@group(0) @binding(2) var<storage, read> s34: array<vec4<f32>>; // up scales rp4
@group(0) @binding(3) var<storage, read> b34: array<vec4<u32>>; // up signs rp4
@group(0) @binding(4) var<storage, read> x: array<vec4<u32>>; // [ncols, N/8] PACKED f16
@group(0) @binding(5) var<storage, read_write> y: array<f32>; // [ncols, M]
@group(0) @binding(6) var<storage, read_write> y16: array<u32>; // packed-f16 twin of y
@group(0) @binding(7) var<uniform> dims: vec4<u32>; // (M, N, _, ncols)
@group(0) @binding(8) var<uniform> epsm: vec4<f32>; // (_, gelu-flag, _, _)
@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 ncols = dims.w;
let nblk = n / 32u;
let nsc = n / 128u;
let xstride = n / 8u;
let g = wid.x;
let col = wid.y;
let xoff = col * xstride;
var agx = 0.0; var agy = 0.0; var agz = 0.0; var agw = 0.0;
var aux = 0.0; var auy = 0.0; var auz = 0.0; var auw = 0.0;
var b = sid;
let xb0 = xoff + b * 4u;
var pp0 = x[xb0]; var pp1 = x[xb0 + 1u];
var pp2 = x[xb0 + 2u]; var pp3 = x[xb0 + 3u];
var wc1 = b14[g * nblk + b]; var dc1 = s14[g * nsc + (b >> 2u)];
var wc3 = b34[g * nblk + b]; var dc3 = s34[g * nsc + (b >> 2u)];
loop {
let bn = b + 32u;
let has_next = bn < nblk;
var np0 = vec4<u32>(0u); var np1 = vec4<u32>(0u);
var np2 = vec4<u32>(0u); var np3 = vec4<u32>(0u);
var nw1 = vec4<u32>(0u); var nd1 = vec4<f32>(0.0);
var nw3 = vec4<u32>(0u); var nd3 = vec4<f32>(0.0);
if (has_next) {
let nxb = xoff + bn * 4u;
np0 = x[nxb]; np1 = x[nxb + 1u];
np2 = x[nxb + 2u]; np3 = x[nxb + 3u];
nw1 = b14[g * nblk + bn]; nd1 = s14[g * nsc + (bn >> 2u)];
nw3 = b34[g * nblk + bn]; nd3 = s34[g * nsc + (bn >> 2u)];
}
let c0 = vec4<f32>(unpack2x16float(pp0.x), unpack2x16float(pp0.y));
let c1 = vec4<f32>(unpack2x16float(pp0.z), unpack2x16float(pp0.w));
let c2 = vec4<f32>(unpack2x16float(pp1.x), unpack2x16float(pp1.y));
let c3 = vec4<f32>(unpack2x16float(pp1.z), unpack2x16float(pp1.w));
let c4 = vec4<f32>(unpack2x16float(pp2.x), unpack2x16float(pp2.y));
let c5 = vec4<f32>(unpack2x16float(pp2.z), unpack2x16float(pp2.w));
let c6 = vec4<f32>(unpack2x16float(pp3.x), unpack2x16float(pp3.y));
let c7 = vec4<f32>(unpack2x16float(pp3.z), unpack2x16float(pp3.w));
{
let w = wc1.x;
let s0 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(0u,1u,2u,3u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s1 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(4u,5u,6u,7u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s2 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(8u,9u,10u,11u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s3 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(12u,13u,14u,15u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s4 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(16u,17u,18u,19u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s5 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(20u,21u,22u,23u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s6 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(24u,25u,26u,27u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s7 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(28u,29u,30u,31u)) & vec4<u32>(1u)) == vec4<u32>(1u));
var t = dot(s0, c0) + dot(s1, c1) + dot(s2, c2) + dot(s3, c3);
t = t + dot(s4, c4) + dot(s5, c5) + dot(s6, c6) + dot(s7, c7);
agx = agx + dc1.x * t;
} {
let w = wc3.x;
let s0 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(0u,1u,2u,3u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s1 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(4u,5u,6u,7u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s2 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(8u,9u,10u,11u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s3 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(12u,13u,14u,15u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s4 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(16u,17u,18u,19u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s5 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(20u,21u,22u,23u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s6 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(24u,25u,26u,27u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s7 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(28u,29u,30u,31u)) & vec4<u32>(1u)) == vec4<u32>(1u));
var t = dot(s0, c0) + dot(s1, c1) + dot(s2, c2) + dot(s3, c3);
t = t + dot(s4, c4) + dot(s5, c5) + dot(s6, c6) + dot(s7, c7);
aux = aux + dc3.x * t;
}
{
let w = wc1.y;
let s0 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(0u,1u,2u,3u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s1 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(4u,5u,6u,7u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s2 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(8u,9u,10u,11u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s3 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(12u,13u,14u,15u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s4 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(16u,17u,18u,19u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s5 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(20u,21u,22u,23u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s6 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(24u,25u,26u,27u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s7 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(28u,29u,30u,31u)) & vec4<u32>(1u)) == vec4<u32>(1u));
var t = dot(s0, c0) + dot(s1, c1) + dot(s2, c2) + dot(s3, c3);
t = t + dot(s4, c4) + dot(s5, c5) + dot(s6, c6) + dot(s7, c7);
agy = agy + dc1.y * t;
} {
let w = wc3.y;
let s0 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(0u,1u,2u,3u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s1 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(4u,5u,6u,7u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s2 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(8u,9u,10u,11u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s3 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(12u,13u,14u,15u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s4 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(16u,17u,18u,19u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s5 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(20u,21u,22u,23u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s6 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(24u,25u,26u,27u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s7 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(28u,29u,30u,31u)) & vec4<u32>(1u)) == vec4<u32>(1u));
var t = dot(s0, c0) + dot(s1, c1) + dot(s2, c2) + dot(s3, c3);
t = t + dot(s4, c4) + dot(s5, c5) + dot(s6, c6) + dot(s7, c7);
auy = auy + dc3.y * t;
}
{
let w = wc1.z;
let s0 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(0u,1u,2u,3u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s1 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(4u,5u,6u,7u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s2 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(8u,9u,10u,11u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s3 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(12u,13u,14u,15u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s4 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(16u,17u,18u,19u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s5 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(20u,21u,22u,23u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s6 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(24u,25u,26u,27u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s7 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(28u,29u,30u,31u)) & vec4<u32>(1u)) == vec4<u32>(1u));
var t = dot(s0, c0) + dot(s1, c1) + dot(s2, c2) + dot(s3, c3);
t = t + dot(s4, c4) + dot(s5, c5) + dot(s6, c6) + dot(s7, c7);
agz = agz + dc1.z * t;
} {
let w = wc3.z;
let s0 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(0u,1u,2u,3u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s1 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(4u,5u,6u,7u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s2 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(8u,9u,10u,11u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s3 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(12u,13u,14u,15u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s4 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(16u,17u,18u,19u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s5 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(20u,21u,22u,23u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s6 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(24u,25u,26u,27u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s7 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(28u,29u,30u,31u)) & vec4<u32>(1u)) == vec4<u32>(1u));
var t = dot(s0, c0) + dot(s1, c1) + dot(s2, c2) + dot(s3, c3);
t = t + dot(s4, c4) + dot(s5, c5) + dot(s6, c6) + dot(s7, c7);
auz = auz + dc3.z * t;
}
{
let w = wc1.w;
let s0 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(0u,1u,2u,3u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s1 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(4u,5u,6u,7u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s2 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(8u,9u,10u,11u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s3 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(12u,13u,14u,15u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s4 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(16u,17u,18u,19u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s5 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(20u,21u,22u,23u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s6 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(24u,25u,26u,27u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s7 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(28u,29u,30u,31u)) & vec4<u32>(1u)) == vec4<u32>(1u));
var t = dot(s0, c0) + dot(s1, c1) + dot(s2, c2) + dot(s3, c3);
t = t + dot(s4, c4) + dot(s5, c5) + dot(s6, c6) + dot(s7, c7);
agw = agw + dc1.w * t;
} {
let w = wc3.w;
let s0 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(0u,1u,2u,3u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s1 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(4u,5u,6u,7u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s2 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(8u,9u,10u,11u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s3 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(12u,13u,14u,15u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s4 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(16u,17u,18u,19u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s5 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(20u,21u,22u,23u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s6 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(24u,25u,26u,27u)) & vec4<u32>(1u)) == vec4<u32>(1u));
let s7 = select(vec4<f32>(-1.0), vec4<f32>(1.0), ((vec4<u32>(w) >> vec4<u32>(28u,29u,30u,31u)) & vec4<u32>(1u)) == vec4<u32>(1u));
var t = dot(s0, c0) + dot(s1, c1) + dot(s2, c2) + dot(s3, c3);
t = t + dot(s4, c4) + dot(s5, c5) + dot(s6, c6) + dot(s7, c7);
auw = auw + dc3.w * t;
}
if (!has_next) { break; }
b = bn;
pp0 = np0; pp1 = np1; pp2 = np2; pp3 = np3;
wc1 = nw1; dc1 = nd1; wc3 = nw3; dc3 = nd3;
}
let tgx = subgroupAdd(agx); let tgy = subgroupAdd(agy); let tgz = subgroupAdd(agz); let tgw = subgroupAdd(agw);
let tux = subgroupAdd(aux); let tuy = subgroupAdd(auy); let tuz = subgroupAdd(auz); let tuw = subgroupAdd(auw);
if (sid == 0u) {
let gv = array<f32, 4>(tgx, tgy, tgz, tgw);
let uv = array<f32, 4>(tux, tuy, tuz, tuw);
var fo = array<f32, 4>(0.0, 0.0, 0.0, 0.0);
for (var r = 0u; r < 4u; r = r + 1u) {
if (g * 4u + r < m && col < ncols) {
let gate = gv[r];
let upv = uv[r];
var act: f32;
if (epsm.y != 0.0) {
let g3 = gate * gate * gate;
let targ = clamp(0.7978845608028654 * (gate + 0.044715 * g3), -20.0, 20.0);
act = 0.5 * gate * (1.0 + tanh(targ));
} else {
act = gate / (1.0 + exp(-gate));
}
let v = act * upv;
y[col * m + g * 4u + r] = v;
fo[r] = v;
}
}
if (g * 4u + 1u < m && col < ncols) { y16[(col * m + g * 4u) / 2u] = pack2x16float(vec2<f32>(fo[0], fo[1])); }
if (g * 4u + 3u < m && col < ncols) { y16[(col * m + g * 4u + 2u) / 2u] = pack2x16float(vec2<f32>(fo[2], fo[3])); }
}
}
"#
.to_string()
}
pub fn gemm_q1_xt_rp4_src() -> String {
let src = gemm_q1_xt_src();
let out = src
.replace(
"@group(0) @binding(0) var<storage, read> scales: array<f16>;",
"@group(0) @binding(0) var<storage, read> scales4: array<vec4<f32>>;",
)
.replace(
"@group(0) @binding(1) var<storage, read> quants: array<u32>;",
"@group(0) @binding(1) var<storage, read> quants4: array<vec4<u32>>;",
)
.replace(
r#" {
let row = wid.x * BM + lid;
if (row < m) {
wsc[lid] = f32(scales[row * nsc + (b >> 2u)]);
wq1[lid] = quants[row * nblk + b];
} else { wsc[lid] = 0.0; wq1[lid] = 0u; }
}"#,
r#" if (lid < 8u) {
let g = wid.x * 8u + lid;
var w4 = vec4<u32>(0u);
var d4 = vec4<f32>(0.0);
if (g < (m + 3u) / 4u) {
w4 = quants4[g * nblk + b];
d4 = scales4[g * nsc + (b >> 2u)];
}
wq1[lid * 4u] = w4.x;
wq1[lid * 4u + 1u] = w4.y;
wq1[lid * 4u + 2u] = w4.z;
wq1[lid * 4u + 3u] = w4.w;
wsc[lid * 4u] = d4.x;
wsc[lid * 4u + 1u] = d4.y;
wsc[lid * 4u + 2u] = d4.z;
wsc[lid * 4u + 3u] = d4.w;
}"#,
);
assert!(
out != src,
"rp4 substitution anchors must match gemm_q1_xt_src"
);
assert!(
!out.contains("array<f16>;"),
"flat scale binding must be replaced"
);
out
}
fn gemm_native_skeleton(w_binding: &str, w_stage_decl: &str, w_stage: &str, inner: &str) -> String {
format!(
r##"{w_binding}
@group(0) @binding(1) var<storage, read> x: array<vec4<f32>>; // [ncols, N/4]
@group(0) @binding(2) var<storage, read_write> y: array<f32>; // [ncols, M]
@group(0) @binding(3) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)
const BM: u32 = 32u;
const BN: u32 = 16u;
var<workgroup> xsv: array<vec4<f32>, 128>;
var<workgroup> xtmp: array<vec4<f32>, 128>;
{w_stage_decl}
@compute @workgroup_size(32)
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 ty = lid / 4u;
let tx = lid % 4u;
let row0 = wid.x * BM + ty * 4u;
let col0 = wid.y * BN;
var acc0 = vec4<f32>(0.0);
var acc1 = vec4<f32>(0.0);
var acc2 = vec4<f32>(0.0);
var acc3 = vec4<f32>(0.0);
for (var b = 0u; b < nblk; b = b + 1u) {{
for (var j = 0u; j < 4u; j = j + 1u) {{
let idx = lid * 4u + j;
let cn = idx / 8u;
let e = idx % 8u;
var v = vec4<f32>(0.0);
if (col0 + cn < ncols) {{ v = x[(col0 + cn) * xstride + b * 8u + e]; }}
xtmp[idx] = v;
}}
workgroupBarrier();
for (var j = 0u; j < 4u; j = j + 1u) {{
let idx = lid * 4u + j;
let kk = idx / 4u;
let c4 = idx % 4u;
let e = kk / 4u;
let comp = kk % 4u;
xsv[idx] = vec4<f32>(
xtmp[(c4 * 4u) * 8u + e][comp],
xtmp[(c4 * 4u + 1u) * 8u + e][comp],
xtmp[(c4 * 4u + 2u) * 8u + e][comp],
xtmp[(c4 * 4u + 3u) * 8u + e][comp],
);
}}
{{
let row = wid.x * BM + lid;
{w_stage}
}}
workgroupBarrier();
{{
{inner}
}}
workgroupBarrier();
}}
{epilogue}
}}
"##,
epilogue = (0..4)
.map(|r| format!(
r##" {{
let row = row0 + {r}u;
if (row < m) {{
for (var c = 0u; c < 4u; c = c + 1u) {{
let col = col0 + tx * 4u + c;
if (col < ncols) {{
let o = col * m + row;
var v = acc{r}.x;
if (c == 1u) {{ v = acc{r}.y; }}
if (c == 2u) {{ v = acc{r}.z; }}
if (c == 3u) {{ v = acc{r}.w; }}
if (dims.z == 1u) {{ y[o] = y[o] + v; }} else {{ y[o] = v; }}
}}
}}
}}
}}"##
))
.collect::<Vec<_>>()
.join("\n")
)
}
pub fn gemm_f16_src() -> String {
let mut inner = String::new();
for r in 0..4 {
for h in 0..4 {
for j in 0..4 {
let e0 = h * 8 + j * 2;
inner.push_str(&format!(
" {{ let wf = unpack2x16float(wq4[(ty * 4u + {r}u) * 4u + {h}u][{j}u]); \
acc{r} = acc{r} + wf.x * xsv[{a}u + tx] + wf.y * xsv[{b}u + tx]; }}\n",
a = e0 * 4,
b = (e0 + 1) * 4,
));
}
}
}
gemm_native_skeleton(
"@group(0) @binding(0) var<storage, read> quants: array<vec4<u32>>; // f16 weights, 8/vec4",
"var<workgroup> wq4: array<vec4<u32>, 128>; // [BM][4] one 32-elem f16 block per row",
r#" for (var h = 0u; h < 4u; h = h + 1u) {
if (row < m) {
wq4[lid * 4u + h] = quants[(row * nblk + b) * 4u + h];
} else {
wq4[lid * 4u + h] = vec4<u32>();
}
}"#,
&inner,
)
}
pub fn gemm_q8_0n_src() -> String {
let mut inner = String::new();
for r in 0..4 {
inner.push_str(&format!(" let d{r} = wsc[ty * 4u + {r}u];\n"));
for h in 0..2 {
for j in 0..4 {
let e0 = (h * 4 + j) * 4;
inner.push_str(&format!(
" {{ let qf = vec4<f32>(unpack4xI8(wq4[(ty * 4u + {r}u) * 2u + {h}u][{j}u])); \
acc{r} = acc{r} + d{r} * (qf.x * xsv[{a}u + tx] + qf.y * xsv[{b}u + tx] + qf.z * xsv[{c}u + tx] + qf.w * xsv[{d}u + tx]); }}\n",
a = e0 * 4,
b = (e0 + 1) * 4,
c = (e0 + 2) * 4,
d = (e0 + 3) * 4,
));
}
}
}
gemm_native_skeleton(
"@group(0) @binding(0) var<storage, read> quants: array<u32>; // 9 words / 32 weights (padded)",
"var<workgroup> wq4: array<vec4<u32>, 64>; // [BM][2] i8 planes\nvar<workgroup> wsc: array<f32, 32>; // [BM] block scales",
r#" if (row < m) {
let base = (row * nblk + b) * 9u;
wsc[lid] = unpack2x16float(quants[base]).x;
for (var h = 0u; h < 2u; h = h + 1u) {
wq4[lid * 2u + h] = vec4<u32>(
quants[base + 1u + h * 4u],
quants[base + 2u + h * 4u],
quants[base + 3u + h * 4u],
quants[base + 4u + h * 4u],
);
}
} else {
wsc[lid] = 0.0;
wq4[lid * 2u] = vec4<u32>();
wq4[lid * 2u + 1u] = vec4<u32>();
}"#,
&inner,
)
}
pub fn gemm_q4_src() -> String {
[Q4_FN, r##"@group(0) @binding(0) var<storage, read> scales: array<f16>;
@group(0) @binding(1) var<storage, read> quants: array<vec4<u32>>; // one vec4 = one 32-elem block
@group(0) @binding(2) var<storage, read> x: array<vec4<f32>>; // [ncols, N/4]
@group(0) @binding(3) var<storage, read_write> y: array<f32>; // [ncols, M]
@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)
const BM: u32 = 32u;
const BN: u32 = 16u;
var<workgroup> xsv: array<vec4<f32>, 128>; // [BK=32][BN/4=4] column-vec4 activation tile
var<workgroup> xtmp: array<vec4<f32>, 128>; // [BN][BK/4=8] coalesced landing tile
var<workgroup> wsc: array<f32, 32>; // [BM] block scales
var<workgroup> wq4: array<vec4<u32>, 32>; // [BM] block quants
@compute @workgroup_size(32)
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 ty = lid / 4u; // 0..7 → 4-row group
let tx = lid % 4u; // 0..3 → one column-vec4 (4 columns)
let row0 = wid.x * BM + ty * 4u;
let col0 = wid.y * BN;
var acc0 = vec4<f32>(0.0);
var acc1 = vec4<f32>(0.0);
var acc2 = vec4<f32>(0.0);
var acc3 = vec4<f32>(0.0);
for (var b = 0u; b < nblk; b = b + 1u) {
// Two-phase staging: coalesced landing (consecutive vec4s along each column's k-run),
// then a shared→shared transpose into column-vec4s; weights staged once per WG.
// BN=16/WG=32 is the MEASURED optimum: the BN=32/WG=64 variant (half the weight DRAM
// traffic) ran 1.5% SLOWER e2e — the projections are no longer weight-bound here.
for (var j = 0u; j < 4u; j = j + 1u) {
let idx = lid * 4u + j;
let cn = idx / 8u;
let e = idx % 8u;
var v = vec4<f32>(0.0);
if (col0 + cn < ncols) { v = x[(col0 + cn) * xstride + b * 8u + e]; }
xtmp[idx] = v;
}
workgroupBarrier();
for (var j = 0u; j < 4u; j = j + 1u) {
let idx = lid * 4u + j;
let kk = idx / 4u;
let c4 = idx % 4u;
let e = kk / 4u;
let comp = kk % 4u;
xsv[idx] = vec4<f32>(
xtmp[(c4 * 4u) * 8u + e][comp],
xtmp[(c4 * 4u + 1u) * 8u + e][comp],
xtmp[(c4 * 4u + 2u) * 8u + e][comp],
xtmp[(c4 * 4u + 3u) * 8u + e][comp],
);
}
{
let row = wid.x * BM + lid;
if (row < m) {
wsc[lid] = f32(scales[row * nblk + b]);
wq4[lid] = quants[row * nblk + b];
} else {
wsc[lid] = 0.0;
wq4[lid] = vec4<u32>();
}
}
workgroupBarrier();
{
let d0 = wsc[ty * 4u + 0u];
let q0 = wq4[ty * 4u + 0u];
let d1 = wsc[ty * 4u + 1u];
let q1 = wq4[ty * 4u + 1u];
let d2 = wsc[ty * 4u + 2u];
let q2 = wq4[ty * 4u + 2u];
let d3 = wsc[ty * 4u + 3u];
let q3 = wq4[ty * 4u + 3u];
let wl00 = d0 * q4_lo(q0.x);
let wh00 = d0 * q4_hi(q0.x);
acc0 = acc0 + wl00.x * xsv[0u + tx] + wl00.y * xsv[4u + tx] + wl00.z * xsv[8u + tx] + wl00.w * xsv[12u + tx];
acc0 = acc0 + wh00.x * xsv[64u + tx] + wh00.y * xsv[68u + tx] + wh00.z * xsv[72u + tx] + wh00.w * xsv[76u + tx];
let wl01 = d0 * q4_lo(q0.y);
let wh01 = d0 * q4_hi(q0.y);
acc0 = acc0 + wl01.x * xsv[16u + tx] + wl01.y * xsv[20u + tx] + wl01.z * xsv[24u + tx] + wl01.w * xsv[28u + tx];
acc0 = acc0 + wh01.x * xsv[80u + tx] + wh01.y * xsv[84u + tx] + wh01.z * xsv[88u + tx] + wh01.w * xsv[92u + tx];
let wl02 = d0 * q4_lo(q0.z);
let wh02 = d0 * q4_hi(q0.z);
acc0 = acc0 + wl02.x * xsv[32u + tx] + wl02.y * xsv[36u + tx] + wl02.z * xsv[40u + tx] + wl02.w * xsv[44u + tx];
acc0 = acc0 + wh02.x * xsv[96u + tx] + wh02.y * xsv[100u + tx] + wh02.z * xsv[104u + tx] + wh02.w * xsv[108u + tx];
let wl03 = d0 * q4_lo(q0.w);
let wh03 = d0 * q4_hi(q0.w);
acc0 = acc0 + wl03.x * xsv[48u + tx] + wl03.y * xsv[52u + tx] + wl03.z * xsv[56u + tx] + wl03.w * xsv[60u + tx];
acc0 = acc0 + wh03.x * xsv[112u + tx] + wh03.y * xsv[116u + tx] + wh03.z * xsv[120u + tx] + wh03.w * xsv[124u + tx];
let wl10 = d1 * q4_lo(q1.x);
let wh10 = d1 * q4_hi(q1.x);
acc1 = acc1 + wl10.x * xsv[0u + tx] + wl10.y * xsv[4u + tx] + wl10.z * xsv[8u + tx] + wl10.w * xsv[12u + tx];
acc1 = acc1 + wh10.x * xsv[64u + tx] + wh10.y * xsv[68u + tx] + wh10.z * xsv[72u + tx] + wh10.w * xsv[76u + tx];
let wl11 = d1 * q4_lo(q1.y);
let wh11 = d1 * q4_hi(q1.y);
acc1 = acc1 + wl11.x * xsv[16u + tx] + wl11.y * xsv[20u + tx] + wl11.z * xsv[24u + tx] + wl11.w * xsv[28u + tx];
acc1 = acc1 + wh11.x * xsv[80u + tx] + wh11.y * xsv[84u + tx] + wh11.z * xsv[88u + tx] + wh11.w * xsv[92u + tx];
let wl12 = d1 * q4_lo(q1.z);
let wh12 = d1 * q4_hi(q1.z);
acc1 = acc1 + wl12.x * xsv[32u + tx] + wl12.y * xsv[36u + tx] + wl12.z * xsv[40u + tx] + wl12.w * xsv[44u + tx];
acc1 = acc1 + wh12.x * xsv[96u + tx] + wh12.y * xsv[100u + tx] + wh12.z * xsv[104u + tx] + wh12.w * xsv[108u + tx];
let wl13 = d1 * q4_lo(q1.w);
let wh13 = d1 * q4_hi(q1.w);
acc1 = acc1 + wl13.x * xsv[48u + tx] + wl13.y * xsv[52u + tx] + wl13.z * xsv[56u + tx] + wl13.w * xsv[60u + tx];
acc1 = acc1 + wh13.x * xsv[112u + tx] + wh13.y * xsv[116u + tx] + wh13.z * xsv[120u + tx] + wh13.w * xsv[124u + tx];
let wl20 = d2 * q4_lo(q2.x);
let wh20 = d2 * q4_hi(q2.x);
acc2 = acc2 + wl20.x * xsv[0u + tx] + wl20.y * xsv[4u + tx] + wl20.z * xsv[8u + tx] + wl20.w * xsv[12u + tx];
acc2 = acc2 + wh20.x * xsv[64u + tx] + wh20.y * xsv[68u + tx] + wh20.z * xsv[72u + tx] + wh20.w * xsv[76u + tx];
let wl21 = d2 * q4_lo(q2.y);
let wh21 = d2 * q4_hi(q2.y);
acc2 = acc2 + wl21.x * xsv[16u + tx] + wl21.y * xsv[20u + tx] + wl21.z * xsv[24u + tx] + wl21.w * xsv[28u + tx];
acc2 = acc2 + wh21.x * xsv[80u + tx] + wh21.y * xsv[84u + tx] + wh21.z * xsv[88u + tx] + wh21.w * xsv[92u + tx];
let wl22 = d2 * q4_lo(q2.z);
let wh22 = d2 * q4_hi(q2.z);
acc2 = acc2 + wl22.x * xsv[32u + tx] + wl22.y * xsv[36u + tx] + wl22.z * xsv[40u + tx] + wl22.w * xsv[44u + tx];
acc2 = acc2 + wh22.x * xsv[96u + tx] + wh22.y * xsv[100u + tx] + wh22.z * xsv[104u + tx] + wh22.w * xsv[108u + tx];
let wl23 = d2 * q4_lo(q2.w);
let wh23 = d2 * q4_hi(q2.w);
acc2 = acc2 + wl23.x * xsv[48u + tx] + wl23.y * xsv[52u + tx] + wl23.z * xsv[56u + tx] + wl23.w * xsv[60u + tx];
acc2 = acc2 + wh23.x * xsv[112u + tx] + wh23.y * xsv[116u + tx] + wh23.z * xsv[120u + tx] + wh23.w * xsv[124u + tx];
let wl30 = d3 * q4_lo(q3.x);
let wh30 = d3 * q4_hi(q3.x);
acc3 = acc3 + wl30.x * xsv[0u + tx] + wl30.y * xsv[4u + tx] + wl30.z * xsv[8u + tx] + wl30.w * xsv[12u + tx];
acc3 = acc3 + wh30.x * xsv[64u + tx] + wh30.y * xsv[68u + tx] + wh30.z * xsv[72u + tx] + wh30.w * xsv[76u + tx];
let wl31 = d3 * q4_lo(q3.y);
let wh31 = d3 * q4_hi(q3.y);
acc3 = acc3 + wl31.x * xsv[16u + tx] + wl31.y * xsv[20u + tx] + wl31.z * xsv[24u + tx] + wl31.w * xsv[28u + tx];
acc3 = acc3 + wh31.x * xsv[80u + tx] + wh31.y * xsv[84u + tx] + wh31.z * xsv[88u + tx] + wh31.w * xsv[92u + tx];
let wl32 = d3 * q4_lo(q3.z);
let wh32 = d3 * q4_hi(q3.z);
acc3 = acc3 + wl32.x * xsv[32u + tx] + wl32.y * xsv[36u + tx] + wl32.z * xsv[40u + tx] + wl32.w * xsv[44u + tx];
acc3 = acc3 + wh32.x * xsv[96u + tx] + wh32.y * xsv[100u + tx] + wh32.z * xsv[104u + tx] + wh32.w * xsv[108u + tx];
let wl33 = d3 * q4_lo(q3.w);
let wh33 = d3 * q4_hi(q3.w);
acc3 = acc3 + wl33.x * xsv[48u + tx] + wl33.y * xsv[52u + tx] + wl33.z * xsv[56u + tx] + wl33.w * xsv[60u + tx];
acc3 = acc3 + wh33.x * xsv[112u + tx] + wh33.y * xsv[116u + tx] + wh33.z * xsv[120u + tx] + wh33.w * xsv[124u + tx];
}
workgroupBarrier();
}
{
let row = row0 + 0u;
if (row < m) {
for (var c = 0u; c < 4u; c = c + 1u) {
let col = col0 + tx * 4u + c;
if (col < ncols) {
let o = col * m + row;
var v = acc0.x;
if (c == 1u) { v = acc0.y; }
if (c == 2u) { v = acc0.z; }
if (c == 3u) { v = acc0.w; }
if (dims.z == 1u) { y[o] = y[o] + v; } else { y[o] = v; }
}
}
}
}
{
let row = row0 + 1u;
if (row < m) {
for (var c = 0u; c < 4u; c = c + 1u) {
let col = col0 + tx * 4u + c;
if (col < ncols) {
let o = col * m + row;
var v = acc1.x;
if (c == 1u) { v = acc1.y; }
if (c == 2u) { v = acc1.z; }
if (c == 3u) { v = acc1.w; }
if (dims.z == 1u) { y[o] = y[o] + v; } else { y[o] = v; }
}
}
}
}
{
let row = row0 + 2u;
if (row < m) {
for (var c = 0u; c < 4u; c = c + 1u) {
let col = col0 + tx * 4u + c;
if (col < ncols) {
let o = col * m + row;
var v = acc2.x;
if (c == 1u) { v = acc2.y; }
if (c == 2u) { v = acc2.z; }
if (c == 3u) { v = acc2.w; }
if (dims.z == 1u) { y[o] = y[o] + v; } else { y[o] = v; }
}
}
}
}
{
let row = row0 + 3u;
if (row < m) {
for (var c = 0u; c < 4u; c = c + 1u) {
let col = col0 + tx * 4u + c;
if (col < ncols) {
let o = col * m + row;
var v = acc3.x;
if (c == 1u) { v = acc3.y; }
if (c == 2u) { v = acc3.z; }
if (c == 3u) { v = acc3.w; }
if (dims.z == 1u) { y[o] = y[o] + v; } else { y[o] = v; }
}
}
}
}
}
"##].concat()
}
pub fn gemm_q4_splitk_src(s_slices: usize) -> String {
let src = gemm_q4_src();
let a1 = " let m = dims.x; let n = dims.y; let ncols = dims.w;\n";
let a2 = " let col0 = wid.y * BN;\n";
let a3 = " for (var b = 0u; b < nblk; b = b + 1u) {\n";
let a4 = "if (dims.z == 1u) { y[o] = y[o] + v; } else { y[o] = v; }";
for a in [a1, a2, a3, a4] {
assert!(src.contains(a), "splitk anchor drifted: {a:?}");
}
src.replace(
a1,
" let m = dims.x; let n = dims.y; let ncols = dims.w;\n let nct = (ncols + 15u) / 16u;\n let sl = wid.y / nct;\n let cwy = wid.y % nct;\n",
)
.replace(a2, " let col0 = cwy * BN;\n")
.replace(
a3,
&format!(
" let per = (nblk + {s}u - 1u) / {s}u;\n let b0 = sl * per;\n let b1 = min(b0 + per, nblk);\n for (var b = b0; b < b1; b = b + 1u) {{\n",
s = s_slices
),
)
.replace(a4, "y[sl * ncols * m + o] = v;")
}
pub fn gemm_q4_splitk_reduce_src(s_slices: usize) -> String {
format!(
r#"@group(0) @binding(0) var<storage, read> part: array<f32>;
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
@group(0) @binding(2) var<uniform> dims: vec4<u32>; // (m, n, acc, ncols)
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {{
let len = dims.w * dims.x;
let i = gid.x + gid.y * nwg.x * 256u;
if (i >= len) {{ return; }}
var s = 0.0;
if (dims.z == 1u) {{ s = y[i]; }}
for (var sl = 0u; sl < {s_slices}u; sl++) {{ s = s + part[sl * len + i]; }}
y[i] = s;
}}
"#
)
}
pub fn gemv_q4_k_src() -> String {
format!(
r#"{Q4_FN}
@group(0) @binding(0) var<storage, read> scales: array<f16>;
@group(0) @binding(1) var<storage, read> quants: array<vec4<u32>>; // one vec4 = one 32-elem block
@group(0) @binding(2) var<storage, read> x: array<vec4<f32>>; // [ncols, N/4]
@group(0) @binding(3) var<storage, read_write> y: array<f32>; // [ncols, M]
@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)
const NR: u32 = 16u; // rows per workgroup
const LANES: u32 = 16u; // threads per row
const KC: u32 = 4u; // columns sharing one weight stream
const TB: u32 = 16u; // blocks per row per tile (= LANES)
const TV: u32 = 128u; // x vec4s per tile per column (TB·32/4)
var<workgroup> xs: array<vec4<f32>, 512>; // KC·TV
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;
var d_cur = 0.0;
var q_cur = vec4<u32>();
if (row < m && lane < nblk) {{
d_cur = f32(scales[row * nblk + lane]);
q_cur = quants[row * nblk + lane];
}}
for (var t = 0u; t < ntiles; t = t + 1u) {{
// Cooperative x load: 512 vec4s (KC columns × TV), 2 per thread; OOB / idle cols → 0.
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();
// Software pipeline: tile t+1's weights are ISSUED here, before tile t's dots —
// the DRAM latency hides behind the arithmetic (same FP order, bitwise-identical).
let bn = (t + 1u) * TB + lane;
var d_nxt = 0.0;
var q_nxt = vec4<u32>();
if (t + 1u < ntiles && row < m && bn < nblk) {{
d_nxt = f32(scales[row * nblk + bn]);
q_nxt = quants[row * nblk + bn];
}}
let b = t * TB + lane;
if (row < m && b < nblk) {{
let d = d_cur;
let q = q_cur;
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;
}}
}}
d_cur = d_nxt;
q_cur = q_nxt;
workgroupBarrier();
}}
// Per-row reduction over the 16 lanes, one column at a time.
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;
if (dims.z == 1u) {{ y[yo] = y[yo] + v; }} else {{ y[yo] = v; }}
}}
workgroupBarrier();
}}
}}
"#
)
}
pub fn gemv_q4_k_lcpp_src() -> String {
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 fma(vec4<f32>(unpack4xU8(word & 0xF0F0F0F0u)), vec4<f32>(0.0625), vec4<f32>(-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>; // (M, N, acc, ncols)
@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 = min(row0 + r, m - 1u);
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) {
let yo = col * m + row0 + r;
if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }
}
}
}
}
"#
.to_string()
}
pub fn gemv_q4k_k_lcpp_src() -> String {
r#"
fn nib4(word: u32, p: u32) -> vec4<f32> {
return vec4<f32>(unpack4xU8((word >> p) & 0x0F0F0F0Fu));
}
// byte j (0..11) of the packed scale region starting at word w0+1
fn scb(base: u32, j: u32) -> u32 { return (blocks[base + 1u + (j >> 2u)] >> (8u * (j & 3u))) & 0xFFu; }
@group(0) @binding(0) var<storage, read> blocks: array<u32>; // 36 words / 256 weights
@group(0) @binding(1) var<storage, read> x: array<vec4<f32>>;
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
@group(0) @binding(3) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)
@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; // 32-element sub-blocks per row
let nsb = n / 256u; // superblocks per row
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 sb = b / 8u; // superblock
let s = b % 8u; // sub-block within it
let g = s / 2u; // 32-byte nibble group
let p = 4u * (s % 2u);
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 ones = vec4<f32>(1.0);
let xsum = dot(v0, ones) + dot(v1, ones) + dot(v2, ones) + dot(v3, ones)
+ dot(v4, ones) + dot(v5, ones) + dot(v6, ones) + dot(v7, ones);
for (var r = 0u; r < 4u; r = r + 1u) {
let row = min(row0 + r, m - 1u);
let w0 = (row * nsb + sb) * 36u;
let dm = unpack2x16float(blocks[w0]);
// 6-bit scale/min for sub-block s (gguf get_scale_min layout)
var sc: u32; var mn: u32;
if (s < 4u) {
sc = scb(w0, s) & 63u;
mn = scb(w0, s + 4u) & 63u;
} else {
sc = (scb(w0, s + 4u) & 0x0Fu) | (((scb(w0, s - 4u) >> 6u) & 3u) << 4u);
mn = (scb(w0, s + 4u) >> 4u) | (((scb(w0, s) >> 6u) & 3u) << 4u);
}
let qw = w0 + 4u + g * 8u;
var qdot = dot(nib4(blocks[qw], p), v0) + dot(nib4(blocks[qw + 1u], p), v1);
qdot = qdot + dot(nib4(blocks[qw + 2u], p), v2) + dot(nib4(blocks[qw + 3u], p), v3);
qdot = qdot + dot(nib4(blocks[qw + 4u], p), v4) + dot(nib4(blocks[qw + 5u], p), v5);
qdot = qdot + dot(nib4(blocks[qw + 6u], p), v6) + dot(nib4(blocks[qw + 7u], p), v7);
acc[r] = acc[r] + dm.x * f32(sc) * qdot - dm.y * f32(mn) * xsum;
}
}
let tot = subgroupAdd(acc);
if (sid == 0u) {
for (var r = 0u; r < 4u; r = r + 1u) {
if (row0 + r < m) {
let yo = col * m + row0 + r;
if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }
}
}
}
}
"#
.to_string()
}
pub fn gemv_q6k_k_lcpp_src() -> String {
r#"
fn nib4(word: u32, p: u32) -> vec4<f32> {
return vec4<f32>(unpack4xU8((word >> p) & 0x0F0F0F0Fu));
}
fn bits4(word: u32, p: u32) -> vec4<f32> {
return vec4<f32>(unpack4xU8((word >> p) & 0x03030303u));
}
@group(0) @binding(0) var<storage, read> blocks: array<u32>; // 53 words / 256 weights (padded)
@group(0) @binding(1) var<storage, read> x: array<vec4<f32>>;
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
@group(0) @binding(3) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)
@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 nsb = n / 256u;
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 sb = b / 8u;
let s = b % 8u;
let h = s / 4u; // 128-element half
let q = s % 4u; // quarter within the half
let np = 4u * (q / 2u); // ql nibble plane shift
let bp = 2u * q; // qh bit-pair shift
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 = min(row0 + r, m - 1u);
let w0 = (row * nsb + sb) * 53u;
let d = unpack2x16float(blocks[w0 + 52u]).x;
// two signed 6-bit-range i8 scales for this sub-block (elements 0-15 / 16-31)
let scw = blocks[w0 + 48u + (h * 8u + q * 2u) / 4u];
let sh0 = 8u * ((h * 8u + q * 2u) % 4u);
let sc0 = f32(bitcast<i32>((scw << (24u - sh0)) & 0xFF000000u) >> 24u);
let sc1 = f32(bitcast<i32>((scw << (16u - sh0)) & 0xFF000000u) >> 24u);
let qlw = w0 + (h * 64u + (q % 2u) * 32u) / 4u;
let qhw = w0 + 32u + (h * 32u) / 4u;
var d0 = 0.0; var d1 = 0.0;
for (var i = 0u; i < 4u; i = i + 1u) {
let qv = nib4(blocks[qlw + i], np) + bits4(blocks[qhw + i], bp) * 16.0 - vec4<f32>(32.0);
let qv2 = nib4(blocks[qlw + 4u + i], np) + bits4(blocks[qhw + 4u + i], bp) * 16.0 - vec4<f32>(32.0);
switch i {
case 0u: { d0 = d0 + dot(qv, v0); d1 = d1 + dot(qv2, v4); }
case 1u: { d0 = d0 + dot(qv, v1); d1 = d1 + dot(qv2, v5); }
case 2u: { d0 = d0 + dot(qv, v2); d1 = d1 + dot(qv2, v6); }
default: { d0 = d0 + dot(qv, v3); d1 = d1 + dot(qv2, v7); }
}
}
acc[r] = acc[r] + d * (sc0 * d0 + sc1 * d1);
}
}
let tot = subgroupAdd(acc);
if (sid == 0u) {
for (var r = 0u; r < 4u; r = r + 1u) {
if (row0 + r < m) {
let yo = col * m + row0 + r;
if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }
}
}
}
}
"#
.to_string()
}
pub fn q6k_pad_blocks(raw: &[u8], nblocks: usize) -> Vec<u8> {
let mut out = vec![0u8; nblocks * 212];
for b in 0..nblocks {
out[b * 212..b * 212 + 210].copy_from_slice(&raw[b * 210..(b + 1) * 210]);
}
out
}
pub fn gemv_q5_0n_k_lcpp_src() -> String {
r#"
fn nib4(word: u32, p: u32) -> vec4<f32> {
return vec4<f32>(unpack4xU8((word >> p) & 0x0F0F0F0Fu));
}
fn hi4(qh: u32, e0: u32) -> vec4<f32> {
return vec4<f32>(vec4<u32>(qh >> e0, qh >> (e0 + 1u), qh >> (e0 + 2u), qh >> (e0 + 3u)) & vec4<u32>(1u));
}
@group(0) @binding(0) var<storage, read> blocks: array<u32>; // 6 words / 32 weights (padded)
@group(0) @binding(1) var<storage, read> x: array<vec4<f32>>;
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
@group(0) @binding(3) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)
@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];
let ones = vec4<f32>(1.0);
let xsum = dot(v0, ones) + dot(v1, ones) + dot(v2, ones) + dot(v3, ones)
+ dot(v4, ones) + dot(v5, ones) + dot(v6, ones) + dot(v7, ones);
for (var r = 0u; r < 4u; r = r + 1u) {
let row = min(row0 + r, m - 1u);
let w0 = (row * nblk + b) * 6u;
let d = unpack2x16float(blocks[w0]).x;
let qh = blocks[w0 + 1u];
var s = dot(nib4(blocks[w0 + 2u], 0u) + hi4(qh, 0u) * 16.0, v0);
s = s + dot(nib4(blocks[w0 + 3u], 0u) + hi4(qh, 4u) * 16.0, v1);
s = s + dot(nib4(blocks[w0 + 4u], 0u) + hi4(qh, 8u) * 16.0, v2);
s = s + dot(nib4(blocks[w0 + 5u], 0u) + hi4(qh, 12u) * 16.0, v3);
s = s + dot(nib4(blocks[w0 + 2u], 4u) + hi4(qh, 16u) * 16.0, v4);
s = s + dot(nib4(blocks[w0 + 3u], 4u) + hi4(qh, 20u) * 16.0, v5);
s = s + dot(nib4(blocks[w0 + 4u], 4u) + hi4(qh, 24u) * 16.0, v6);
s = s + dot(nib4(blocks[w0 + 5u], 4u) + hi4(qh, 28u) * 16.0, v7);
// fold the −16 zero-point through the x-sum instead of per element
acc[r] = acc[r] + d * (s - 16.0 * xsum);
}
}
let tot = subgroupAdd(acc);
if (sid == 0u) {
for (var r = 0u; r < 4u; r = r + 1u) {
if (row0 + r < m) {
let yo = col * m + row0 + r;
if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }
}
}
}
}
"#
.to_string()
}
pub fn gemv_q8_0n_k_lcpp_src() -> String {
r#"
fn q8b(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xI8(word)); }
@group(0) @binding(0) var<storage, read> blocks: array<u32>; // 9 words / 32 weights (padded)
@group(0) @binding(1) var<storage, read> x: array<vec4<f32>>;
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
@group(0) @binding(3) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)
@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 = min(row0 + r, m - 1u);
let w0 = (row * nblk + b) * 9u;
let d = unpack2x16float(blocks[w0]).x;
var s = dot(q8b(blocks[w0 + 1u]), v0) + dot(q8b(blocks[w0 + 2u]), v1);
s = s + dot(q8b(blocks[w0 + 3u]), v2) + dot(q8b(blocks[w0 + 4u]), v3);
s = s + dot(q8b(blocks[w0 + 5u]), v4) + dot(q8b(blocks[w0 + 6u]), v5);
s = s + dot(q8b(blocks[w0 + 7u]), v6) + dot(q8b(blocks[w0 + 8u]), 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) {
let yo = col * m + row0 + r;
if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }
}
}
}
}
"#
.to_string()
}
pub fn mlp_gate_q8_0n_lcpp_src() -> String {
r#"
fn q8b(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xI8(word)); }
@group(0) @binding(0) var<storage, read> w1: array<u32>; // gate, 9 words/32 weights
@group(0) @binding(1) var<storage, read> w3: array<u32>; // up, 9 words/32 weights
@group(0) @binding(2) var<storage, read> x: array<vec4<f32>>; // [ncols, N/4] PRE-NORMED
@group(0) @binding(3) var<storage, read_write> y: array<f32>; // [ncols, M]
@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, _, ncols)
@group(0) @binding(5) var<uniform> epsm: vec4<f32>; // (_, gelu-flag, _, _)
@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 ag = vec4<f32>(0.0);
var au = 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 = min(row0 + r, m - 1u);
let w0 = (row * nblk + b) * 9u;
{
let d = unpack2x16float(w1[w0]).x;
var sv = dot(q8b(w1[w0 + 1u]), v0) + dot(q8b(w1[w0 + 2u]), v1);
sv = sv + dot(q8b(w1[w0 + 3u]), v2) + dot(q8b(w1[w0 + 4u]), v3);
sv = sv + dot(q8b(w1[w0 + 5u]), v4) + dot(q8b(w1[w0 + 6u]), v5);
sv = sv + dot(q8b(w1[w0 + 7u]), v6) + dot(q8b(w1[w0 + 8u]), v7);
ag[r] = ag[r] + d * sv;
}
{
let d = unpack2x16float(w3[w0]).x;
var sv = dot(q8b(w3[w0 + 1u]), v0) + dot(q8b(w3[w0 + 2u]), v1);
sv = sv + dot(q8b(w3[w0 + 3u]), v2) + dot(q8b(w3[w0 + 4u]), v3);
sv = sv + dot(q8b(w3[w0 + 5u]), v4) + dot(q8b(w3[w0 + 6u]), v5);
sv = sv + dot(q8b(w3[w0 + 7u]), v6) + dot(q8b(w3[w0 + 8u]), v7);
au[r] = au[r] + d * sv;
}
}
}
let tg = subgroupAdd(ag);
let tu = subgroupAdd(au);
if (sid == 0u) {
for (var r = 0u; r < 4u; r = r + 1u) {
if (row0 + r < m) {
let gate = tg[r];
let upv = tu[r];
var act: f32;
if (epsm.y != 0.0) {
let g3 = gate * gate * gate;
let targ = clamp(0.7978845608028654 * (gate + 0.044715 * g3), -20.0, 20.0);
act = 0.5 * gate * (1.0 + tanh(targ));
} else {
act = gate / (1.0 + exp(-gate));
}
y[col * m + row0 + r] = act * upv;
}
}
}
}
"#
.to_string()
}
pub fn q8_0n_lmhead_lcpp_src() -> String {
let src = gemv_q8_0n_k_lcpp_src();
for frag in [
"@group(0) @binding(3) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)",
" let row0 = wid.x * 4u;\n let col = wid.y;",
" if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }",
] {
assert!(src.contains(frag), "q8_0n lcpp source drifted: {frag}");
}
src.replace(
"@group(0) @binding(3) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)",
"@group(0) @binding(3) var<uniform> dims: vec4<u32>; // (M, N, ncols, gy_rows)",
)
.replace(
" let row0 = wid.x * 4u;\n let col = wid.y;",
" let col = wid.y / dims.w;\n let row0 = ((wid.y % dims.w) * 32768u + wid.x) * 4u;",
)
.replace(
" if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }",
" y[yo] = tot[r];",
)
}
pub fn q5_0_to_q8_0n(raw: &[u8], nblocks: usize) -> Vec<u8> {
let mut out = vec![0u8; nblocks * 36];
for b in 0..nblocks {
let (s, d) = (b * 22, b * 36);
out[d..d + 2].copy_from_slice(&raw[s..s + 2]); let qh = u32::from_le_bytes([raw[s + 2], raw[s + 3], raw[s + 4], raw[s + 5]]);
for l in 0..16 {
let byte = raw[s + 6 + l];
let lo = (byte & 0x0F) | ((((qh >> l) & 1) as u8) << 4);
let hi = (byte >> 4) | ((((qh >> (l + 16)) & 1) as u8) << 4);
out[d + 4 + l] = (lo as i8 - 16) as u8;
out[d + 4 + l + 16] = (hi as i8 - 16) as u8;
}
}
out
}
pub fn q4_0_to_q8_0n(raw: &[u8], nblocks: usize) -> Vec<u8> {
let mut out = vec![0u8; nblocks * 36];
for b in 0..nblocks {
let (s, d) = (b * 18, b * 36);
out[d..d + 2].copy_from_slice(&raw[s..s + 2]);
for l in 0..16 {
let byte = raw[s + 2 + l];
out[d + 4 + l] = ((byte & 0x0F) as i8 - 8) as u8;
out[d + 4 + l + 16] = ((byte >> 4) as i8 - 8) as u8;
}
}
out
}
pub fn iq4_nl_pad_blocks(raw: &[u8], nblocks: usize) -> Vec<u8> {
let mut out = Vec::with_capacity(nblocks * 20);
for b in raw[..nblocks * 18].chunks_exact(18) {
out.extend_from_slice(&b[0..2]);
out.extend_from_slice(&[0, 0]);
out.extend_from_slice(&b[2..18]);
}
out
}
pub fn gemv_iq4_nl_k_lcpp_src() -> String {
r#"
const KV = array<f32,16>(-127.0, -104.0, -83.0, -65.0, -49.0, -35.0, -22.0, -10.0,
1.0, 13.0, 25.0, 38.0, 53.0, 69.0, 89.0, 113.0);
fn kv4(word: u32, p: u32) -> vec4<f32> {
let q = unpack4xU8((word >> p) & 0x0F0F0F0Fu);
return vec4<f32>(KV[q.x], KV[q.y], KV[q.z], KV[q.w]);
}
@group(0) @binding(0) var<storage, read> blocks: array<u32>; // 5 words / 32 weights (padded)
@group(0) @binding(1) var<storage, read> x: array<vec4<f32>>;
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
@group(0) @binding(3) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)
@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 = min(row0 + r, m - 1u);
let w0 = (row * nblk + b) * 5u;
let d = unpack2x16float(blocks[w0]).x;
var s = dot(kv4(blocks[w0 + 1u], 0u), v0) + dot(kv4(blocks[w0 + 2u], 0u), v1);
s = s + dot(kv4(blocks[w0 + 3u], 0u), v2) + dot(kv4(blocks[w0 + 4u], 0u), v3);
s = s + dot(kv4(blocks[w0 + 1u], 4u), v4) + dot(kv4(blocks[w0 + 2u], 4u), v5);
s = s + dot(kv4(blocks[w0 + 3u], 4u), v6) + dot(kv4(blocks[w0 + 4u], 4u), 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) {
let yo = col * m + row0 + r;
if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }
}
}
}
}
"#
.to_string()
}
pub fn gemv_iq4_xs_k_lcpp_src() -> String {
r#"
const KV = array<f32,16>(-127.0, -104.0, -83.0, -65.0, -49.0, -35.0, -22.0, -10.0,
1.0, 13.0, 25.0, 38.0, 53.0, 69.0, 89.0, 113.0);
fn kv4(word: u32, p: u32) -> vec4<f32> {
let q = unpack4xU8((word >> p) & 0x0F0F0F0Fu);
return vec4<f32>(KV[q.x], KV[q.y], KV[q.z], KV[q.w]);
}
@group(0) @binding(0) var<storage, read> blocks: array<u32>; // 34 words / 256 weights
@group(0) @binding(1) var<storage, read> x: array<vec4<f32>>;
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
@group(0) @binding(3) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)
@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 / 256u;
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) {
for (var r = 0u; r < 4u; r = r + 1u) {
let row = min(row0 + r, m - 1u);
let base = (row * nblk + b) * 34u;
let w0 = blocks[base];
let d = unpack2x16float(w0).x;
let sh = w0 >> 16u;
let sl = blocks[base + 1u];
for (var sb = 0u; sb < 8u; sb = sb + 1u) {
let lo = (sl >> (8u * (sb / 2u) + 4u * (sb % 2u))) & 0xFu;
let hi = (sh >> (2u * sb)) & 3u;
let dl = d * (f32(lo | (hi << 4u)) - 32.0);
let q0 = base + 2u + sb * 4u;
let xv = xoff + b * 64u + sb * 8u;
var s = dot(kv4(blocks[q0], 0u), x[xv]) + dot(kv4(blocks[q0 + 1u], 0u), x[xv + 1u]);
s = s + dot(kv4(blocks[q0 + 2u], 0u), x[xv + 2u]) + dot(kv4(blocks[q0 + 3u], 0u), x[xv + 3u]);
s = s + dot(kv4(blocks[q0], 4u), x[xv + 4u]) + dot(kv4(blocks[q0 + 1u], 4u), x[xv + 5u]);
s = s + dot(kv4(blocks[q0 + 2u], 4u), x[xv + 6u]) + dot(kv4(blocks[q0 + 3u], 4u), x[xv + 7u]);
acc[r] = acc[r] + dl * s;
}
}
}
let tot = subgroupAdd(acc);
if (sid == 0u) {
for (var r = 0u; r < 4u; r = r + 1u) {
if (row0 + r < m) {
let yo = col * m + row0 + r;
if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }
}
}
}
}
"#
.to_string()
}
pub fn gemv_iq_grid_src(ty: u32) -> String {
let (words, body): (u32, &str) = match ty {
16 => (
17,
r#"
for (var g = 0u; g < 8u; g = g + 1u) {
let wlo = blocks[base + 1u + 2u * g];
let whi = blocks[base + 2u + 2u * g];
let db = d * (0.5 + f32(whi >> 28u)) * 0.25;
for (var j = 0u; j < 4u; j = j + 1u) {
let row = (wlo >> (8u * j)) & 0xFFu;
let sgn = ks((whi >> (7u * j)) & 0x7Fu);
let xi = xblk + g * 8u + j * 2u;
s = s + db * (dot(g4(grid[row * 2u]) * sgn4(sgn, 0u), x[xi])
+ dot(g4(grid[row * 2u + 1u]) * sgn4(sgn, 4u), x[xi + 1u]));
}
}"#,
),
17 => (
19,
r#"
for (var w = 0u; w < 32u; w = w + 1u) {
let v = (blocks[base + 1u + w / 2u] >> (16u * (w % 2u))) & 0xFFFFu;
let sidx = w / 2u;
let sb = (blocks[base + 17u + (sidx >> 3u)]
>> (8u * ((sidx >> 1u) & 3u) + 4u * (sidx & 1u))) & 0xFu;
let db = d * (0.5 + f32(sb)) * 0.25;
let row = v & 511u;
let sgn = ks(v >> 9u);
let xi = xblk + w * 2u;
s = s + db * (dot(g4(grid[row * 2u]) * sgn4(sgn, 0u), x[xi])
+ dot(g4(grid[row * 2u + 1u]) * sgn4(sgn, 4u), x[xi + 1u]));
}"#,
),
22 => (
21,
r#"
for (var b8 = 0u; b8 < 32u; b8 = b8 + 1u) {
let qsb = (blocks[base + 1u + b8 / 4u] >> (8u * (b8 % 4u))) & 0xFFu;
let qhb = (blocks[base + 17u + b8 / 16u] >> (8u * ((b8 / 4u) % 4u))) & 0xFFu;
let row = qsb | (((qhb >> (2u * (b8 % 4u))) & 3u) << 8u);
let sgb = (blocks[base + 9u + b8 / 4u] >> (8u * (b8 % 4u))) & 0xFFu;
let sidx = b8 / 2u;
let sc = (blocks[base + 19u + (sidx >> 3u)]
>> (8u * ((sidx >> 1u) & 3u) + 4u * (sidx & 1u))) & 0xFu;
let db = d * (0.5 + f32(sc)) * 0.25;
let xi = xblk + b8 * 2u;
s = s + db * (dot(g4(grid[row * 2u]) * sgn4(sgb, 0u), x[xi])
+ dot(g4(grid[row * 2u + 1u]) * sgn4(sgb, 4u), x[xi + 1u]));
}"#,
),
18 => (
25,
r#"
for (var g = 0u; g < 8u; g = g + 1u) {
let w = blocks[base + 17u + g];
let db = d * (0.5 + f32(w >> 28u)) * 0.5;
for (var j = 0u; j < 4u; j = j + 1u) {
let sgn = ks((w >> (7u * j)) & 0x7Fu);
for (var t = 0u; t < 2u; t = t + 1u) {
let qi = g * 8u + j * 2u + t;
let row = (blocks[base + 1u + qi / 4u] >> (8u * (qi % 4u))) & 0xFFu;
s = s + db * dot(g4(grid[row]) * sgn4(sgn, 4u * t), x[xblk + qi]);
}
}
}"#,
),
21 => (
28,
r#"
for (var i = 0u; i < 64u; i = i + 1u) {
let sc = (blocks[base + 27u] >> (4u * (i / 8u))) & 0xFu;
let db = d * f32(1u + 2u * sc);
let qsb = (blocks[base + 1u + i / 4u] >> (8u * (i % 4u))) & 0xFFu;
let hb = (blocks[base + 17u + i / 32u] >> (8u * ((i / 8u) % 4u))) & 0xFFu;
let row = qsb | (((hb >> (i % 8u)) & 1u) << 8u);
let sgb = (blocks[base + 19u + i / 8u] >> (8u * ((i / 2u) % 4u))) & 0xFFu;
s = s + db * dot(g4(grid[row]) * sgn4(sgb, (i & 1u) * 4u), x[xblk + i]);
}"#,
),
19 => (
13,
r#"
for (var g = 0u; g < 8u; g = g + 1u) {
let h = (blocks[base + 9u + g / 2u] >> (16u * (g % 2u))) & 0xFFFFu;
let dl = d * f32(2u * ((h >> 12u) & 7u) + 1u);
var delta = 0.125;
if ((h & 0x8000u) != 0u) { delta = -0.125; }
let dv = vec4<f32>(delta);
for (var j = 0u; j < 4u; j = j + 1u) {
let qi = g * 4u + j;
let qsb = (blocks[base + 1u + qi / 4u] >> (8u * (qi % 4u))) & 0xFFu;
let row = qsb | (((h >> (3u * j)) & 7u) << 8u);
let xi = xblk + g * 8u + j * 2u;
s = s + dl * (dot(g4(grid[row * 2u]) + dv, x[xi])
+ dot(g4(grid[row * 2u + 1u]) + dv, x[xi + 1u]));
}
}"#,
),
29 => (
14,
r#"
for (var s16 = 0u; s16 < 16u; s16 = s16 + 1u) {
let scw = (blocks[base + 12u + s16 / 8u] >> (16u * ((s16 / 4u) % 2u))) & 0xFFFFu;
let sc3 = (scw >> (3u * (s16 % 4u))) & 7u;
let dl = d * f32(2u * sc3 + 1u);
for (var h2 = 0u; h2 < 2u; h2 = h2 + 1u) {
let b8 = s16 * 2u + h2;
let nibb = (blocks[base + 8u + b8 / 8u] >> (8u * ((b8 / 2u) % 4u))) & 0xFFu;
let nib = (nibb >> (4u * (b8 % 2u))) & 0xFu;
let qsb = (blocks[base + b8 / 4u] >> (8u * (b8 % 4u))) & 0xFFu;
let row = qsb | ((nib & 7u) << 8u);
var delta = 0.125;
if ((nib & 8u) != 0u) { delta = -0.125; }
let dv = vec4<f32>(delta);
let xi = xblk + b8 * 2u;
s = s + dl * (dot(g4(grid[row * 2u]) + dv, x[xi])
+ dot(g4(grid[row * 2u + 1u]) + dv, x[xi + 1u]));
}
}"#,
),
other => panic!("gemv_iq_grid_src: not a grid-IQ type: {other}"),
};
let d_expr = if ty == 29 {
r#"let sw0 = blocks[base + 12u]; let sw1 = blocks[base + 13u];
let sc0 = sw0 & 0xFFFFu; let sc1 = sw0 >> 16u;
let sc2 = sw1 & 0xFFFFu; let sc3 = sw1 >> 16u;
let dbits = ((sc0 & 0xF000u) >> 12u) | ((sc1 & 0xF000u) >> 8u)
| ((sc2 & 0xF000u) >> 4u) | (sc3 & 0xF000u);
let d = unpack2x16float(dbits).x;"#
} else {
"let d = unpack2x16float(blocks[base]).x;"
};
format!(
r#"
fn ks(i: u32) -> u32 {{ return i | ((countOneBits(i) & 1u) << 7u); }}
fn g4(gword: u32) -> vec4<f32> {{ return vec4<f32>(unpack4xI8(gword)); }}
fn sgn4(sb: u32, p: u32) -> vec4<f32> {{
return vec4<f32>(1.0) - 2.0 * vec4<f32>(vec4<u32>(sb >> p, sb >> (p + 1u), sb >> (p + 2u), sb >> (p + 3u)) & vec4<u32>(1u));
}}
@group(0) @binding(0) var<storage, read> grid: array<u32>; // packed i8 codebook rows
@group(0) @binding(1) var<storage, read> blocks: array<u32>; // {words} words / 256 weights (padded)
@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>; // (M, N, acc, ncols)
@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 / 256u;
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) {{
for (var r = 0u; r < 4u; r = r + 1u) {{
let rr = min(row0 + r, m - 1u);
let base = (rr * nblk + b) * {words}u;
let xblk = xoff + b * 64u;
{d_expr}
var s = 0.0;
{body}
acc[r] = acc[r] + s;
}}
}}
let tot = subgroupAdd(acc);
if (sid == 0u) {{
for (var r = 0u; r < 4u; r = r + 1u) {{
if (row0 + r < m) {{
let yo = col * m + row0 + r;
if (dims.z == 1u) {{ y[yo] = y[yo] + tot[r]; }} else {{ y[yo] = tot[r]; }}
}}
}}
}}
}}
"#
)
}
pub fn iq_grid_geom(ty: u32) -> (usize, &'static [i8], usize) {
match ty {
16 => (17, &iq_tables::IQ2_XXS_GRID[..], 66),
17 => (19, &iq_tables::IQ2_XS_GRID[..], 74),
22 => (21, &iq_tables::IQ2_S_GRID[..], 82),
18 => (25, &iq_tables::IQ3_XXS_GRID[..], 98),
21 => (28, &iq_tables::IQ3_S_GRID[..], 110),
19 => (13, &iq_tables::IQ1_GRID[..], 50),
29 => (14, &iq_tables::IQ1_GRID[..], 56),
other => panic!("not a grid-IQ type: {other}"),
}
}
pub fn iq_pad_blocks(raw: &[u8], nblocks: usize, in_bytes: usize, out_bytes: usize) -> Vec<u8> {
assert!(out_bytes % 4 == 0 && out_bytes >= in_bytes + 2);
let mut out = Vec::with_capacity(nblocks * out_bytes);
for b in raw[..nblocks * in_bytes].chunks_exact(in_bytes) {
out.extend_from_slice(&b[0..2]);
out.extend_from_slice(&[0, 0]);
out.extend_from_slice(&b[2..in_bytes]);
out.resize(out.len() + (out_bytes - 2 - in_bytes), 0);
}
out
}
pub fn q5_0_pad_blocks(raw: &[u8], nblocks: usize) -> Vec<u8> {
let mut out = vec![0u8; nblocks * 24];
for b in 0..nblocks {
let (s, d) = (b * 22, b * 24);
out[d..d + 2].copy_from_slice(&raw[s..s + 2]); out[d + 4..d + 8].copy_from_slice(&raw[s + 2..s + 6]); out[d + 8..d + 24].copy_from_slice(&raw[s + 6..s + 22]); }
out
}
pub fn q8_0_pad_blocks(raw: &[u8], nblocks: usize) -> Vec<u8> {
let mut out = vec![0u8; nblocks * 36];
for b in 0..nblocks {
let (s, d) = (b * 34, b * 36);
out[d..d + 2].copy_from_slice(&raw[s..s + 2]);
out[d + 4..d + 36].copy_from_slice(&raw[s + 2..s + 34]);
}
out
}
pub fn gemv_q5k_k_lcpp_src() -> String {
let src = gemv_q4k_k_lcpp_src();
for frag in [
"let qw = w0 + 4u + g * 8u;",
"let w0 = (row * nsb + sb) * 36u;",
"var qdot = dot(nib4(blocks[qw], p), v0) + dot(nib4(blocks[qw + 1u], p), v1);",
"qdot = qdot + dot(nib4(blocks[qw + 2u], p), v2) + dot(nib4(blocks[qw + 3u], p), v3);",
"qdot = qdot + dot(nib4(blocks[qw + 4u], p), v4) + dot(nib4(blocks[qw + 5u], p), v5);",
"qdot = qdot + dot(nib4(blocks[qw + 6u], p), v6) + dot(nib4(blocks[qw + 7u], p), v7);",
] {
assert!(src.contains(frag), "q4k source drifted: {frag}");
}
src.replace(
"let w0 = (row * nsb + sb) * 36u;",
"let w0 = (row * nsb + sb) * 44u;",
)
.replace("let qw = w0 + 4u + g * 8u;", "let qw = w0 + 12u + g * 8u;
let hw = w0 + 4u;")
.replace(
"var qdot = dot(nib4(blocks[qw], p), v0) + dot(nib4(blocks[qw + 1u], p), v1);",
"var qdot = dot(nib4(blocks[qw], p) + bits1(blocks[hw], s) * 16.0, v0) + dot(nib4(blocks[qw + 1u], p) + bits1(blocks[hw + 1u], s) * 16.0, v1);",
)
.replace(
"qdot = qdot + dot(nib4(blocks[qw + 2u], p), v2) + dot(nib4(blocks[qw + 3u], p), v3);",
"qdot = qdot + dot(nib4(blocks[qw + 2u], p) + bits1(blocks[hw + 2u], s) * 16.0, v2) + dot(nib4(blocks[qw + 3u], p) + bits1(blocks[hw + 3u], s) * 16.0, v3);",
)
.replace(
"qdot = qdot + dot(nib4(blocks[qw + 4u], p), v4) + dot(nib4(blocks[qw + 5u], p), v5);",
"qdot = qdot + dot(nib4(blocks[qw + 4u], p) + bits1(blocks[hw + 4u], s) * 16.0, v4) + dot(nib4(blocks[qw + 5u], p) + bits1(blocks[hw + 5u], s) * 16.0, v5);",
)
.replace(
"qdot = qdot + dot(nib4(blocks[qw + 6u], p), v6) + dot(nib4(blocks[qw + 7u], p), v7);",
"qdot = qdot + dot(nib4(blocks[qw + 6u], p) + bits1(blocks[hw + 6u], s) * 16.0, v6) + dot(nib4(blocks[qw + 7u], p) + bits1(blocks[hw + 7u], s) * 16.0, v7);",
)
.replace(
"fn nib4(word: u32, p: u32) -> vec4<f32> {",
"fn bits1(word: u32, s: u32) -> vec4<f32> {
return vec4<f32>(vec4<u32>(word >> s, word >> (s + 8u), word >> (s + 16u), word >> (s + 24u)) & vec4<u32>(1u));
}
fn nib4(word: u32, p: u32) -> vec4<f32> {",
)
}
pub fn gemv_f16_k_lcpp_src() -> String {
r#"
fn f16x4(a: u32, b: u32) -> vec4<f32> {
let lo = unpack2x16float(a);
let hi = unpack2x16float(b);
return vec4<f32>(lo.x, lo.y, hi.x, hi.y);
}
@group(0) @binding(0) var<storage, read> w: array<vec4<u32>>; // f16 weights, 8 per vec4
@group(0) @binding(1) var<storage, read> x: array<vec4<f32>>;
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
@group(0) @binding(3) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)
@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 = min(row0 + r, m - 1u);
let wb = (row * nblk + b) * 4u;
let q0 = w[wb]; let q1 = w[wb + 1u];
let q2 = w[wb + 2u]; let q3 = w[wb + 3u];
var s = dot(f16x4(q0.x, q0.y), v0) + dot(f16x4(q0.z, q0.w), v1);
s = s + dot(f16x4(q1.x, q1.y), v2) + dot(f16x4(q1.z, q1.w), v3);
s = s + dot(f16x4(q2.x, q2.y), v4) + dot(f16x4(q2.z, q2.w), v5);
s = s + dot(f16x4(q3.x, q3.y), v6) + dot(f16x4(q3.z, q3.w), v7);
acc[r] = acc[r] + s;
}
}
let tot = subgroupAdd(acc);
if (sid == 0u) {
for (var r = 0u; r < 4u; r = r + 1u) {
if (row0 + r < m) {
let yo = col * m + row0 + r;
if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }
}
}
}
}
"#
.to_string()
}
pub fn f16_lmhead_lcpp_src() -> String {
let src = gemv_f16_k_lcpp_src();
for frag in [
"@group(0) @binding(3) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)",
" let row0 = wid.x * 4u;\n let col = wid.y;",
" if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }",
] {
assert!(src.contains(frag), "f16 lcpp source drifted: {frag}");
}
src.replace(
"@group(0) @binding(3) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)",
"@group(0) @binding(3) var<uniform> dims: vec4<u32>; // (M, N, ncols, gy_rows)",
)
.replace(
" let row0 = wid.x * 4u;\n let col = wid.y;",
" let col = wid.y / dims.w;\n let row0 = ((wid.y % dims.w) * 32768u + wid.x) * 4u;",
)
.replace(
" if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }",
" y[yo] = tot[r];",
)
}
pub fn mlp_gate_f16_lcpp_src() -> String {
r#"
fn f16x4(a: u32, b: u32) -> vec4<f32> {
let lo = unpack2x16float(a);
let hi = unpack2x16float(b);
return vec4<f32>(lo.x, lo.y, hi.x, hi.y);
}
@group(0) @binding(0) var<storage, read> w1: array<vec4<u32>>; // gate, f16 (8 per vec4)
@group(0) @binding(1) var<storage, read> w3: array<vec4<u32>>; // up, f16 (8 per vec4)
@group(0) @binding(2) var<storage, read> x: array<vec4<f32>>; // [ncols, N/4] PRE-NORMED
@group(0) @binding(3) var<storage, read_write> y: array<f32>; // [ncols, M]
@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, _, ncols)
@group(0) @binding(5) var<uniform> epsm: vec4<f32>; // (_, gelu-flag, _, _)
@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 ag = vec4<f32>(0.0);
var au = 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 = min(row0 + r, m - 1u);
let wb = (row * nblk + b) * 4u;
{
let q0 = w1[wb]; let q1 = w1[wb + 1u];
let q2 = w1[wb + 2u]; let q3 = w1[wb + 3u];
var sv = dot(f16x4(q0.x, q0.y), v0) + dot(f16x4(q0.z, q0.w), v1);
sv = sv + dot(f16x4(q1.x, q1.y), v2) + dot(f16x4(q1.z, q1.w), v3);
sv = sv + dot(f16x4(q2.x, q2.y), v4) + dot(f16x4(q2.z, q2.w), v5);
sv = sv + dot(f16x4(q3.x, q3.y), v6) + dot(f16x4(q3.z, q3.w), v7);
ag[r] = ag[r] + sv;
}
{
let q0 = w3[wb]; let q1 = w3[wb + 1u];
let q2 = w3[wb + 2u]; let q3 = w3[wb + 3u];
var sv = dot(f16x4(q0.x, q0.y), v0) + dot(f16x4(q0.z, q0.w), v1);
sv = sv + dot(f16x4(q1.x, q1.y), v2) + dot(f16x4(q1.z, q1.w), v3);
sv = sv + dot(f16x4(q2.x, q2.y), v4) + dot(f16x4(q2.z, q2.w), v5);
sv = sv + dot(f16x4(q3.x, q3.y), v6) + dot(f16x4(q3.z, q3.w), v7);
au[r] = au[r] + sv;
}
}
}
let tg = subgroupAdd(ag);
let tu = subgroupAdd(au);
if (sid == 0u) {
for (var r = 0u; r < 4u; r = r + 1u) {
if (row0 + r < m) {
let gate = tg[r];
let upv = tu[r];
var act: f32;
if (epsm.y != 0.0) {
let g3 = gate * gate * gate;
let targ = clamp(0.7978845608028654 * (gate + 0.044715 * g3), -20.0, 20.0);
act = 0.5 * gate * (1.0 + tanh(targ));
} else {
act = gate / (1.0 + exp(-gate));
}
y[col * m + row0 + r] = act * upv;
}
}
}
}
"#
.to_string()
}
pub fn gemv_q1_k_lcpp_src() -> String {
r#"enable f16;
fn q1s(word: u32, sh: u32) -> vec4<f32> {
let bits = (vec4<u32>(word) >> vec4<u32>(sh, sh + 1u, sh + 2u, sh + 3u)) & vec4<u32>(1u);
return select(vec4<f32>(-1.0), vec4<f32>(1.0), bits == vec4<u32>(1u));
}
@group(0) @binding(0) var<storage, read> scales: array<f16>; // [m, n/128]
@group(0) @binding(1) var<storage, read> bits: array<u32>; // [m, n/32] sign words
@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>; // (M, N, acc, ncols)
@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 nsc = n / 128u;
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 = min(row0 + r, m - 1u);
let d = f32(scales[row * nsc + (b >> 2u)]);
let w = bits[row * nblk + b];
var s = dot(q1s(w, 0u), v0) + dot(q1s(w, 4u), v1);
s = s + dot(q1s(w, 8u), v2) + dot(q1s(w, 12u), v3);
s = s + dot(q1s(w, 16u), v4) + dot(q1s(w, 20u), v5);
s = s + dot(q1s(w, 24u), v6) + dot(q1s(w, 28u), 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) {
let yo = col * m + row0 + r;
if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }
}
}
}
}
"#
.to_string()
}
pub fn mlp_gate_q1_f16x_src() -> String {
let src = mlp_gate_q1_lcpp_src();
let out = src
.replace(
"@group(0) @binding(4) var<storage, read> x: array<vec4<f32>>; // [ncols, N/4] PRE-NORMED",
"@group(0) @binding(4) var<storage, read> x: array<vec4<u32>>; // [ncols, N/8] PACKED f16",
)
.replace(" let xstride = n / 4u;", " let xstride = n / 8u;")
.replace(
r#" 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];"#,
r#" let xb = xoff + b * 4u;
let q0 = x[xb]; let q1 = x[xb + 1u];
let q2 = x[xb + 2u]; let q3 = x[xb + 3u];
let v0 = vec4<f32>(unpack2x16float(q0.x), unpack2x16float(q0.y));
let v1 = vec4<f32>(unpack2x16float(q0.z), unpack2x16float(q0.w));
let v2 = vec4<f32>(unpack2x16float(q1.x), unpack2x16float(q1.y));
let v3 = vec4<f32>(unpack2x16float(q1.z), unpack2x16float(q1.w));
let v4 = vec4<f32>(unpack2x16float(q2.x), unpack2x16float(q2.y));
let v5 = vec4<f32>(unpack2x16float(q2.z), unpack2x16float(q2.w));
let v6 = vec4<f32>(unpack2x16float(q3.x), unpack2x16float(q3.y));
let v7 = vec4<f32>(unpack2x16float(q3.z), unpack2x16float(q3.w));"#,
);
assert!(
out != src,
"f16x substitution anchors must match mlp_gate_q1_lcpp_src"
);
out
}
pub fn mlp_gate_q1_lcpp_src() -> String {
r#"enable f16;
fn q1s(word: u32, sh: u32) -> vec4<f32> {
let bits = (vec4<u32>(word) >> vec4<u32>(sh, sh + 1u, sh + 2u, sh + 3u)) & vec4<u32>(1u);
return select(vec4<f32>(-1.0), vec4<f32>(1.0), bits == vec4<u32>(1u));
}
@group(0) @binding(0) var<storage, read> s1: array<f16>; // gate scales [M, N/128]
@group(0) @binding(1) var<storage, read> b1: array<u32>; // gate signs [M, N/32]
@group(0) @binding(2) var<storage, read> s3: array<f16>; // up scales
@group(0) @binding(3) var<storage, read> b3: array<u32>; // up signs
@group(0) @binding(4) var<storage, read> x: array<vec4<f32>>; // [ncols, N/4] PRE-NORMED
@group(0) @binding(5) var<storage, read_write> y: array<f32>; // [ncols, M]
@group(0) @binding(6) var<uniform> dims: vec4<u32>; // (M, N, _, ncols)
@group(0) @binding(7) var<uniform> epsm: vec4<f32>; // (_, gelu-flag, _, _)
@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 nsc = n / 128u;
let xstride = n / 4u;
let row0 = wid.x * 4u;
let col = wid.y;
let xoff = col * xstride;
var ag = vec4<f32>(0.0);
var au = 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];
let sblk = b >> 2u;
for (var r = 0u; r < 4u; r = r + 1u) {
let row = min(row0 + r, m - 1u);
{
let d = f32(s1[row * nsc + sblk]);
let w = b1[row * nblk + b];
var sv = dot(q1s(w, 0u), v0) + dot(q1s(w, 4u), v1);
sv = sv + dot(q1s(w, 8u), v2) + dot(q1s(w, 12u), v3);
sv = sv + dot(q1s(w, 16u), v4) + dot(q1s(w, 20u), v5);
sv = sv + dot(q1s(w, 24u), v6) + dot(q1s(w, 28u), v7);
ag[r] = ag[r] + d * sv;
}
{
let d = f32(s3[row * nsc + sblk]);
let w = b3[row * nblk + b];
var sv = dot(q1s(w, 0u), v0) + dot(q1s(w, 4u), v1);
sv = sv + dot(q1s(w, 8u), v2) + dot(q1s(w, 12u), v3);
sv = sv + dot(q1s(w, 16u), v4) + dot(q1s(w, 20u), v5);
sv = sv + dot(q1s(w, 24u), v6) + dot(q1s(w, 28u), v7);
au[r] = au[r] + d * sv;
}
}
}
let tg = subgroupAdd(ag);
let tu = subgroupAdd(au);
if (sid == 0u) {
for (var r = 0u; r < 4u; r = r + 1u) {
if (row0 + r < m) {
let gate = tg[r];
let upv = tu[r];
var act: f32;
if (epsm.y != 0.0) {
let g3 = gate * gate * gate;
let targ = clamp(0.7978845608028654 * (gate + 0.044715 * g3), -20.0, 20.0);
act = 0.5 * gate * (1.0 + tanh(targ));
} else {
act = gate / (1.0 + exp(-gate));
}
y[col * m + row0 + r] = act * upv;
}
}
}
}
"#
.to_string()
}
pub(crate) fn gemv_q4_k_nbar_src() -> String {
format!(
r#"{Q4_FN}
@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>; // (M, N, acc, ncols)
const KC: u32 = 4u;
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;
var acc = vec4<f32>(0.0);
for (var b = lane; b < nblk; b = b + 16u) {{
var d = 0.0;
var q = vec4<u32>();
if (row < m) {{
d = f32(scales[row * nblk + b]);
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;
}}
}}
}}
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;
if (dims.z == 1u) {{ y[yo] = y[yo] + v; }} else {{ y[yo] = v; }}
}}
workgroupBarrier();
}}
}}
"#
)
}
pub fn gemv_q4_k_sg_src() -> String {
gemv_sg_generic(4)
}
pub(crate) fn gemv_q4_k_sg16_src() -> String {
gemv_sg_generic(16)
}
pub fn gemv_q4_k_sg_solo_src() -> String {
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 fma(vec4<f32>(unpack4xU8(word & 0xF0F0F0F0u)), vec4<f32>(0.0625), vec4<f32>(-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>; // (M, N, acc, ncols=1)
@compute @workgroup_size(256)
fn main(
@builtin(workgroup_id) wid: vec3<u32>,
@builtin(local_invocation_id) lid3: vec3<u32>,
@builtin(subgroup_invocation_id) sid: u32,
@builtin(subgroup_size) ssz: u32,
) {
let nsg = 256u / ssz;
let sg = lid3.x / ssz;
let m = dims.x; let n = dims.y;
let row_raw = wid.x * nsg + sg;
let row = min(row_raw, m - 1u);
let nblk = n / 32u;
var acc = 0.0;
for (var b = sid; b < nblk; b = b + ssz) {
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 = b * 8u;
var s = dot(l0, x[xb]) + dot(h0, x[xb + 4u]);
s = s + dot(l1, x[xb + 1u]) + dot(h1, x[xb + 5u]);
s = s + dot(l2, x[xb + 2u]) + dot(h2, x[xb + 6u]);
s = s + dot(l3, x[xb + 3u]) + dot(h3, x[xb + 7u]);
acc = acc + d * s;
}
let tot = subgroupAdd(acc);
if (sid == 0u && row_raw < m) {
if (dims.z == 1u) { y[row] = y[row] + tot; } else { y[row] = tot; }
}
}
"#
.to_string()
}
fn gemv_sg_generic(kc: usize) -> String {
format!(
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 fma(vec4<f32>(unpack4xU8(word & 0xF0F0F0F0u)), vec4<f32>(0.0625), vec4<f32>(-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>; // (M, N, acc, ncols)
const KC: u32 = {kc}u;
@compute @workgroup_size(256)
fn main(
@builtin(workgroup_id) wid: vec3<u32>,
@builtin(local_invocation_id) lid3: vec3<u32>,
@builtin(subgroup_invocation_id) sid: u32,
@builtin(subgroup_size) ssz: u32,
) {{
let nsg = 256u / ssz;
let sg = lid3.x / ssz;
let m = dims.x; let n = dims.y; let ncols = dims.w;
let row_raw = wid.x * nsg + sg;
let row = min(row_raw, m - 1u);
let col0 = wid.y * KC;
let nblk = n / 32u;
let xstride = n / 4u;
var acc = array<f32, {kc}>();
for (var b = sid; b < nblk; b = b + ssz) {{
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 = 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;
}}
}}
}}
var tot = array<f32, {kc}>();
for (var cc = 0u; cc < KC; cc = cc + 1u) {{ tot[cc] = subgroupAdd(acc[cc]); }}
if (sid == 0u && row_raw < m) {{
for (var cc = 0u; cc < KC; cc = cc + 1u) {{
if (col0 + cc < ncols) {{
let yo = (col0 + cc) * m + row;
if (dims.z == 1u) {{ y[yo] = y[yo] + tot[cc]; }} else {{ y[yo] = tot[cc]; }}
}}
}}
}}
}}
"#
)
}
pub fn gemv_q8_k_sg_solo_src() -> String {
r#"enable f16;
fn q8b(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xI8(word)); }
@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>; // (M, N, acc, ncols=1)
@compute @workgroup_size(256)
fn main(
@builtin(workgroup_id) wid: vec3<u32>,
@builtin(local_invocation_id) lid3: vec3<u32>,
@builtin(subgroup_invocation_id) sid: u32,
@builtin(subgroup_size) ssz: u32,
) {
let nsg = 256u / ssz;
let sg = lid3.x / ssz;
let m = dims.x; let n = dims.y;
let row_raw = wid.x * nsg + sg;
let row = min(row_raw, m - 1u);
let nblk = n / 32u;
var acc = 0.0;
for (var b = sid; b < nblk; b = b + ssz) {
let d = f32(scales[row * nblk + b]);
let qa = quants[(row * nblk + b) * 2u];
let qb = quants[(row * nblk + b) * 2u + 1u];
let v0 = q8b(qa.x); let v1 = q8b(qa.y); let v2 = q8b(qa.z); let v3 = q8b(qa.w);
let v4 = q8b(qb.x); let v5 = q8b(qb.y); let v6 = q8b(qb.z); let v7 = q8b(qb.w);
let xb = b * 8u;
var s = dot(v0, x[xb]) + dot(v1, x[xb + 1u]) + dot(v2, x[xb + 2u]) + dot(v3, x[xb + 3u]);
s = s + dot(v4, x[xb + 4u]) + dot(v5, x[xb + 5u]) + dot(v6, x[xb + 6u]) + dot(v7, x[xb + 7u]);
acc = acc + d * s;
}
let tot = subgroupAdd(acc);
if (sid == 0u && row_raw < m) {
if (dims.z == 1u) { y[row] = y[row] + tot; } else { y[row] = tot; }
}
}
"#
.to_string()
}
pub fn gemv_q8_k_sg_src() -> String {
r#"enable f16;
fn q8b(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xI8(word)); }
@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>; // (M, N, acc, ncols)
const KC: u32 = 4u;
@compute @workgroup_size(256)
fn main(
@builtin(workgroup_id) wid: vec3<u32>,
@builtin(local_invocation_id) lid3: vec3<u32>,
@builtin(subgroup_invocation_id) sid: u32,
@builtin(subgroup_size) ssz: u32,
) {
let nsg = 256u / ssz;
let sg = lid3.x / ssz;
let m = dims.x; let n = dims.y; let ncols = dims.w;
let row_raw = wid.x * nsg + sg;
let row = min(row_raw, m - 1u);
let col0 = wid.y * KC;
let nblk = n / 32u;
let xstride = n / 4u;
var acc = array<f32, 4>();
for (var b = sid; b < nblk; b = b + ssz) {
let d = f32(scales[row * nblk + b]);
let qa = quants[(row * nblk + b) * 2u];
let qb = quants[(row * nblk + b) * 2u + 1u];
let v0 = q8b(qa.x); let v1 = q8b(qa.y); let v2 = q8b(qa.z); let v3 = q8b(qa.w);
let v4 = q8b(qb.x); let v5 = q8b(qb.y); let v6 = q8b(qb.z); let v7 = q8b(qb.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(v0, x[base]) + dot(v1, x[base + 1u]) + dot(v2, x[base + 2u]) + dot(v3, x[base + 3u]);
s = s + dot(v4, x[base + 4u]) + dot(v5, x[base + 5u]) + dot(v6, x[base + 6u]) + dot(v7, x[base + 7u]);
acc[cc] = acc[cc] + d * s;
}
}
}
var tot = array<f32, 4>();
for (var cc = 0u; cc < KC; cc = cc + 1u) { tot[cc] = subgroupAdd(acc[cc]); }
if (sid == 0u && row_raw < m) {
for (var cc = 0u; cc < KC; cc = cc + 1u) {
if (col0 + cc < ncols) {
let yo = (col0 + cc) * m + row;
if (dims.z == 1u) { y[yo] = y[yo] + tot[cc]; } else { y[yo] = tot[cc]; }
}
}
}
}
"#
.to_string()
}
pub fn gemv_q8_k_src() -> String {
r#"enable f16;
fn q8b(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xI8(word)); }
@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>; // (M, N, acc, ncols)
const KC: u32 = 4u;
const ROWS: u32 = 16u;
const LANES: u32 = 16u;
var<workgroup> red: array<vec4<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 rl = lid / LANES;
let lane = lid % LANES;
let m = dims.x; let n = dims.y; let ncols = dims.w;
let row_raw = wid.x * ROWS + rl;
let row = min(row_raw, m - 1u);
let col0 = wid.y * KC;
let nblk = n / 32u;
let xstride = n / 4u;
var acc = vec4<f32>(0.0);
for (var b = lane; b < nblk; b = b + LANES) {
let d = f32(scales[row * nblk + b]);
let qa = quants[(row * nblk + b) * 2u];
let qb = quants[(row * nblk + b) * 2u + 1u];
let v0 = q8b(qa.x); let v1 = q8b(qa.y); let v2 = q8b(qa.z); let v3 = q8b(qa.w);
let v4 = q8b(qb.x); let v5 = q8b(qb.y); let v6 = q8b(qb.z); let v7 = q8b(qb.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(v0, x[base]) + dot(v1, x[base + 1u]) + dot(v2, x[base + 2u]) + dot(v3, x[base + 3u]);
s = s + dot(v4, x[base + 4u]) + dot(v5, x[base + 5u]) + dot(v6, x[base + 6u]) + dot(v7, x[base + 7u]);
acc[cc] = acc[cc] + d * s;
}
}
}
red[lid] = acc;
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_raw < m) {
let tot = red[lid] + red[lid + 1u];
for (var cc = 0u; cc < KC; cc = cc + 1u) {
if (col0 + cc < ncols) {
let yo = (col0 + cc) * m + row;
if (dims.z == 1u) { y[yo] = y[yo] + tot[cc]; } else { y[yo] = tot[cc]; }
}
}
}
}
"#
.to_string()
}
pub(crate) fn mlp_gate_q8_k_src() -> String {
r#"enable f16;
fn q8b(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xI8(word)); }
@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> wn: array<vec4<f32>>;
@group(0) @binding(6) var<storage, read_write> y: array<f32>;
@group(0) @binding(7) var<uniform> dims: vec4<u32>; // (M, N, _, ncols)
@group(0) @binding(8) var<uniform> epsm: vec4<f32>;
const KC: u32 = 4u;
const ROWS: u32 = 16u;
const LANES: u32 = 16u;
var<workgroup> redg: array<vec4<f32>, 256>;
var<workgroup> redu: array<vec4<f32>, 256>;
var<workgroup> reds: array<vec4<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 rl = lid / LANES;
let lane = lid % LANES;
let m = dims.x; let n = dims.y; let ncols = dims.w;
let row_raw = wid.x * ROWS + rl;
let row = min(row_raw, m - 1u);
let col0 = wid.y * KC;
let nblk = n / 32u;
let xstride = n / 4u;
var ag = vec4<f32>(0.0);
var au = vec4<f32>(0.0);
var sq = vec4<f32>(0.0);
for (var b = lane; b < nblk; b = b + LANES) {
let d1 = f32(s1[row * nblk + b]);
let d3 = f32(s3[row * nblk + b]);
let qa = q1[(row * nblk + b) * 2u];
let qa2 = q1[(row * nblk + b) * 2u + 1u];
let qb = q3[(row * nblk + b) * 2u];
let qb2 = q3[(row * nblk + b) * 2u + 1u];
let a0 = q8b(qa.x); let a1 = q8b(qa.y); let a2 = q8b(qa.z); let a3 = q8b(qa.w);
let a4 = q8b(qa2.x); let a5 = q8b(qa2.y); let a6 = q8b(qa2.z); let a7 = q8b(qa2.w);
let b0 = q8b(qb.x); let b1 = q8b(qb.y); let b2 = q8b(qb.z); let b3 = q8b(qb.w);
let b4 = q8b(qb2.x); let b5 = q8b(qb2.y); let b6 = q8b(qb2.z); let b7 = q8b(qb2.w);
let xb = b * 8u;
for (var cc = 0u; cc < KC; cc = cc + 1u) {
if (col0 + cc < ncols) {
let base = (col0 + cc) * xstride + xb;
let v0 = x[base] * wn[xb]; let v4 = x[base + 4u] * wn[xb + 4u];
let v1 = x[base + 1u] * wn[xb + 1u]; let v5 = x[base + 5u] * wn[xb + 5u];
let v2 = x[base + 2u] * wn[xb + 2u]; let v6 = x[base + 6u] * wn[xb + 6u];
let v3 = x[base + 3u] * wn[xb + 3u]; let v7 = x[base + 7u] * wn[xb + 7u];
var g = dot(a0, v0) + dot(a4, v4) + dot(a1, v1) + dot(a5, v5);
g = g + dot(a2, v2) + dot(a6, v6) + dot(a3, v3) + dot(a7, v7);
var u = dot(b0, v0) + dot(b4, v4) + dot(b1, v1) + dot(b5, v5);
u = u + dot(b2, v2) + dot(b6, v6) + dot(b3, v3) + dot(b7, v7);
ag[cc] = ag[cc] + d1 * g;
au[cc] = au[cc] + d3 * u;
let r0 = x[base]; let r1 = x[base + 1u]; let r2 = x[base + 2u]; let r3 = x[base + 3u];
let r4 = x[base + 4u]; let r5 = x[base + 5u]; let r6 = x[base + 6u]; let r7 = x[base + 7u];
sq[cc] = sq[cc] + dot(r0, r0) + dot(r1, r1) + dot(r2, r2) + dot(r3, r3)
+ dot(r4, r4) + dot(r5, r5) + dot(r6, r6) + dot(r7, r7);
}
}
}
redg[lid] = ag; redu[lid] = au; reds[lid] = sq;
workgroupBarrier();
if (lane < 8u) { redg[lid] += redg[lid + 8u]; redu[lid] += redu[lid + 8u]; reds[lid] += reds[lid + 8u]; }
workgroupBarrier();
if (lane < 4u) { redg[lid] += redg[lid + 4u]; redu[lid] += redu[lid + 4u]; reds[lid] += reds[lid + 4u]; }
workgroupBarrier();
if (lane < 2u) { redg[lid] += redg[lid + 2u]; redu[lid] += redu[lid + 2u]; reds[lid] += reds[lid + 2u]; }
workgroupBarrier();
if (lane == 0u && row_raw < m) {
let tg = redg[lid] + redg[lid + 1u];
let tu = redu[lid] + redu[lid + 1u];
let ts = reds[lid] + reds[lid + 1u];
for (var cc = 0u; cc < KC; cc = cc + 1u) {
if (col0 + cc < ncols) {
let inv = 1.0 / sqrt(ts[cc] / f32(n) + epsm.x);
let gate = inv * tg[cc];
let upv = inv * tu[cc];
var act: f32;
if (epsm.y != 0.0) {
let g3 = gate * gate * gate;
let targ = clamp(0.7978845608028654 * (gate + 0.044715 * g3), -20.0, 20.0);
act = 0.5 * gate * (1.0 + tanh(targ));
} else {
act = gate / (1.0 + exp(-gate));
}
y[(col0 + cc) * m + row] = act * upv;
}
}
}
}
"#
.to_string()
}
pub(crate) fn mlp_gate_q8_k_sg_src() -> String {
r#"enable f16;
fn q8b(word: u32) -> vec4<f32> { return vec4<f32>(unpack4xI8(word)); }
@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> wn: array<vec4<f32>>;
@group(0) @binding(6) var<storage, read_write> y: array<f32>;
@group(0) @binding(7) var<uniform> dims: vec4<u32>; // (M, N, _, ncols)
@group(0) @binding(8) var<uniform> epsm: vec4<f32>;
const KC: u32 = 4u;
@compute @workgroup_size(256)
fn main(
@builtin(workgroup_id) wid: vec3<u32>,
@builtin(local_invocation_id) lid3: vec3<u32>,
@builtin(subgroup_invocation_id) sid: u32,
@builtin(subgroup_size) ssz: u32,
) {
let nsg = 256u / ssz;
let sg = lid3.x / ssz;
let m = dims.x; let n = dims.y; let ncols = dims.w;
let row_raw = wid.x * nsg + sg;
let row = min(row_raw, m - 1u);
let col0 = wid.y * KC;
let nblk = n / 32u;
let xstride = n / 4u;
var ag = vec4<f32>(0.0);
var au = vec4<f32>(0.0);
var sq = vec4<f32>(0.0);
for (var b = sid; b < nblk; b = b + ssz) {
let d1 = f32(s1[row * nblk + b]);
let d3 = f32(s3[row * nblk + b]);
let qa = q1[(row * nblk + b) * 2u];
let qa2 = q1[(row * nblk + b) * 2u + 1u];
let qb = q3[(row * nblk + b) * 2u];
let qb2 = q3[(row * nblk + b) * 2u + 1u];
let a0 = q8b(qa.x); let a1 = q8b(qa.y); let a2 = q8b(qa.z); let a3 = q8b(qa.w);
let a4 = q8b(qa2.x); let a5 = q8b(qa2.y); let a6 = q8b(qa2.z); let a7 = q8b(qa2.w);
let b0 = q8b(qb.x); let b1 = q8b(qb.y); let b2 = q8b(qb.z); let b3 = q8b(qb.w);
let b4 = q8b(qb2.x); let b5 = q8b(qb2.y); let b6 = q8b(qb2.z); let b7 = q8b(qb2.w);
let xb = b * 8u;
for (var cc = 0u; cc < KC; cc = cc + 1u) {
if (col0 + cc < ncols) {
let base = (col0 + cc) * xstride + xb;
let v0 = x[base] * wn[xb]; let v4 = x[base + 4u] * wn[xb + 4u];
let v1 = x[base + 1u] * wn[xb + 1u]; let v5 = x[base + 5u] * wn[xb + 5u];
let v2 = x[base + 2u] * wn[xb + 2u]; let v6 = x[base + 6u] * wn[xb + 6u];
let v3 = x[base + 3u] * wn[xb + 3u]; let v7 = x[base + 7u] * wn[xb + 7u];
var g = dot(a0, v0) + dot(a4, v4) + dot(a1, v1) + dot(a5, v5);
g = g + dot(a2, v2) + dot(a6, v6) + dot(a3, v3) + dot(a7, v7);
var u = dot(b0, v0) + dot(b4, v4) + dot(b1, v1) + dot(b5, v5);
u = u + dot(b2, v2) + dot(b6, v6) + dot(b3, v3) + dot(b7, v7);
ag[cc] = ag[cc] + d1 * g;
au[cc] = au[cc] + d3 * u;
let r0 = x[base]; let r1 = x[base + 1u]; let r2 = x[base + 2u]; let r3 = x[base + 3u];
let r4 = x[base + 4u]; let r5 = x[base + 5u]; let r6 = x[base + 6u]; let r7 = x[base + 7u];
sq[cc] = sq[cc] + dot(r0, r0) + dot(r1, r1) + dot(r2, r2) + dot(r3, r3)
+ dot(r4, r4) + dot(r5, r5) + dot(r6, r6) + dot(r7, r7);
}
}
}
let tg = subgroupAdd(ag);
let tu = subgroupAdd(au);
let ts = subgroupAdd(sq);
if (sid == 0u && row_raw < m) {
for (var cc = 0u; cc < KC; cc = cc + 1u) {
if (col0 + cc < ncols) {
let inv = 1.0 / sqrt(ts[cc] / f32(n) + epsm.x);
let gate = inv * tg[cc];
let upv = inv * tu[cc];
var act: f32;
if (epsm.y != 0.0) {
let g3 = gate * gate * gate;
let targ = clamp(0.7978845608028654 * (gate + 0.044715 * g3), -20.0, 20.0);
act = 0.5 * gate * (1.0 + tanh(targ));
} else {
act = gate / (1.0 + exp(-gate));
}
y[(col0 + cc) * m + row] = act * upv;
}
}
}
}
"#
.to_string()
}
pub(crate) fn q4_gemv_norm32_k_sg_src() -> String {
gn32_sg_generic(4)
}
#[allow(dead_code)]
pub(crate) fn q4_gemv_norm32_k_sg16_src() -> String {
gn32_sg_generic(16)
}
fn gn32_sg_generic(kc: usize) -> String {
format!(
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 fma(vec4<f32>(unpack4xU8(word & 0xF0F0F0F0u)), vec4<f32>(0.0625), vec4<f32>(-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> wn: array<vec4<f32>>;
@group(0) @binding(4) var<storage, read_write> y: array<f32>;
@group(0) @binding(5) var<uniform> dims: vec4<u32>; // (M, N, _, ncols)
@group(0) @binding(6) var<uniform> epsm: vec4<f32>;
const KC: u32 = {kc}u;
@compute @workgroup_size(256)
fn main(
@builtin(workgroup_id) wid: vec3<u32>,
@builtin(local_invocation_id) lid3: vec3<u32>,
@builtin(subgroup_invocation_id) sid: u32,
@builtin(subgroup_size) ssz: u32,
) {{
let nsg = 256u / ssz;
let sg = lid3.x / ssz;
let m = dims.x; let n = dims.y; let ncols = dims.w;
let row_raw = wid.x * nsg + sg;
let row = min(row_raw, m - 1u);
let col0 = wid.y * KC;
let nblk = n / 32u;
let xstride = n / 4u;
var acc = array<f32, {kc}>();
var sq = array<f32, {kc}>();
for (var b = sid; b < nblk; b = b + ssz) {{
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 = b * 8u;
for (var cc = 0u; cc < KC; cc = cc + 1u) {{
if (col0 + cc < ncols) {{
let base = (col0 + cc) * xstride + xb;
var s = 0.0;
var qq = 0.0;
for (var i = 0u; i < 8u; i = i + 1u) {{
let xv = x[base + i];
qq = qq + dot(xv, xv);
let nv = xv * wn[xb + i];
switch i {{
case 0u: {{ s = s + dot(l0, nv); }}
case 1u: {{ s = s + dot(l1, nv); }}
case 2u: {{ s = s + dot(l2, nv); }}
case 3u: {{ s = s + dot(l3, nv); }}
case 4u: {{ s = s + dot(h0, nv); }}
case 5u: {{ s = s + dot(h1, nv); }}
case 6u: {{ s = s + dot(h2, nv); }}
default: {{ s = s + dot(h3, nv); }}
}}
}}
acc[cc] = acc[cc] + d * s;
sq[cc] = sq[cc] + qq;
}}
}}
}}
var tot = array<f32, {kc}>();
var sqt = array<f32, {kc}>();
for (var cc = 0u; cc < KC; cc = cc + 1u) {{
tot[cc] = subgroupAdd(acc[cc]);
sqt[cc] = subgroupAdd(sq[cc]);
}}
if (sid == 0u && row_raw < m) {{
for (var cc = 0u; cc < KC; cc = cc + 1u) {{
if (col0 + cc < ncols) {{
let inv = 1.0 / sqrt(sqt[cc] / f32(n) + epsm.x);
y[(col0 + cc) * m + row] = inv * tot[cc];
}}
}}
}}
}}
"#
)
}
pub(crate) fn mlp_gate_q4_k_sg_src() -> String {
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 fma(vec4<f32>(unpack4xU8(word & 0xF0F0F0F0u)), vec4<f32>(0.0625), vec4<f32>(-8.0)); }
@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> wn: array<vec4<f32>>;
@group(0) @binding(6) var<storage, read_write> y: array<f32>;
@group(0) @binding(7) var<uniform> dims: vec4<u32>; // (M, N, _, ncols)
@group(0) @binding(8) var<uniform> epsm: vec4<f32>;
const KC: u32 = 4u;
@compute @workgroup_size(256)
fn main(
@builtin(workgroup_id) wid: vec3<u32>,
@builtin(local_invocation_id) lid3: vec3<u32>,
@builtin(subgroup_invocation_id) sid: u32,
@builtin(subgroup_size) ssz: u32,
) {
let nsg = 256u / ssz;
let sg = lid3.x / ssz;
let m = dims.x; let n = dims.y; let ncols = dims.w;
let row_raw = wid.x * nsg + sg;
let row = min(row_raw, m - 1u);
let col0 = wid.y * KC;
let nblk = n / 32u;
let xstride = n / 4u;
var ag = vec4<f32>(0.0);
var au = vec4<f32>(0.0);
var sq = vec4<f32>(0.0);
for (var b = sid; b < nblk; b = b + ssz) {
let d1 = f32(s1[row * nblk + b]);
let d3 = f32(s3[row * nblk + b]);
let qa = q1[row * nblk + b];
let qb = q3[row * nblk + b];
let a0 = q4_lo(qa.x); let a4 = q4_hi(qa.x);
let a1 = q4_lo(qa.y); let a5 = q4_hi(qa.y);
let a2 = q4_lo(qa.z); let a6 = q4_hi(qa.z);
let a3 = q4_lo(qa.w); let a7 = q4_hi(qa.w);
let b0 = q4_lo(qb.x); let b4 = q4_hi(qb.x);
let b1 = q4_lo(qb.y); let b5 = q4_hi(qb.y);
let b2 = q4_lo(qb.z); let b6 = q4_hi(qb.z);
let b3 = q4_lo(qb.w); let b7 = q4_hi(qb.w);
let xb = b * 8u;
for (var cc = 0u; cc < KC; cc = cc + 1u) {
if (col0 + cc < ncols) {
let base = (col0 + cc) * xstride + xb;
let v0 = x[base] * wn[xb]; let v4 = x[base + 4u] * wn[xb + 4u];
let v1 = x[base + 1u] * wn[xb + 1u]; let v5 = x[base + 5u] * wn[xb + 5u];
let v2 = x[base + 2u] * wn[xb + 2u]; let v6 = x[base + 6u] * wn[xb + 6u];
let v3 = x[base + 3u] * wn[xb + 3u]; let v7 = x[base + 7u] * wn[xb + 7u];
var g = dot(a0, v0) + dot(a4, v4) + dot(a1, v1) + dot(a5, v5);
g = g + dot(a2, v2) + dot(a6, v6) + dot(a3, v3) + dot(a7, v7);
var u = dot(b0, v0) + dot(b4, v4) + dot(b1, v1) + dot(b5, v5);
u = u + dot(b2, v2) + dot(b6, v6) + dot(b3, v3) + dot(b7, v7);
ag[cc] = ag[cc] + d1 * g;
au[cc] = au[cc] + d3 * u;
let r0 = x[base]; let r1 = x[base + 1u]; let r2 = x[base + 2u]; let r3 = x[base + 3u];
let r4 = x[base + 4u]; let r5 = x[base + 5u]; let r6 = x[base + 6u]; let r7 = x[base + 7u];
sq[cc] = sq[cc] + dot(r0, r0) + dot(r1, r1) + dot(r2, r2) + dot(r3, r3)
+ dot(r4, r4) + dot(r5, r5) + dot(r6, r6) + dot(r7, r7);
}
}
}
let tg = subgroupAdd(ag);
let tu = subgroupAdd(au);
let ts = subgroupAdd(sq);
if (sid == 0u && row_raw < m) {
for (var cc = 0u; cc < KC; cc = cc + 1u) {
if (col0 + cc < ncols) {
let inv = 1.0 / sqrt(ts[cc] / f32(n) + epsm.x);
let gate = inv * tg[cc];
let upv = inv * tu[cc];
var act: f32;
if (epsm.y != 0.0) {
let g3 = gate * gate * gate;
let targ = clamp(0.7978845608028654 * (gate + 0.044715 * g3), -20.0, 20.0);
act = 0.5 * gate * (1.0 + tanh(targ));
} else {
act = gate / (1.0 + exp(-gate));
}
y[(col0 + cc) * m + row] = act * upv;
}
}
}
}
"#
.to_string()
}
pub(crate) fn q4_gemv_norm32_k_src() -> String {
format!(
r#"{Q4_FN}
@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>>; // [ncols, N/4]
@group(0) @binding(3) var<storage, read> wn: array<vec4<f32>>;
@group(0) @binding(4) var<storage, read_write> y: array<f32>; // [ncols, M]
@group(0) @binding(5) var<uniform> dims: vec4<u32>; // (M, N, _, ncols)
@group(0) @binding(6) var<uniform> epsm: vec4<f32>;
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>; // KC·TV, PRE-normed (x·wn)
var<workgroup> red: array<f32, 256>;
var<workgroup> sqc: array<f32, 4>; // per-column Σx² (raw x)
@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);
var sq = vec4<f32>(0.0); // per-thread partial Σx² per column
let ntiles = (nblk + TB - 1u) / TB;
var d_cur = 0.0;
var q_cur = vec4<u32>();
if (row < m && lane < nblk) {{
d_cur = f32(scales[row * nblk + lane]);
q_cur = quants[row * nblk + lane];
}}
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);
var w = vec4<f32>(0.0);
if (col0 + cc < ncols && e < xstride) {{ v = x[(col0 + cc) * xstride + e]; w = wn[e]; }}
sq[cc] = sq[cc] + dot(v, v);
xs[idx] = v * w;
}}
workgroupBarrier();
// Software pipeline: tile t+1's weights are ISSUED here, before tile t's dots —
// the DRAM latency hides behind the arithmetic (same FP order, bitwise-identical).
let bn = (t + 1u) * TB + lane;
var d_nxt = 0.0;
var q_nxt = vec4<u32>();
if (t + 1u < ntiles && row < m && bn < nblk) {{
d_nxt = f32(scales[row * nblk + bn]);
q_nxt = quants[row * nblk + bn];
}}
let b = t * TB + lane;
if (row < m && b < nblk) {{
let d = d_cur;
let q = q_cur;
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;
}}
}}
d_cur = d_nxt;
q_cur = q_nxt;
workgroupBarrier();
}}
// Per-column Σx²: tree over all 256 threads (each element loaded exactly once).
for (var cc = 0u; cc < KC; cc = cc + 1u) {{
red[lid] = sq[cc];
workgroupBarrier();
for (var st = 128u; st > 0u; st = st >> 1u) {{
if (lid < st) {{ red[lid] = red[lid] + red[lid + st]; }}
workgroupBarrier();
}}
if (lid == 0u) {{ sqc[cc] = red[0]; }}
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 inv = 1.0 / sqrt(sqc[cc] / f32(n) + epsm.x);
y[(col0 + cc) * m + row] = inv * (red[lid] + red[lid + 1u]);
}}
workgroupBarrier();
}}
}}
"#
)
}
pub(crate) fn mlp_gate_q4_lcpp_src() -> String {
format!(
r#"{Q4_FN}
@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>>; // [ncols, N/4] PRE-NORMED
@group(0) @binding(5) var<storage, read_write> y: array<f32>; // [ncols, M]
@group(0) @binding(6) var<uniform> dims: vec4<u32>; // (M, N, _, ncols)
@group(0) @binding(7) var<uniform> epsm: vec4<f32>; // (_, gelu-flag, _, _)
@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 ag = vec4<f32>(0.0);
var au = 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 = min(row0 + r, m - 1u);
{{
let d = f32(s1[row * nblk + b]);
let q = q1[row * nblk + b];
var sv = dot(q4_lo(q.x), v0) + dot(q4_hi(q.x), v4);
sv = sv + dot(q4_lo(q.y), v1) + dot(q4_hi(q.y), v5);
sv = sv + dot(q4_lo(q.z), v2) + dot(q4_hi(q.z), v6);
sv = sv + dot(q4_lo(q.w), v3) + dot(q4_hi(q.w), v7);
ag[r] = ag[r] + d * sv;
}}
{{
let d = f32(s3[row * nblk + b]);
let q = q3[row * nblk + b];
var sv = dot(q4_lo(q.x), v0) + dot(q4_hi(q.x), v4);
sv = sv + dot(q4_lo(q.y), v1) + dot(q4_hi(q.y), v5);
sv = sv + dot(q4_lo(q.z), v2) + dot(q4_hi(q.z), v6);
sv = sv + dot(q4_lo(q.w), v3) + dot(q4_hi(q.w), v7);
au[r] = au[r] + d * sv;
}}
}}
}}
let tg = subgroupAdd(ag);
let tu = subgroupAdd(au);
if (sid == 0u) {{
for (var r = 0u; r < 4u; r = r + 1u) {{
if (row0 + r < m) {{
let gate = tg[r];
let upv = tu[r];
var act: f32;
if (epsm.y != 0.0) {{
let g3 = gate * gate * gate;
let targ = clamp(0.7978845608028654 * (gate + 0.044715 * g3), -20.0, 20.0);
act = 0.5 * gate * (1.0 + tanh(targ));
}} else {{
act = gate / (1.0 + exp(-gate));
}}
y[col * m + row0 + r] = act * upv;
}}
}}
}}
}}
"#
)
}
pub(crate) fn mlp_gate_q4_k_src() -> String {
format!(
r#"{Q4_FN}
@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>>; // [ncols, N/4]
@group(0) @binding(5) var<storage, read> wn: array<vec4<f32>>;
@group(0) @binding(6) var<storage, read_write> y: array<f32>; // [ncols, M]
@group(0) @binding(7) var<uniform> dims: vec4<u32>; // (M, N, _, ncols)
@group(0) @binding(8) var<uniform> epsm: vec4<f32>;
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>;
var<workgroup> sqc: array<f32, 4>;
@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 ag = vec4<f32>(0.0);
var au = vec4<f32>(0.0);
var sq = 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);
var w = vec4<f32>(0.0);
if (col0 + cc < ncols && e < xstride) {{ v = x[(col0 + cc) * xstride + e]; w = wn[e]; }}
sq[cc] = sq[cc] + dot(v, v);
xs[idx] = v * w;
}}
workgroupBarrier();
let b = t * TB + lane;
if (row < m && b < nblk) {{
let d1 = f32(s1[row * nblk + b]);
let d3 = f32(s3[row * nblk + b]);
let qa = q1[row * nblk + b];
let qb = q3[row * nblk + b];
let a0 = q4_lo(qa.x); let a4 = q4_hi(qa.x);
let a1 = q4_lo(qa.y); let a5 = q4_hi(qa.y);
let a2 = q4_lo(qa.z); let a6 = q4_hi(qa.z);
let a3 = q4_lo(qa.w); let a7 = q4_hi(qa.w);
let b0 = q4_lo(qb.x); let b4 = q4_hi(qb.x);
let b1 = q4_lo(qb.y); let b5 = q4_hi(qb.y);
let b2 = q4_lo(qb.z); let b6 = q4_hi(qb.z);
let b3 = q4_lo(qb.w); let b7 = q4_hi(qb.w);
let xb = lane * 8u;
for (var cc = 0u; cc < KC; cc = cc + 1u) {{
let base = cc * TV + xb;
let v0 = xs[base]; let v4 = xs[base + 4u];
let v1 = xs[base + 1u]; let v5 = xs[base + 5u];
let v2 = xs[base + 2u]; let v6 = xs[base + 6u];
let v3 = xs[base + 3u]; let v7 = xs[base + 7u];
var g = dot(a0, v0) + dot(a4, v4) + dot(a1, v1) + dot(a5, v5);
g = g + dot(a2, v2) + dot(a6, v6) + dot(a3, v3) + dot(a7, v7);
var u = dot(b0, v0) + dot(b4, v4) + dot(b1, v1) + dot(b5, v5);
u = u + dot(b2, v2) + dot(b6, v6) + dot(b3, v3) + dot(b7, v7);
ag[cc] = ag[cc] + d1 * g;
au[cc] = au[cc] + d3 * u;
}}
}}
workgroupBarrier();
}}
for (var cc = 0u; cc < KC; cc = cc + 1u) {{
red[lid] = sq[cc];
workgroupBarrier();
for (var st = 128u; st > 0u; st = st >> 1u) {{
if (lid < st) {{ red[lid] = red[lid] + red[lid + st]; }}
workgroupBarrier();
}}
if (lid == 0u) {{ sqc[cc] = red[0]; }}
workgroupBarrier();
}}
for (var cc = 0u; cc < KC; cc = cc + 1u) {{
red[lid] = ag[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();
let gsum = red[lid] + red[lid + 1u];
workgroupBarrier();
red[lid] = au[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 usum = red[lid] + red[lid + 1u];
let inv = 1.0 / sqrt(sqc[cc] / f32(n) + epsm.x);
let gate = inv * gsum;
let upv = inv * usum;
var act: f32;
if (epsm.y != 0.0) {{
let g3 = gate * gate * gate;
let targ = clamp(0.7978845608028654 * (gate + 0.044715 * g3), -20.0, 20.0);
act = 0.5 * gate * (1.0 + tanh(targ));
}} else {{
act = gate / (1.0 + exp(-gate)); // SiLU
}}
y[(col0 + cc) * m + row] = act * upv;
}}
workgroupBarrier();
}}
}}
"#
)
}
pub(crate) fn conv_dot_k_src(eps: f32) -> String {
format!(
r#"{Q4_FN}
@group(0) @binding(0) var<storage, read> s_in: array<f16>;
@group(0) @binding(1) var<storage, read> q_in: array<u32>;
@group(0) @binding(2) var<storage, read> x: array<vec4<f32>>; // cur_k [k, h]
@group(0) @binding(3) var<storage, read> wn: array<vec4<f32>>;
@group(0) @binding(4) var<storage, read_write> bcx: array<f32>; // [k, 3h] (B|C|x gates)
@group(0) @binding(5) var<uniform> dims: vec4<u32>; // (hidden, conv_l, _, _)
const EPS: f32 = {eps};
const WG: u32 = 32u;
var<workgroup> psq: array<f32, 32>;
var<workgroup> pb: array<f32, 32>;
var<workgroup> pc: array<f32, 32>;
var<workgroup> px: array<f32, 32>;
@compute @workgroup_size(32)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid3: vec3<u32>) {{
let dim_out = wid.x; let lid = lid3.x; let col = wid.y;
let h = dims.x;
if (dim_out >= h) {{ return; }}
let nblk = h / 32u;
let xoff = col * (h / 4u);
let rb = dim_out; let rc = h + dim_out; let rx = 2u * h + dim_out;
var sq = 0.0; var ab = 0.0; var ac = 0.0; var ax = 0.0;
for (var b = lid; b < nblk; b = b + WG) {{
let db = f32(s_in[rb * nblk + b]); let dc = f32(s_in[rc * nblk + b]); let dx = f32(s_in[rx * nblk + b]);
let qbb = (rb * nblk + b) * 4u; let qbc = (rc * nblk + b) * 4u; let qbx = (rx * nblk + b) * 4u;
let xb = b * 8u;
for (var wi = 0u; wi < 4u; wi = wi + 1u) {{
let nlo = x[xoff + xb + wi] * wn[xb + wi];
let nhi = x[xoff + xb + 4u + wi] * wn[xb + 4u + wi];
sq = sq + dot(x[xoff + xb + wi], x[xoff + xb + wi]) + dot(x[xoff + xb + 4u + wi], x[xoff + xb + 4u + wi]);
ab = ab + db * (dot(q4_lo(q_in[qbb + wi]), nlo) + dot(q4_hi(q_in[qbb + wi]), nhi));
ac = ac + dc * (dot(q4_lo(q_in[qbc + wi]), nlo) + dot(q4_hi(q_in[qbc + wi]), nhi));
ax = ax + dx * (dot(q4_lo(q_in[qbx + wi]), nlo) + dot(q4_hi(q_in[qbx + wi]), nhi));
}}
}}
psq[lid] = sq; pb[lid] = ab; pc[lid] = ac; px[lid] = ax;
workgroupBarrier();
var stride = WG / 2u;
loop {{
if (lid < stride) {{ psq[lid] = psq[lid] + psq[lid + stride]; pb[lid] = pb[lid + stride] + pb[lid]; pc[lid] = pc[lid + stride] + pc[lid]; px[lid] = px[lid + stride] + px[lid]; }}
workgroupBarrier();
if (stride == 1u) {{ break; }}
stride = stride / 2u;
}}
if (lid == 0u) {{
let inv = 1.0 / sqrt(psq[0] / f32(h) + EPS);
let base = col * 3u * h;
bcx[base + dim_out] = inv * pb[0];
bcx[base + h + dim_out] = inv * pc[0];
bcx[base + 2u * h + dim_out] = inv * px[0];
}}
}}
"#
)
}
pub(crate) const CONV_MIX_K: &str = r#"
@group(0) @binding(0) var<storage, read> bcx: array<f32>; // [k, 3h]
@group(0) @binding(1) var<storage, read> cw: array<f32>; // [h, conv_l]
@group(0) @binding(2) var<storage, read_write> state: array<f32>; // [slots, h, conv_l]
@group(0) @binding(3) var<storage, read_write> y: array<f32>; // [k, h]
@group(0) @binding(4) var<storage, read_write> snap: array<f32>; // [k, h, conv_l]
@group(0) @binding(5) var<storage, read> cmeta: array<u32>; // [steps, stride, 2]: (pos, ring slot)
@group(0) @binding(6) var<storage, read> cnt: array<u32>; // chained-step counter
@group(0) @binding(7) var<uniform> dims: vec4<u32>; // (h, conv_l, k, step_stride)
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
let h = dims.x; let l = dims.y; let k = dims.z;
let c = gid.x;
if (c >= h) { return; }
for (var p = 0u; p < k; p = p + 1u) {
let sbase = cmeta[(cnt[0] * dims.w + p) * 2u + 1u] * h * l + c * l;
let base = p * 3u * h;
let bx = bcx[base + c] * bcx[base + 2u * h + c];
var conv_out = bx * cw[c * l + l - 1u];
for (var tap = 0u; tap + 1u < l; tap = tap + 1u) {
let nxt = state[sbase + tap + 1u];
state[sbase + tap] = nxt;
conv_out = conv_out + nxt * cw[c * l + tap];
}
state[sbase + l - 1u] = bx;
y[p * h + c] = bcx[base + h + c] * conv_out;
for (var tap = 0u; tap < l; tap = tap + 1u) { snap[(p * h + c) * l + tap] = state[sbase + tap]; }
}
}
"#;
const CONV_FUSED_F16: &str = r#"enable f16;
@group(0) @binding(0) var<storage, read> w: array<vec4<f16>>; // in_proj [3*hidden, hidden]
@group(0) @binding(1) var<storage, read> x: array<vec4<f32>>; // cur [hidden]
@group(0) @binding(2) var<storage, read> wn: array<vec4<f32>>; // operator_norm [hidden]
@group(0) @binding(3) var<storage, read> cw: array<f32>; // conv_w [hidden, conv_l]
@group(0) @binding(4) var<storage, read_write> state: array<f32>; // conv_state [hidden, conv_l]
@group(0) @binding(5) var<storage, read_write> y: array<f32>; // conv_y [hidden]
@group(0) @binding(6) var<uniform> dims: vec4<u32>; // (hidden, conv_l, _, _)
@group(0) @binding(7) var<uniform> epsm: vec4<f32>;
const WG: u32 = 32u;
var<workgroup> psq: array<f32, 32>;
var<workgroup> p0: array<f32, 32>;
var<workgroup> p1: array<f32, 32>;
var<workgroup> p2: array<f32, 32>;
@compute @workgroup_size(32)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid3: vec3<u32>) {
let dim_out = wid.x; let lid = lid3.x;
let h = dims.x; let l = dims.y;
if (dim_out >= h) { return; }
let n4 = h / 4u;
let base_b = dim_out * n4;
let base_c = (h + dim_out) * n4;
let base_x = (2u * h + dim_out) * n4;
var sq = 0.0; var ab = 0.0; var ac = 0.0; var ax = 0.0;
for (var j = lid; j < n4; j = j + WG) {
let xv = x[j];
let xn = vec4<f16>(xv * wn[j]);
sq = sq + dot(xv, xv);
ab = ab + f32(dot(w[base_b + j], xn));
ac = ac + f32(dot(w[base_c + j], xn));
ax = ax + f32(dot(w[base_x + j], xn));
}
psq[lid] = sq; p0[lid] = ab; p1[lid] = ac; p2[lid] = ax;
workgroupBarrier();
var stride = WG / 2u;
loop {
if (lid < stride) {
psq[lid] = psq[lid] + psq[lid + stride];
p0[lid] = p0[lid] + p0[lid + stride];
p1[lid] = p1[lid] + p1[lid + stride];
p2[lid] = p2[lid] + p2[lid + stride];
}
workgroupBarrier();
if (stride == 1u) { break; }
stride = stride / 2u;
}
if (lid == 0u) {
let inv = 1.0 / sqrt(psq[0] / f32(h) + epsm.x);
let b_gate = inv * p0[0];
let c_gate = inv * p1[0];
let x_val = inv * p2[0];
let bx = b_gate * x_val;
let sbase = dim_out * l;
var conv_out = bx * cw[sbase + l - 1u];
for (var tap = 0u; tap + 1u < l; tap = tap + 1u) {
let nxt = state[sbase + tap + 1u];
state[sbase + tap] = nxt;
conv_out = conv_out + nxt * cw[sbase + tap];
}
state[sbase + l - 1u] = bx;
y[dim_out] = c_gate * conv_out;
}
}
"#;
pub struct ConvFusedF16 {
pipeline: wgpu::ComputePipeline,
}
impl ConvFusedF16 {
pub fn new(ctx: &GpuCtx) -> Self {
let m = ctx
.device
.shader_module_tuned(wgpu::ShaderModuleDescriptor {
label: Some("conv_fused_f16"),
source: wgpu::ShaderSource::Wgsl(CONV_FUSED_F16.into()),
});
Self {
pipeline: ctx
.device
.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some("conv_fused_f16"),
layout: None,
module: &m,
entry_point: Some("main"),
compilation_options: wgpu::PipelineCompilationOptions::default(),
cache: None,
}),
}
}
pub fn pipeline(&self) -> &wgpu::ComputePipeline {
&self.pipeline
}
#[allow(clippy::too_many_arguments)]
pub fn make(
&self,
ctx: &GpuCtx,
in_proj: &crate::weights::F16,
x: &wgpu::Buffer,
wn: &wgpu::Buffer,
cw: &wgpu::Buffer,
state: &wgpu::Buffer,
y: &wgpu::Buffer,
hidden: u32,
conv_l: u32,
eps: f32,
) -> (wgpu::BindGroup, u32, u32) {
let dims = ctx
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("dims"),
contents: bytemuck::cast_slice(&[hidden, conv_l, 0u32, 0u32]),
usage: wgpu::BufferUsages::UNIFORM,
});
let meta = ctx
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("eps"),
contents: bytemuck::cast_slice(&[eps, 0.0, 0.0, 0.0]),
usage: wgpu::BufferUsages::UNIFORM,
});
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: in_proj.buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: wn.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: cw.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: state.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 5,
resource: y.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 6,
resource: dims.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 7,
resource: meta.as_entire_binding(),
},
],
});
(bg, hidden, 1)
}
}
fn q4_lmhead_tiled_src() -> String {
format!(
r#"{Q4_FN}
@group(0) @binding(0) var<storage, read> scales: array<f16>;
@group(0) @binding(1) var<storage, read> quants: array<u32>;
@group(0) @binding(2) var<storage, read> x: array<vec4<f32>>; // [ncols, N/4]
@group(0) @binding(3) var<storage, read_write> y: array<f32>; // [ncols, M]
@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, ncols, gy_rows)
const WG: u32 = 32u;
const NR: u32 = 4u; // rows per workgroup
const KC: u32 = 4u; // columns per workgroup — one weight stream feeds up to 4 batch columns
var<workgroup> partial: array<f32, 512>; // WG*NR*KC
@compute @workgroup_size(32)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid3: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {{
let lid = lid3.x;
// Row tiles ride (wid.x, wid.y % gy_rows); column tiles ride wid.y / gy_rows — flattening
// both into gy keeps the dispatch 2-D. Per (row, col) lane the FP sequence is EXACTLY the
// single-column kernel's (s over wi, then acc += d·s, same reduction tree), so any column of
// a batched dispatch is bitwise-equal to a solo dispatch of that column.
let gy_rows = dims.w;
let dim0 = (wid.x + (wid.y % gy_rows) * nwg.x) * NR;
let col0 = (wid.y / gy_rows) * KC;
let ncols = dims.z;
let nblk = dims.y / 32u;
let xstride = dims.y / 4u;
var acc: array<f32, 16>; // [NR][KC]
for (var i = 0u; i < 16u; i = i + 1u) {{ acc[i] = 0.0; }}
for (var b = lid; b < nblk; b = b + WG) {{
let xb = b * 8u;
for (var n = 0u; n < NR; n = n + 1u) {{
let row = dim0 + n;
let d = f32(scales[row * nblk + b]);
let qb = (row * nblk + b) * 4u;
var s: array<f32, 4>;
for (var c = 0u; c < KC; c = c + 1u) {{ s[c] = 0.0; }}
for (var wi = 0u; wi < 4u; wi = wi + 1u) {{
let lo = q4_lo(quants[qb + wi]);
let hi = q4_hi(quants[qb + wi]);
for (var c = 0u; c < KC; c = c + 1u) {{
let xo = (col0 + c) * xstride + xb;
s[c] = s[c] + dot(lo, x[xo + wi]) + dot(hi, x[xo + 4u + wi]);
}}
}}
for (var c = 0u; c < KC; c = c + 1u) {{
acc[n * KC + c] = acc[n * KC + c] + d * s[c];
}}
}}
}}
for (var i = 0u; i < 16u; i = i + 1u) {{ partial[lid * 16u + i] = acc[i]; }}
workgroupBarrier();
var stride = WG / 2u;
loop {{
if (lid < stride) {{
for (var i = 0u; i < 16u; i = i + 1u) {{
partial[lid * 16u + i] = partial[lid * 16u + i] + partial[(lid + stride) * 16u + i];
}}
}}
workgroupBarrier();
if (stride == 1u) {{ break; }}
stride = stride / 2u;
}}
if (lid < 16u) {{
let n = lid / KC;
let c = lid % KC;
let row = dim0 + n;
let col = col0 + c;
if (row < dims.x && col < ncols) {{ y[col * dims.x + row] = partial[lid]; }}
}}
}}
"#
)
}
fn q4_lmhead_tiled_kc16_src() -> String {
let src = q4_lmhead_tiled_src();
for frag in [
"const KC: u32 = 4u; // columns per workgroup — one weight stream feeds up to 4 batch columns",
"const NR: u32 = 4u; // rows per workgroup",
"var s: array<f32, 4>;",
] {
assert!(src.contains(frag), "tiled head source drifted: {frag}");
}
src.replace(
"const KC: u32 = 4u; // columns per workgroup — one weight stream feeds up to 4 batch columns",
"const KC: u32 = 16u; // columns per workgroup — one weight stream feeds up to 16 batch columns",
)
.replace(
"const NR: u32 = 4u; // rows per workgroup",
"const NR: u32 = 1u; // rows per workgroup (NR·KC stays 16 — see the derivation comment)",
)
.replace("var s: array<f32, 4>;", "var s: array<f32, 16>;")
}
pub fn test_head_kc4() -> String {
q4_lmhead_tiled_src()
}
pub fn test_head_kc16() -> String {
q4_lmhead_tiled_kc16_src()
}
#[allow(dead_code)]
fn q4_lmhead_sg_src() -> String {
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 fma(vec4<f32>(unpack4xU8(word & 0xF0F0F0F0u)), vec4<f32>(0.0625), vec4<f32>(-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>>; // [ncols, N/4]
@group(0) @binding(3) var<storage, read_write> y: array<f32>; // [ncols, M]
@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, ncols, gy_rows)
const KC: u32 = 4u;
@compute @workgroup_size(256)
fn main(
@builtin(workgroup_id) wid: vec3<u32>,
@builtin(local_invocation_id) lid3: vec3<u32>,
@builtin(num_workgroups) nwg: vec3<u32>,
@builtin(subgroup_invocation_id) sid: u32,
@builtin(subgroup_size) ssz: u32,
) {
let nsg = 256u / ssz;
let sg = lid3.x / ssz;
let gy_rows = dims.w;
let m = dims.x; let n = dims.y; let ncols = dims.z;
let row_raw = (wid.x + (wid.y % gy_rows) * nwg.x) * nsg + sg;
let row = min(row_raw, m - 1u);
let col0 = (wid.y / gy_rows) * KC;
let nblk = n / 32u;
let xstride = n / 4u;
var acc = vec4<f32>(0.0);
for (var b = sid; b < nblk; b = b + ssz) {
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 = b * 8u;
for (var cc = 0u; cc < KC; cc = cc + 1u) {
if (col0 + cc < ncols) {
let base = (col0 + cc) * xstride + xb;
var sv = dot(l0, x[base]) + dot(h0, x[base + 4u]);
sv = sv + dot(l1, x[base + 1u]) + dot(h1, x[base + 5u]);
sv = sv + dot(l2, x[base + 2u]) + dot(h2, x[base + 6u]);
sv = sv + dot(l3, x[base + 3u]) + dot(h3, x[base + 7u]);
acc[cc] = acc[cc] + d * sv;
}
}
}
let tot = subgroupAdd(acc);
if (sid == 0u && row_raw < m) {
for (var cc = 0u; cc < KC; cc = cc + 1u) {
if (col0 + cc < ncols) { y[(col0 + cc) * m + row] = tot[cc]; }
}
}
}
"#
.to_string()
}
#[allow(dead_code)]
pub(crate) fn mlp_gate_q4_k_sg16_src() -> String {
let src = mlp_gate_q4_k_sg_src();
for frag in [
"const KC: u32 = 4u;",
"var ag = vec4<f32>(0.0);",
"var au = vec4<f32>(0.0);",
"var sq = vec4<f32>(0.0);",
"let tg = subgroupAdd(ag);",
"let tu = subgroupAdd(au);",
"let ts = subgroupAdd(sq);",
] {
assert!(src.contains(frag), "mlp sg source drifted: {frag}");
}
src.replace("const KC: u32 = 4u;", "const KC: u32 = 16u;")
.replace("var ag = vec4<f32>(0.0);", "var ag = array<f32, 16>();")
.replace("var au = vec4<f32>(0.0);", "var au = array<f32, 16>();")
.replace("var sq = vec4<f32>(0.0);", "var sq = array<f32, 16>();")
.replace(
"let tg = subgroupAdd(ag);",
"var tg = array<f32, 16>();\n for (var cc = 0u; cc < KC; cc = cc + 1u) { tg[cc] = subgroupAdd(ag[cc]); }",
)
.replace(
"let tu = subgroupAdd(au);",
"var tu = array<f32, 16>();\n for (var cc = 0u; cc < KC; cc = cc + 1u) { tu[cc] = subgroupAdd(au[cc]); }",
)
.replace(
"let ts = subgroupAdd(sq);",
"var ts = array<f32, 16>();\n for (var cc = 0u; cc < KC; cc = cc + 1u) { ts[cc] = subgroupAdd(sq[cc]); }",
)
}
fn q4_lmhead_nbar_src() -> String {
let src = gemv_q4_k_nbar_src();
for frag in [
"@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)",
"let m = dims.x; let n = dims.y; let ncols = dims.w;",
"if (dims.z == 1u) {{ y[yo] = y[yo] + v; }} else {{ y[yo] = v; }}",
] {
let frag = frag.replace("{{", "{").replace("}}", "}");
assert!(src.contains(&frag), "nbar source drifted: {frag}");
}
src.replace(
"let m = dims.x; let n = dims.y; let ncols = dims.w;",
"let m = dims.x; let n = dims.y; let ncols = dims.z;",
)
.replace(
"if (dims.z == 1u) { y[yo] = y[yo] + v; } else { y[yo] = v; }",
"y[yo] = v;",
)
}
pub fn q1_lmhead_lcpp_src() -> String {
let src = gemv_q1_k_lcpp_src();
for frag in [
"@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)",
" let m = dims.x; let n = dims.y;",
" let row0 = wid.x * 4u;\n let col = wid.y;",
" if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }",
] {
assert!(src.contains(frag), "q1 lcpp source drifted: {frag}");
}
src.replace(
"@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)",
"@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, ncols, gy_rows)",
)
.replace(
" let row0 = wid.x * 4u;\n let col = wid.y;",
" let col = wid.y / dims.w;\n let row0 = ((wid.y % dims.w) * 32768u + wid.x) * 4u;",
)
.replace(
" if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }",
" y[yo] = tot[r];",
)
}
fn q1_lmhead_lcpp_nr_src(nr: u32) -> String {
assert!(
nr >= 4 && nr.is_multiple_of(4),
"nr must be a positive multiple of 4"
);
let g = nr / 4;
let mut accs = String::new();
for i in 0..g {
accs.push_str(&format!(" var acc{i} = vec4<f32>(0.0);\n"));
}
let mut body = String::new();
for i in 0..g {
body.push_str(&format!(
r#" for (var r = 0u; r < 4u; r = r + 1u) {{
let row = min(row0 + {off}u + r, m - 1u);
let d = f32(scales[row * nsc + (b >> 2u)]);
let w = bits[row * nblk + b];
var s = dot(q1s(w, 0u), v0) + dot(q1s(w, 4u), v1);
s = s + dot(q1s(w, 8u), v2) + dot(q1s(w, 12u), v3);
s = s + dot(q1s(w, 16u), v4) + dot(q1s(w, 20u), v5);
s = s + dot(q1s(w, 24u), v6) + dot(q1s(w, 28u), v7);
acc{i}[r] = acc{i}[r] + d * s;
}}
"#,
off = i * 4
));
}
let mut tail = String::new();
for i in 0..g {
tail.push_str(&format!(" let tot{i} = subgroupAdd(acc{i});\n"));
}
let mut writes = String::new();
for i in 0..g {
writes.push_str(&format!(
r#" for (var r = 0u; r < 4u; r = r + 1u) {{
if (row0 + {off}u + r < m) {{
y[col * m + row0 + {off}u + r] = tot{i}[r];
}}
}}
"#,
off = i * 4
));
}
format!(
r#"enable f16;
fn q1s(word: u32, sh: u32) -> vec4<f32> {{
let bits = (vec4<u32>(word) >> vec4<u32>(sh, sh + 1u, sh + 2u, sh + 3u)) & vec4<u32>(1u);
return select(vec4<f32>(-1.0), vec4<f32>(1.0), bits == vec4<u32>(1u));
}}
@group(0) @binding(0) var<storage, read> scales: array<f16>; // [m, n/128]
@group(0) @binding(1) var<storage, read> bits: array<u32>; // [m, n/32]
@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>; // (M, N, ncols, gy_rows)
@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 nsc = n / 128u;
let xstride = n / 4u;
let col = wid.y / dims.w;
let row0 = ((wid.y % dims.w) * 32768u + wid.x) * {nr}u;
let xoff = col * xstride;
{accs} 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];
{body} }}
{tail} if (sid == 0u) {{
{writes} }}
}}
"#
)
}
fn q4_lmhead_lcpp_src() -> String {
let src = gemv_q4_k_lcpp_src();
for frag in [
"@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)",
" let m = dims.x; let n = dims.y;",
" let row0 = wid.x * 4u;\n let col = wid.y;",
" if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }",
] {
assert!(src.contains(frag), "lcpp source drifted: {frag}");
}
src.replace(
"@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, acc, ncols)",
"@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, ncols, gy_rows)",
)
.replace(
" let row0 = wid.x * 4u;\n let col = wid.y;",
" let col = wid.y / dims.w;\n let row0 = ((wid.y % dims.w) * 32768u + wid.x) * 4u;",
)
.replace(
" if (dims.z == 1u) { y[yo] = y[yo] + tot[r]; } else { y[yo] = tot[r]; }",
" y[yo] = tot[r];",
)
}
fn q4_lmhead_lcpp_nr_src(nr: u32) -> String {
assert!(
nr >= 4 && nr.is_multiple_of(4),
"nr must be a positive multiple of 4"
);
let g = nr / 4;
let mut accs = String::new();
for i in 0..g {
accs.push_str(&format!(" var acc{i} = vec4<f32>(0.0);\n"));
}
let mut body = String::new();
for i in 0..g {
body.push_str(&format!(
r#" for (var r = 0u; r < 4u; r = r + 1u) {{
let row = min(row0 + {off}u + r, m - 1u);
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{i}[r] = acc{i}[r] + d * s;
}}
"#,
off = i * 4
));
}
let mut tail = String::new();
for i in 0..g {
tail.push_str(&format!(" let tot{i} = subgroupAdd(acc{i});\n"));
}
let mut writes = String::new();
for i in 0..g {
writes.push_str(&format!(
r#" for (var r = 0u; r < 4u; r = r + 1u) {{
if (row0 + {off}u + r < m) {{
y[col * m + row0 + {off}u + r] = tot{i}[r];
}}
}}
"#,
off = i * 4
));
}
format!(
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 fma(vec4<f32>(unpack4xU8(word & 0xF0F0F0F0u)), vec4<f32>(0.0625), vec4<f32>(-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>; // (M, N, ncols, gy_rows)
@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 col = wid.y / dims.w;
let row0 = ((wid.y % dims.w) * 32768u + wid.x) * {nr}u;
let xoff = col * xstride;
{accs} 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];
{body} }}
{tail} if (sid == 0u) {{
{writes} }}
}}
"#
)
}
#[allow(dead_code)]
fn q4_lmhead_lcpp_banded_src() -> String {
format!(
r#"{Q4_FN}
@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>>; // [ncols, N/4]
@group(0) @binding(3) var<storage, read_write> y: array<f32>; // [ncols, M]
@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, ncols, gy_rows)
@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 nc = min(dims.z, 8u);
let nblk = n / 32u;
let xstride = n / 4u;
let row0 = (wid.y * 32768u + wid.x) * 4u;
var acc: array<array<f32, 4>, 8>;
for (var c = 0u; c < 8u; c = c + 1u) {{
acc[c][0] = 0.0; acc[c][1] = 0.0; acc[c][2] = 0.0; acc[c][3] = 0.0;
}}
for (var b = sid; b < nblk; b = b + 32u) {{
// Weights for the 4 rows of this block — read ONCE, dotted with every column.
for (var r = 0u; r < 4u; r = r + 1u) {{
let row = min(row0 + r, m - 1u);
let d = f32(scales[row * nblk + b]);
let q = quants[row * nblk + b];
let ql0 = q4_lo(q.x); let qh0 = q4_hi(q.x);
let ql1 = q4_lo(q.y); let qh1 = q4_hi(q.y);
let ql2 = q4_lo(q.z); let qh2 = q4_hi(q.z);
let ql3 = q4_lo(q.w); let qh3 = q4_hi(q.w);
for (var c = 0u; c < nc; c = c + 1u) {{
let xb = c * xstride + b * 8u;
var s = dot(ql0, x[xb]) + dot(qh0, x[xb + 4u]);
s = s + dot(ql1, x[xb + 1u]) + dot(qh1, x[xb + 5u]);
s = s + dot(ql2, x[xb + 2u]) + dot(qh2, x[xb + 6u]);
s = s + dot(ql3, x[xb + 3u]) + dot(qh3, x[xb + 7u]);
acc[c][r] = acc[c][r] + d * s;
}}
}}
}}
for (var c = 0u; c < nc; c = c + 1u) {{
let t0 = subgroupAdd(acc[c][0]);
let t1 = subgroupAdd(acc[c][1]);
let t2 = subgroupAdd(acc[c][2]);
let t3 = subgroupAdd(acc[c][3]);
if (sid == 0u) {{
if (row0 < m) {{ y[c * m + row0] = t0; }}
if (row0 + 1u < m) {{ y[c * m + row0 + 1u] = t1; }}
if (row0 + 2u < m) {{ y[c * m + row0 + 2u] = t2; }}
if (row0 + 3u < m) {{ y[c * m + row0 + 3u] = t3; }}
}}
}}
}}
"#
)
}
pub struct Q4LmHead {
pipeline: wgpu::ComputePipeline,
pipeline_nbar: wgpu::ComputePipeline,
nbar: bool,
pipeline_lcpp: Option<wgpu::ComputePipeline>,
pipeline_kc16: wgpu::ComputePipeline,
force_kc: Option<u32>,
band: u32,
rows_per_wg: u32,
head_nr: u32,
kind: crate::weights::WKind,
}
impl Q4LmHead {
pub fn new(ctx: &GpuCtx, kind: crate::weights::WKind) -> Self {
use crate::weights::WKind;
assert!(
matches!(kind, WKind::Q4_0 | WKind::F16 | WKind::Q8_0N),
"no lmhead kernel for {kind:?} yet (Q4_0 / f16 / native Q8_0N only)"
);
let wf16 = kind != WKind::Q4_0; let (label, src, rows_per_wg) = ("q4_lmhead_tiled", q4_lmhead_tiled_src(), 4);
let has_lcpp = ctx.subgroups32_effective();
let head_nr = std::env::var("OSFKB_HEAD_NR")
.ok()
.and_then(|v| v.parse().ok())
.filter(|nr: &u32| *nr >= 4 && (*nr).is_multiple_of(4) && *nr <= 32)
.unwrap_or(8);
let head_nr = if wf16 { 4 } else { head_nr };
let mk = |label: &str, src: String| {
let m = ctx
.device
.shader_module_tuned(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: wgpu::PipelineCompilationOptions::default(),
cache: None,
})
};
Self {
pipeline: mk(label, src),
pipeline_nbar: mk("q4_lmhead_nbar", q4_lmhead_nbar_src()),
nbar: has_lcpp && std::env::var("OSFKB_HEAD_NBAR").ok().as_deref() != Some("0"),
pipeline_lcpp: has_lcpp.then(|| {
if kind == WKind::Q8_0N {
mk("q8_0n_lmhead_lcpp", q8_0n_lmhead_lcpp_src())
} else if wf16 {
mk("f16_lmhead_lcpp", f16_lmhead_lcpp_src())
} else if std::env::var("OSFKB_Q1_WEIGHTS").ok().as_deref() == Some("1") {
if head_nr > 4 {
mk("q1_lmhead_lcpp_nr", q1_lmhead_lcpp_nr_src(head_nr))
} else {
mk("q1_lmhead_lcpp", q1_lmhead_lcpp_src())
}
} else if head_nr != 4 {
mk("q4_lmhead_lcpp_nr", q4_lmhead_lcpp_nr_src(head_nr))
} else {
mk("q4_lmhead_lcpp", q4_lmhead_lcpp_src())
}
}),
pipeline_kc16: mk("q4_lmhead_tiled_kc16", q4_lmhead_tiled_kc16_src()),
band: std::env::var("OSFKB_HEAD_BAND")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(u32::MAX),
force_kc: std::env::var("OSFKB_HEAD_KC")
.ok()
.and_then(|v| v.parse().ok())
.filter(|kc| *kc == 4 || *kc == 16),
rows_per_wg,
head_nr,
kind,
}
}
pub fn pipeline(&self) -> &wgpu::ComputePipeline {
&self.pipeline
}
fn kc_for(&self, ncols: u32) -> u32 {
match self.force_kc {
Some(16) if ncols > 1 => 16,
_ => 4,
}
}
pub fn set_single_stream(&mut self, on: bool) {
self.nbar = on && self.pipeline_lcpp.is_some();
}
pub fn set_band(&mut self, band: u32) {
self.band = band;
}
pub fn active_variant(&self, ncols: u32) -> &'static str {
if self.nbar && ncols <= self.band {
if self.pipeline_lcpp.is_some() {
return "lcpp";
}
if ncols <= 2 {
return "nbar";
}
}
if self.kc_for(ncols) == 16 {
"tiled_kc16"
} else {
"tiled"
}
}
pub fn pipeline_for(&self, ncols: u32) -> &wgpu::ComputePipeline {
if self.kind != crate::weights::WKind::Q4_0 {
return self.pipeline_lcpp.as_ref().expect(
"no-scale-table heads require the lcpp family (validated 32-wide subgroups)",
);
}
if self.nbar && ncols <= self.band {
if let Some(pl) = &self.pipeline_lcpp {
return pl;
}
if ncols <= 2 {
return &self.pipeline_nbar;
}
}
if self.kc_for(ncols) == 16 {
&self.pipeline_kc16
} else {
&self.pipeline
}
}
#[allow(clippy::too_many_arguments)]
pub fn make(
&self,
ctx: &GpuCtx,
w: &crate::weights::Q4,
x: &wgpu::Buffer,
y: &wgpu::Buffer,
m: u32,
n: u32,
ncols: u32,
) -> (wgpu::BindGroup, u32, u32) {
let kc = self.kc_for(ncols);
let lcpp = self.nbar && ncols <= self.band && self.pipeline_lcpp.is_some();
let rows = if self.nbar && ncols <= self.band {
if self.pipeline_lcpp.is_some() {
self.head_nr
} else if ncols <= 2 {
16
} else {
self.rows_per_wg
}
} else if kc == 16 {
1
} else {
self.rows_per_wg
};
let nwg_rows = m.div_ceil(rows);
let gx = nwg_rows.min(32768);
let gy_rows = nwg_rows.div_ceil(gx);
let gy = if lcpp {
gy_rows * ncols
} else {
gy_rows * ncols.div_ceil(kc)
};
let dims = ctx
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("dims"),
contents: bytemuck::cast_slice(&[m, n, ncols, gy_rows]),
usage: wgpu::BufferUsages::UNIFORM,
});
let bufs: [&wgpu::Buffer; 5] = if self.kind != crate::weights::WKind::Q4_0 {
[&w.quants, x, y, &dims, &dims]
} else {
[&w.scales, &w.quants, x, y, &dims]
};
let used = if self.kind != crate::weights::WKind::Q4_0 {
4
} else {
5
};
let entries: Vec<wgpu::BindGroupEntry> = bufs[..used]
.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: &self.pipeline_for(ncols).get_bind_group_layout(0),
entries: &entries,
});
(bg, gx, gy)
}
}
fn q4_lmhead_sparse_src() -> String {
format!(
r#"{Q4_FN}
@group(0) @binding(0) var<storage, read> scales: array<f16>;
@group(0) @binding(1) var<storage, read> quants: array<u32>;
@group(0) @binding(2) var<storage, read> x: array<vec4<f32>>;
@group(0) @binding(3) var<storage, read_write> svals: array<f32>; // svals[slot] = logit of ids[slot]
@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M=vocab, N=hidden, _, _)
@group(0) @binding(5) var<storage, read> smeta: array<u32>; // [4] = frozen allowed-count
@group(0) @binding(6) var<storage, read> ids: array<u32>; // compacted allowed token ids
const WG: u32 = 32u;
const NR: u32 = 8u;
var<workgroup> partial: array<f32, 256>; // WG*NR
@compute @workgroup_size(32)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid3: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {{
let lid = lid3.x;
let slot0 = (wid.x + wid.y * nwg.x) * NR;
let cnt = smeta[4];
let nblk = dims.y / 32u;
var acc: array<f32, 8>;
for (var n = 0u; n < NR; n = n + 1u) {{ acc[n] = 0.0; }}
for (var b = lid; b < nblk; b = b + WG) {{
let xb = b * 8u;
for (var n = 0u; n < NR; n = n + 1u) {{
let slot = slot0 + n;
var row = 0u;
if (slot < cnt) {{ row = ids[slot]; }}
let d = f32(scales[row * nblk + b]);
let qb = (row * nblk + b) * 4u;
var s = 0.0;
for (var wi = 0u; wi < 4u; wi = wi + 1u) {{
s = s + dot(q4_lo(quants[qb + wi]), x[xb + wi]) + dot(q4_hi(quants[qb + wi]), x[xb + 4u + wi]);
}}
acc[n] = acc[n] + d * s;
}}
}}
for (var n = 0u; n < NR; n = n + 1u) {{ partial[lid * NR + n] = acc[n]; }}
workgroupBarrier();
var stride = WG / 2u;
loop {{
if (lid < stride) {{ for (var n = 0u; n < NR; n = n + 1u) {{ partial[lid * NR + n] = partial[lid * NR + n] + partial[(lid + stride) * NR + n]; }} }}
workgroupBarrier();
if (stride == 1u) {{ break; }}
stride = stride / 2u;
}}
if (lid < NR) {{ let slot = slot0 + lid; if (slot < cnt) {{ svals[slot] = partial[lid]; }} }}
}}
"#
)
}
#[allow(dead_code)]
fn q4_lmhead_sparse_sg_src() -> String {
let src = q4_lmhead_sg_src();
let frags = [
"@group(0) @binding(4) var<uniform> dims: vec4<u32>; // (M, N, ncols, gy_rows)",
" let row_raw = (wid.x + (wid.y % gy_rows) * nwg.x) * nsg + sg;",
" let row = min(row_raw, m - 1u);",
" if (sid == 0u && row_raw < m) {",
" for (var cc = 0u; cc < KC; cc = cc + 1u) {",
" if (col0 + cc < ncols) { y[(col0 + cc) * m + row] = tot[cc]; }",
];
for f in frags {
assert!(src.contains(f), "full sg head source drifted: {f}");
}
src.replace(
"@group(0) @binding(3) var<storage, read_write> y: array<f32>;",
"@group(0) @binding(3) var<storage, read_write> y: array<f32>;\n@group(0) @binding(5) var<storage, read> smeta: array<u32>;\n@group(0) @binding(6) var<storage, read> ids: array<u32>;",
)
.replace(
" let gy_rows = dims.w;",
" let gy_rows = max(dims.w, 1u);",
)
.replace(
" let ncols = dims.z;",
" let ncols = max(dims.z, 1u);",
)
.replace(
" let row_raw = (wid.x + (wid.y % gy_rows) * nwg.x) * nsg + sg;\n let row = min(row_raw, m - 1u);",
" let slot = (wid.x + (wid.y % gy_rows) * nwg.x) * nsg + sg;\n let row_raw = slot;\n var row = 0u;\n if (slot < smeta[4]) { row = ids[slot]; }",
)
.replace(
" if (sid == 0u && row_raw < m) {\n for (var cc = 0u; cc < KC; cc = cc + 1u) {\n if (col0 + cc < ncols) { y[(col0 + cc) * m + row] = tot[cc]; }",
" if (sid == 0u && slot < smeta[4]) {\n for (var cc = 0u; cc < KC; cc = cc + 1u) {\n if (col0 + cc < ncols) { y[(col0 + cc) * m + slot] = tot[cc]; }",
)
}
pub struct Q4LmHeadSparse {
pipeline: wgpu::ComputePipeline,
}
impl Q4LmHeadSparse {
pub fn new(ctx: &GpuCtx) -> Self {
let (label, src) = ("q4_lmhead_sparse", q4_lmhead_sparse_src());
let m = ctx
.device
.shader_module_tuned(wgpu::ShaderModuleDescriptor {
label: Some(label),
source: wgpu::ShaderSource::Wgsl(src.into()),
});
Self {
pipeline: ctx
.device
.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some(label),
layout: None,
module: &m,
entry_point: Some("main"),
compilation_options: wgpu::PipelineCompilationOptions::default(),
cache: None,
}),
}
}
pub fn pipeline(&self) -> &wgpu::ComputePipeline {
&self.pipeline
}
#[allow(clippy::too_many_arguments)]
pub fn make(
&self,
ctx: &GpuCtx,
w: &crate::weights::Q4,
x: &wgpu::Buffer,
svals: &wgpu::Buffer,
smeta: &wgpu::Buffer,
ids: &wgpu::Buffer,
m: u32,
n: u32,
) -> wgpu::BindGroup {
self.make_at(ctx, w, x, 0, x.size(), svals, smeta, ids, m, n)
}
#[allow(clippy::too_many_arguments)]
pub fn make_at(
&self,
ctx: &GpuCtx,
w: &crate::weights::Q4,
x: &wgpu::Buffer,
x_off: u64,
x_size: u64,
svals: &wgpu::Buffer,
smeta: &wgpu::Buffer,
ids: &wgpu::Buffer,
m: u32,
n: u32,
) -> wgpu::BindGroup {
let dims = ctx
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("dims"),
contents: bytemuck::cast_slice(&[m, n, 0u32, 0u32]),
usage: wgpu::BufferUsages::UNIFORM,
});
ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("q4_lmhead_sparse"),
layout: &self.pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: w.scales.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: w.quants.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: wgpu::BindingResource::Buffer(wgpu::BufferBinding {
buffer: x,
offset: x_off,
size: std::num::NonZeroU64::new(x_size),
}),
},
wgpu::BindGroupEntry {
binding: 3,
resource: svals.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: dims.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 5,
resource: smeta.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 6,
resource: ids.as_entire_binding(),
},
],
})
}
}
#[cfg(test)]
mod dp4a_probe {
use crate::forward::ShaderModuleTuned as _;
#[test]
fn dp4a_supported() {
let Ok(ctx) = super::GpuCtx::new() else {
eprintln!("SKIP: no gpu");
return;
};
let src = r#"@group(0) @binding(0) var<storage, read_write> y: array<i32>;
@compute @workgroup_size(1)
fn main() {
let a: u32 = 0x01020304u;
let b: u32 = 0x05060708u;
y[0] = dot4I8Packed(a, b);
y[1] = i32(dot4U8Packed(a, b));
}"#;
let scope = ctx.device.push_error_scope(wgpu::ErrorFilter::Validation);
let m = ctx
.device
.shader_module_tuned(wgpu::ShaderModuleDescriptor {
label: Some("dp4a"),
source: wgpu::ShaderSource::Wgsl(src.into()),
});
let _p = ctx
.device
.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some("dp4a"),
layout: None,
module: &m,
entry_point: Some("main"),
compilation_options: Default::default(),
cache: None,
});
let err = pollster::block_on(scope.pop());
match err {
None => eprintln!("DP4A OK on {} (dot4I8Packed available)", ctx.backend),
Some(e) => eprintln!("DP4A UNAVAILABLE on {}: {e:?}", ctx.backend),
}
}
}
#[cfg(test)]
mod gemm_q1_dev {
use crate::forward::ShaderModuleTuned as _;
#[test]
fn rp4_shaders_validate() {
let Ok(ctx) = crate::GpuCtx::new() else {
eprintln!("SKIP");
return;
};
for (name, src) in [
("q1_rp4_repack", crate::q1_rp4_repack_src()),
("gemv_q1_rp4_f16x", crate::gemv_q1_rp4_f16x_src()),
("mlp_gate_q1_rp4_f16x", crate::mlp_gate_q1_rp4_f16x_src()),
("gemm_q1_xt_rp4", crate::gemm_q1_xt_rp4_src()),
] {
let md = ctx
.device
.shader_module_tuned(wgpu::ShaderModuleDescriptor {
label: Some(name),
source: wgpu::ShaderSource::Wgsl(src.as_str().into()),
});
let _pl = ctx
.device
.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some(name),
layout: None,
module: &md,
entry_point: Some("main"),
compilation_options: Default::default(),
cache: None,
});
eprintln!("{name} pipeline OK");
}
}
#[test]
fn gemm_q1_shader_validates() {
let Ok(ctx) = super::GpuCtx::new() else {
eprintln!("SKIP: no gpu");
return;
};
let src = super::gemm_q1_src();
let scope = ctx.device.push_error_scope(wgpu::ErrorFilter::Validation);
let m = ctx
.device
.shader_module_tuned(wgpu::ShaderModuleDescriptor {
label: Some("gemm_q1"),
source: wgpu::ShaderSource::Wgsl(src.into()),
});
let _p = ctx
.device
.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some("gemm_q1"),
layout: None,
module: &m,
entry_point: Some("main"),
compilation_options: Default::default(),
cache: None,
});
let err = pollster::block_on(scope.pop());
assert!(err.is_none(), "gemm_q1 failed validation: {err:?}");
eprintln!("gemm_q1 pipeline OK on {}", ctx.backend);
}
}