use anyhow::{Context, Result};
use wgpu::util::DeviceExt;
use crate::backend::wgsl_pp::Preprocessor;
use crate::tensor::DType;
use half::f16;
pub mod io_stats {
use std::sync::atomic::{AtomicU64, Ordering::Relaxed};
static SUBMITS: AtomicU64 = AtomicU64::new(0);
static READBACKS: AtomicU64 = AtomicU64::new(0);
static READBACK_BYTES: AtomicU64 = AtomicU64::new(0);
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct GpuIoStats {
pub submits: u64,
pub readbacks: u64,
pub readback_bytes: u64,
}
impl GpuIoStats {
pub fn per_token(&self, tokens: u64) -> (f64, f64, f64) {
if tokens == 0 {
return (0.0, 0.0, 0.0);
}
let t = tokens as f64;
(
self.submits as f64 / t,
self.readbacks as f64 / t,
self.readback_bytes as f64 / t,
)
}
}
pub(crate) fn record_submit() {
SUBMITS.fetch_add(1, Relaxed);
}
pub(crate) fn record_readback(bytes: u64) {
READBACKS.fetch_add(1, Relaxed);
READBACK_BYTES.fetch_add(bytes, Relaxed);
}
pub fn snapshot() -> GpuIoStats {
GpuIoStats {
submits: SUBMITS.load(Relaxed),
readbacks: READBACKS.load(Relaxed),
readback_bytes: READBACK_BYTES.load(Relaxed),
}
}
pub fn reset() {
SUBMITS.store(0, Relaxed);
READBACKS.store(0, Relaxed);
READBACK_BYTES.store(0, Relaxed);
}
}
pub struct GpuContext {
pub device: wgpu::Device,
pub queue: wgpu::Queue,
pub adapter_name: String,
pub backend: String,
pub max_storage_buffer_binding_size: u64,
pub max_buffer_size: u64,
pub min_storage_buffer_offset_alignment: u64,
pub preprocessor: Preprocessor,
pub profiler: Option<GpuProfiler>,
staging: std::sync::Mutex<Option<wgpu::Buffer>>,
staging_size: std::sync::atomic::AtomicU64,
}
pub struct GpuTensor {
pub buffer: wgpu::Buffer,
pub dtype: DType,
pub shape: Vec<usize>,
}
impl GpuTensor {
pub fn numel(&self) -> usize {
self.shape.iter().product()
}
pub fn size_bytes(&self) -> usize {
let block_size = self.dtype.block_size();
let raw_size = self.numel().div_ceil(block_size) * self.dtype.block_bytes();
raw_size.div_ceil(4) * 4
}
}
pub struct GpuProfiler {
query_set: wgpu::QuerySet,
resolve_buf: wgpu::Buffer,
read_buf: wgpu::Buffer,
timestamp_period: f32, spans: std::sync::Mutex<Vec<(String, u32, u32)>>,
next_query: std::sync::atomic::AtomicU32,
max_queries: u32,
}
pub(crate) struct PendingReadback {
#[cfg(not(target_arch = "wasm32"))]
device: wgpu::Device,
staging: wgpu::Buffer,
size: u64,
rx: futures_channel::oneshot::Receiver<Result<(), wgpu::BufferAsyncError>>,
}
impl PendingReadback {
pub(crate) async fn recv(self) -> Result<Vec<u8>> {
#[cfg(not(target_arch = "wasm32"))]
self.device.poll(wgpu::Maintain::Wait);
self.rx
.await
.map_err(|_| anyhow::anyhow!("GPU readback channel closed"))?
.map_err(|e| anyhow::anyhow!("GPU readback failed: {e:?}"))?;
let slice = self.staging.slice(0..self.size);
let data = slice.get_mapped_range();
let bytes = data.to_vec();
drop(data);
self.staging.unmap();
Ok(bytes)
}
}
impl GpuContext {
pub(crate) fn submit_encoder(&self, enc: wgpu::CommandEncoder) {
io_stats::record_submit();
self.queue.submit(Some(enc.finish()));
}
#[cfg(not(target_arch = "wasm32"))]
pub fn new() -> Result<Self> {
pollster::block_on(Self::new_async())
}
#[cfg(target_arch = "wasm32")]
pub fn new() -> Result<Self> {
anyhow::bail!(
"GpuContext::new() is unavailable on wasm32; use GpuContext::new_async().await"
)
}
pub async fn new_async() -> Result<Self> {
let instance = wgpu::Instance::new(&wgpu::InstanceDescriptor {
backends: wgpu::Backends::all(),
..Default::default()
});
let adapter = instance
.request_adapter(&wgpu::RequestAdapterOptions {
power_preference: wgpu::PowerPreference::HighPerformance,
..Default::default()
})
.await
.context("no GPU adapter found")?;
let adapter_name = adapter.get_info().name.clone();
let backend = format!("{:?}", adapter.get_info().backend);
let profile_requested = std::env::var("CERA_GPU_PROFILE").as_deref() == Ok("1");
let has_timestamps =
profile_requested && adapter.features().contains(wgpu::Features::TIMESTAMP_QUERY);
let mut features = wgpu::Features::empty();
if has_timestamps {
features |= wgpu::Features::TIMESTAMP_QUERY;
}
if adapter.features().contains(wgpu::Features::SHADER_F16) {
features |= wgpu::Features::SHADER_F16;
}
let adapter_limits = adapter.limits();
let (device, queue) = adapter
.request_device(
&wgpu::DeviceDescriptor {
label: Some("cera-gpu"),
required_features: features,
required_limits: adapter_limits.clone(),
memory_hints: wgpu::MemoryHints::Performance,
},
None,
)
.await
.map_err(|e| anyhow::anyhow!("failed to request GPU device: {e}"))?;
let profiler = if has_timestamps {
let max_queries = 512u32; let timestamp_period = queue.get_timestamp_period();
let query_set = device.create_query_set(&wgpu::QuerySetDescriptor {
label: Some("profiler"),
ty: wgpu::QueryType::Timestamp,
count: max_queries,
});
let buf_size = (max_queries as u64) * 8; let resolve_buf = device.create_buffer(&wgpu::BufferDescriptor {
label: Some("profiler-resolve"),
size: buf_size,
usage: wgpu::BufferUsages::QUERY_RESOLVE | wgpu::BufferUsages::COPY_SRC,
mapped_at_creation: false,
});
let read_buf = device.create_buffer(&wgpu::BufferDescriptor {
label: Some("profiler-read"),
size: buf_size,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
tracing::info!("GPU timestamp profiling enabled (period={timestamp_period}ns/tick)");
Some(GpuProfiler {
query_set,
resolve_buf,
read_buf,
timestamp_period,
spans: std::sync::Mutex::new(Vec::new()),
next_query: std::sync::atomic::AtomicU32::new(0),
max_queries,
})
} else {
tracing::info!("GPU timestamp profiling not available");
None
};
let mut preprocessor = Preprocessor::new();
preprocessor.add_include("common_decls.tmpl", shaders::COMMON_DECLS);
preprocessor.add_include("mul_mat_decls.tmpl", shaders::MUL_MAT_DECLS);
tracing::info!(
adapter = %adapter_name,
backend = %backend,
max_storage_buffer_binding_size = adapter_limits.max_storage_buffer_binding_size,
max_buffer_size = adapter_limits.max_buffer_size,
min_subgroup_size = adapter_limits.min_subgroup_size,
"GPU initialized"
);
Ok(Self {
device,
queue,
adapter_name,
backend,
max_storage_buffer_binding_size: adapter_limits.max_storage_buffer_binding_size as u64,
max_buffer_size: adapter_limits.max_buffer_size,
min_storage_buffer_offset_alignment: adapter_limits.min_storage_buffer_offset_alignment
as u64,
preprocessor,
profiler,
staging: std::sync::Mutex::new(None),
staging_size: std::sync::atomic::AtomicU64::new(0),
})
}
pub fn upload_storage(&self, data: &[u8], label: &str) -> wgpu::Buffer {
self.assert_within_max_buffer(data.len() as u64, label);
self.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some(label),
contents: data,
usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_SRC
| wgpu::BufferUsages::COPY_DST,
})
}
pub fn upload_f32(&self, data: &[f32], label: &str) -> wgpu::Buffer {
self.upload_storage(bytemuck::cast_slice(data), label)
}
pub fn upload_f32_as_f16(&self, data: &[f32], label: &str) -> wgpu::Buffer {
let byte_size = (data.len() * 2) as u64;
let aligned_size = byte_size.div_ceil(4) * 4;
self.assert_within_max_buffer(aligned_size, label);
let buffer = self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some(label),
size: aligned_size,
usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_SRC
| wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let chunk_size = 512 * 1024;
for (i, chunk) in data.chunks(chunk_size).enumerate() {
let f16_chunk: Vec<f16> = chunk.iter().map(|&x| f16::from_f32(x)).collect();
let chunk_byte_size = (f16_chunk.len() * 2) as u64;
let aligned_chunk_size = chunk_byte_size.div_ceil(4) * 4;
if aligned_chunk_size > chunk_byte_size {
let mut padded = f16_chunk;
padded.push(f16::ZERO);
self.queue.write_buffer(
&buffer,
(i * chunk_size * 2) as u64,
bytemuck::cast_slice(&padded),
);
} else {
self.queue.write_buffer(
&buffer,
(i * chunk_size * 2) as u64,
bytemuck::cast_slice(&f16_chunk),
);
}
}
buffer
}
pub fn upload_f16(&self, data: &[f16], label: &str) -> wgpu::Buffer {
let size = (data.len() * 2) as u64;
let aligned_size = size.div_ceil(4) * 4;
let buffer = self.create_storage_rw(aligned_size, label);
self.queue
.write_buffer(&buffer, 0, bytemuck::cast_slice(data));
buffer
}
fn assert_within_max_buffer(&self, size: u64, label: &str) {
assert!(
size <= self.max_buffer_size,
"wgpu buffer '{label}' is {size} bytes, exceeding adapter \
max_buffer_size {}; a smaller context or paged KV is required",
self.max_buffer_size
);
}
pub fn create_storage_rw(&self, size: u64, label: &str) -> wgpu::Buffer {
self.assert_within_max_buffer(size, label);
self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some(label),
size,
usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_SRC
| wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
})
}
pub fn download_f32(&self, buffer: &wgpu::Buffer, count: usize) -> Vec<f32> {
use std::sync::atomic::Ordering;
let size = (count * std::mem::size_of::<f32>()) as u64;
let staging_guard = {
let mut guard = self.staging.lock().expect("staging mutex poisoned");
if guard.as_ref().map(|b| b.size() < size).unwrap_or(true) {
*guard = Some(self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("staging-download"),
size,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
}));
self.staging_size.store(size, Ordering::Relaxed);
}
guard
};
let staging = staging_guard.as_ref().unwrap();
let mut encoder = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("download"),
});
encoder.copy_buffer_to_buffer(buffer, 0, staging, 0, size);
self.submit_encoder(encoder);
io_stats::record_readback(size);
let slice = staging.slice(0..size);
let (tx, rx) = std::sync::mpsc::channel();
slice.map_async(wgpu::MapMode::Read, move |result| {
tx.send(result).ok();
});
self.device.poll(wgpu::Maintain::Wait);
rx.recv()
.expect("GPU readback channel closed")
.expect("GPU readback failed");
let data = slice.get_mapped_range();
let result: Vec<f32> = bytemuck::cast_slice(&data).to_vec();
drop(data);
staging.unmap();
result
}
pub(crate) fn begin_download(&self, buffer: &wgpu::Buffer, size: u64) -> PendingReadback {
let staging = self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("staging-download-async"),
size,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let mut encoder = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("download-async"),
});
encoder.copy_buffer_to_buffer(buffer, 0, &staging, 0, size);
self.submit_encoder(encoder);
io_stats::record_readback(size);
let (tx, rx) = futures_channel::oneshot::channel();
staging
.slice(0..size)
.map_async(wgpu::MapMode::Read, move |result| {
let _ = tx.send(result);
});
PendingReadback {
#[cfg(not(target_arch = "wasm32"))]
device: self.device.clone(),
staging,
size,
rx,
}
}
pub async fn download_f32_async(
&self,
buffer: &wgpu::Buffer,
count: usize,
) -> Result<Vec<f32>> {
let size = (count * std::mem::size_of::<f32>()) as u64;
let bytes = self.begin_download(buffer, size).recv().await?;
let mut out = vec![0.0f32; count];
bytemuck::cast_slice_mut(&mut out).copy_from_slice(&bytes);
Ok(out)
}
pub async fn download_u32_async(
&self,
buffer: &wgpu::Buffer,
count: usize,
) -> Result<Vec<u32>> {
let size = (count * std::mem::size_of::<u32>()) as u64;
let bytes = self.begin_download(buffer, size).recv().await?;
let mut out = vec![0u32; count];
bytemuck::cast_slice_mut(&mut out).copy_from_slice(&bytes);
Ok(out)
}
pub fn download_u32(&self, buffer: &wgpu::Buffer, count: usize) -> Vec<u32> {
use std::sync::atomic::Ordering;
let size = (count * std::mem::size_of::<u32>()) as u64;
let staging_guard = {
let mut guard = self.staging.lock().expect("staging mutex poisoned");
if guard.as_ref().map(|b| b.size() < size).unwrap_or(true) {
*guard = Some(self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("staging-download"),
size,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
}));
self.staging_size.store(size, Ordering::Relaxed);
}
guard
};
let staging = staging_guard.as_ref().unwrap();
let mut encoder = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("download_u32"),
});
encoder.copy_buffer_to_buffer(buffer, 0, staging, 0, size);
self.submit_encoder(encoder);
io_stats::record_readback(size);
let slice = staging.slice(0..size);
let (tx, rx) = std::sync::mpsc::channel();
slice.map_async(wgpu::MapMode::Read, move |r| {
tx.send(r).ok();
});
self.device.poll(wgpu::Maintain::Wait);
rx.recv()
.expect("GPU readback channel closed")
.expect("GPU readback failed");
let data = slice.get_mapped_range();
let result: Vec<u32> = bytemuck::cast_slice(&data).to_vec();
drop(data);
staging.unmap();
result
}
pub fn download_f16_as_f32(&self, buffer: &wgpu::Buffer, count: usize) -> Vec<f32> {
use std::sync::atomic::Ordering;
let size = (count * std::mem::size_of::<f16>()) as u64;
let aligned_size = size.div_ceil(4) * 4;
let staging_guard = {
let mut guard = self.staging.lock().expect("staging mutex poisoned");
if guard
.as_ref()
.map(|b| b.size() < aligned_size)
.unwrap_or(true)
{
*guard = Some(self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("staging-download"),
size: aligned_size,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
}));
self.staging_size.store(aligned_size, Ordering::Relaxed);
}
guard
};
let staging = staging_guard.as_ref().unwrap();
let mut encoder = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("download_f16"),
});
encoder.copy_buffer_to_buffer(buffer, 0, staging, 0, aligned_size);
self.submit_encoder(encoder);
io_stats::record_readback(aligned_size);
let slice = staging.slice(0..aligned_size);
let (tx, rx) = std::sync::mpsc::channel();
slice.map_async(wgpu::MapMode::Read, move |r| {
tx.send(r).ok();
});
self.device.poll(wgpu::Maintain::Wait);
rx.recv()
.expect("GPU readback channel closed")
.expect("GPU readback failed");
let data = slice.get_mapped_range();
let f16_data: &[f16] = bytemuck::cast_slice(&data[0..size as usize]);
let result: Vec<f32> = f16_data.iter().map(|&x| x.to_f32()).collect();
drop(data);
staging.unmap();
result
}
pub fn create_pipeline(
&self,
shader_source: &str,
entry_point: &str,
label: &str,
) -> wgpu::ComputePipeline {
self.create_pipeline_with_defines(shader_source, entry_point, label, &[])
}
pub fn create_pipeline_with_defines(
&self,
shader_source: &str,
entry_point: &str,
label: &str,
defines: &[(&str, &str)],
) -> wgpu::ComputePipeline {
let preprocessed = self
.preprocessor
.preprocess(shader_source, defines)
.with_context(|| format!("failed to preprocess shader: {label}"))
.expect("shader preprocessing failed");
let module = self
.device
.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some(label),
source: wgpu::ShaderSource::Wgsl(preprocessed.into()),
});
self.device
.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some(label),
layout: None, module: &module,
entry_point: Some(entry_point),
compilation_options: wgpu::PipelineCompilationOptions {
zero_initialize_workgroup_memory: false,
..Default::default()
},
cache: None,
})
}
pub fn begin_profile_span(&self, label: &str) -> Option<wgpu::ComputePassTimestampWrites<'_>> {
use std::sync::atomic::Ordering;
let profiler = self.profiler.as_ref()?;
let idx = profiler.next_query.load(Ordering::Relaxed);
if idx + 2 > profiler.max_queries {
return None; }
profiler.next_query.store(idx + 2, Ordering::Relaxed);
profiler
.spans
.lock()
.expect("profiler mutex poisoned")
.push((label.to_string(), idx, idx + 1));
Some(wgpu::ComputePassTimestampWrites {
query_set: &profiler.query_set,
beginning_of_pass_write_index: Some(idx),
end_of_pass_write_index: Some(idx + 1),
})
}
pub fn reset_profiler(&self) {
use std::sync::atomic::Ordering;
if let Some(profiler) = &self.profiler {
profiler.next_query.store(0, Ordering::Relaxed);
profiler
.spans
.lock()
.expect("profiler mutex poisoned")
.clear();
}
}
pub fn finish_profiler(&self) {
use std::sync::atomic::Ordering;
let profiler = match &self.profiler {
Some(p) => p,
None => return,
};
let n_queries = profiler.next_query.load(Ordering::Relaxed);
if n_queries == 0 {
return;
}
let mut enc = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
enc.resolve_query_set(&profiler.query_set, 0..n_queries, &profiler.resolve_buf, 0);
enc.copy_buffer_to_buffer(
&profiler.resolve_buf,
0,
&profiler.read_buf,
0,
(n_queries as u64) * 8,
);
self.submit_encoder(enc);
let slice = profiler.read_buf.slice(..((n_queries as u64) * 8));
let (tx, rx) = std::sync::mpsc::channel();
slice.map_async(wgpu::MapMode::Read, move |r| {
tx.send(r).ok();
});
self.device.poll(wgpu::Maintain::Wait);
rx.recv().unwrap().unwrap();
let data = slice.get_mapped_range();
let timestamps: &[u64] = bytemuck::cast_slice(&data);
let period_ns = profiler.timestamp_period as f64;
let spans = profiler.spans.lock().expect("profiler mutex poisoned");
let mut totals: std::collections::HashMap<String, (f64, usize)> =
std::collections::HashMap::new();
for (label, start_idx, end_idx) in spans.iter() {
let start = timestamps[*start_idx as usize];
let end = timestamps[*end_idx as usize];
let us = (end.wrapping_sub(start)) as f64 * period_ns / 1000.0;
let entry = totals.entry(label.clone()).or_insert((0.0, 0));
entry.0 += us;
entry.1 += 1;
}
let mut sorted: Vec<_> = totals.into_iter().collect();
sorted.sort_by(|a, b| b.1.0.partial_cmp(&a.1.0).unwrap());
let total_us: f64 = sorted.iter().map(|(_, (us, _))| us).sum();
eprintln!("── GPU Profile ({total_us:.0}µs total) ──");
for (label, (us, count)) in &sorted {
let pct = us / total_us * 100.0;
eprintln!(" {label:20} {us:8.0}µs ({count:3}×) {pct:5.1}%");
}
drop(data);
profiler.read_buf.unmap();
}
}
pub mod shaders {
pub const COMMON_DECLS: &str = include_str!("shaders/common_decls.tmpl");
pub const MUL_MAT_DECLS: &str = include_str!("shaders/mul_mat_decls.tmpl");
pub const MUL_MAT_REG_TILE: &str = include_str!("shaders/mul_mat_reg_tile.wgsl");
pub const GEMV_F32: &str = include_str!("shaders/gemv_f32.wgsl");
pub const GEMM_F32: &str = include_str!("shaders/gemm_f32.wgsl");
pub const GEMV_Q4_0: &str = include_str!("shaders/gemv_q4_0.wgsl");
pub const GEMV_Q4_0_FAST: &str = include_str!("shaders/gemv_q4_0_fast.wgsl");
pub const GEMV_Q4_K: &str = include_str!("shaders/gemv_q4_k.wgsl");
pub const GEMV_Q5_K: &str = include_str!("shaders/gemv_q5_k.wgsl");
pub const GEMV_Q6_K: &str = include_str!("shaders/gemv_q6_k.wgsl");
pub const GEMV_Q8_0: &str = include_str!("shaders/gemv_q8_0.wgsl");
pub const ELEMENTWISE: &str = include_str!("shaders/elementwise.wgsl");
pub const SCALE_F32: &str = include_str!("shaders/scale_f32.wgsl");
pub const RMSNORM: &str = include_str!("shaders/rmsnorm.wgsl");
pub const RMSNORM_BATCH: &str = include_str!("shaders/rmsnorm_batch.wgsl");
pub const QK_NORM_ROPE_BATCH: &str = include_str!("shaders/qk_norm_rope_batch.wgsl");
pub const CONV1D_FUSED_BATCH: &str = include_str!("shaders/conv1d_fused_batch.wgsl");
pub const PER_HEAD_RMSNORM: &str = include_str!("shaders/per_head_rmsnorm.wgsl");
pub const SOFTMAX: &str = include_str!("shaders/softmax.wgsl");
pub const ARGMAX_F32: &str = include_str!("shaders/argmax_f32.wgsl");
pub const ROPE: &str = include_str!("shaders/rope.wgsl");
pub const KV_SHIFT: &str = include_str!("shaders/kv_shift.wgsl");
pub const FLASH_ATTENTION: &str = include_str!("shaders/flash_attention.wgsl");
pub const ATTENTION_PREFILL: &str = include_str!("shaders/attention_prefill.wgsl");
pub const CONV1D: &str = include_str!("shaders/conv1d.wgsl");
pub const CONV1D_FUSED: &str = include_str!("shaders/conv1d_fused.wgsl");
pub const LAYERNORM_BATCH: &str = include_str!("shaders/layernorm_batch.wgsl");
pub const GELU: &str = include_str!("shaders/gelu.wgsl");
pub const BIAS_ADD: &str = include_str!("shaders/bias_add.wgsl");
pub const VIT_ATTENTION: &str = include_str!("shaders/vit_attention.wgsl");
pub const VIT_ATTENTION_TILED: &str = include_str!("shaders/vit_attention_tiled.wgsl");
}
pub const MAX_WG: u32 = 65535;
pub fn kv_shift_workgroups(total_threads: u32) -> (u32, u32, u32) {
let wg = total_threads.div_ceil(256);
(wg.min(MAX_WG), wg.div_ceil(MAX_WG), 1)
}
pub fn gemv_row_workgroups(count: u32) -> (u32, u32, u32) {
(count.min(MAX_WG), count.div_ceil(MAX_WG), 1)
}
#[derive(Copy, Clone)]
pub struct KvShiftParams {
pub n_keep: u32,
pub shift: u32,
pub retained: u32,
pub n_kv_heads: u32,
pub head_dim: u32,
pub freq_base_bits: u32,
pub rope_type: u32,
pub has_freq_factors: u32,
}
impl KvShiftParams {
pub fn to_u32_array(self) -> [u32; 8] {
[
self.n_keep,
self.shift,
self.retained,
self.n_kv_heads,
self.head_dim,
self.freq_base_bits,
self.rope_type,
self.has_freq_factors,
]
}
pub fn total_threads(self) -> u32 {
self.retained * self.n_kv_heads * (self.head_dim / 2)
}
pub fn dispatch_dims(self) -> (u32, u32, u32) {
kv_shift_workgroups(self.total_threads())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn kv_shift_workgroups_recovers_flat_index_bijectively() {
let totals = [
1u32,
255,
256,
257,
304, MAX_WG * 256, MAX_WG * 256 + 1, (MAX_WG + 7) * 256, 3 * MAX_WG * 256 - 100, ];
for &total in &totals {
let wg = total.div_ceil(256);
let (gx, gy, gz) = kv_shift_workgroups(total);
assert_eq!(gz, 1, "z extent is always 1 (total={total})");
if gy > 1 {
assert_eq!(
gx, MAX_WG,
"X must be pinned to MAX_WG when Y>1 (total={total})"
);
}
assert!(
(gx as u64) * (gy as u64) >= wg as u64,
"grid {gx}x{gy} under-covers wg={wg} (total={total})",
);
for fw in 0..wg {
let x = fw % MAX_WG;
let y = fw / MAX_WG;
assert!(
x < gx && y < gy,
"flat {fw} maps to ({x},{y}) outside grid {gx}x{gy} (total={total})",
);
assert_eq!(
x + y * MAX_WG,
fw,
"get_wid recovery is not the inverse (total={total})"
);
}
}
}
#[test]
fn kv_shift_params_dispatch_dims_match_total_threads() {
let p = KvShiftParams {
n_keep: 2,
shift: 3,
retained: 19,
n_kv_heads: 2,
head_dim: 16,
freq_base_bits: 10_000.0f32.to_bits(),
rope_type: 0,
has_freq_factors: 0,
};
assert_eq!(p.total_threads(), 19 * 2 * (16 / 2)); assert_eq!(p.dispatch_dims(), kv_shift_workgroups(p.total_threads()));
assert_eq!(p.dispatch_dims(), (2, 1, 1));
}
#[test]
fn gemv_row_workgroups_flattening_is_gap_free() {
let counts = [
1u32,
MAX_WG - 1,
MAX_WG, MAX_WG + 1, 2 * MAX_WG + 7, 3 * MAX_WG - 100, ];
for &rg in &counts {
let (gx, gy, gz) = gemv_row_workgroups(rg);
assert_eq!(gz, 1, "z extent is always 1 (row_groups={rg})");
if gy > 1 {
assert_eq!(
gx, MAX_WG,
"X must be pinned to MAX_WG when Y>1 (row_groups={rg})"
);
}
assert!(
(gx as u64) * (gy as u64) >= rg as u64,
"grid {gx}x{gy} under-covers row_groups={rg}",
);
for fw in 0..rg {
let x = fw % MAX_WG;
let y = fw / MAX_WG;
assert!(
x < gx && y < gy,
"flat {fw} maps to ({x},{y}) outside grid {gx}x{gy} (row_groups={rg})",
);
assert_eq!(
x + y * MAX_WG,
fw,
"get_wid recovery is not the inverse (row_groups={rg})"
);
}
}
}
#[test]
fn test_gpu_context_init() {
let ctx = GpuContext::new();
match ctx {
Ok(ctx) => {
println!("GPU: {} ({})", ctx.adapter_name, ctx.backend);
assert!(!ctx.adapter_name.is_empty());
}
Err(e) => {
println!("No GPU available (expected in CI): {e}");
}
}
}
#[test]
fn test_gpu_upload_download_roundtrip() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return, };
let data: Vec<f32> = (0..257).map(|i| i as f32 * 0.1).collect();
let buf = ctx.upload_f32(&data, "test");
let result = ctx.download_f32(&buf, data.len());
assert_eq!(data, result);
}
#[test]
fn test_gpu_async_download_roundtrip() {
let ctx = match pollster::block_on(GpuContext::new_async()) {
Ok(ctx) => ctx,
Err(_) => return, };
let data: Vec<f32> = (0..257).map(|i| i as f32 * 0.1).collect();
let buf = ctx.upload_f32(&data, "test_async");
let result = pollster::block_on(ctx.download_f32_async(&buf, data.len()))
.expect("async readback failed");
assert_eq!(data, result);
}
#[test]
fn test_gpu_f16_roundtrip() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return, };
let data: Vec<f32> = (0..257).map(|i| i as f32 * 0.1).collect();
let buf = ctx.upload_f32_as_f16(&data, "test_f16");
let result = ctx.download_f16_as_f32(&buf, data.len());
for i in 0..data.len() {
let diff = (data[i] - result[i]).abs();
assert!(
diff < 2e-2,
"f16 mismatch at {i}: {} vs {}",
data[i],
result[i]
);
}
}
#[test]
fn test_gpu_gemv_f32() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m = 4u32;
let k = 8u32;
let a: Vec<f32> = (0..m * k).map(|i| (i as f32 - 16.0) * 0.1).collect();
let x: Vec<f32> = (0..k).map(|i| (i as f32 + 1.0) * 0.5).collect();
let mut expected = vec![0.0f32; m as usize];
for i in 0..m as usize {
for j in 0..k as usize {
expected[i] += a[i * k as usize + j] * x[j];
}
}
let a_buf = ctx.upload_f32(&a, "A");
let x_buf = ctx.upload_f32(&x, "x");
let y_buf = ctx.create_storage_rw((m as u64) * 4, "y");
let params = [m, k, 0u32, 0u32];
let params_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(shaders::GEMV_F32, "gemv_f32", "gemv_f32");
let bind_group = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("gemv_f32"),
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: params_buf.as_entire_binding(),
},
],
});
let mut encoder = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
{
let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("gemv_f32"),
timestamp_writes: None,
});
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.dispatch_workgroups(m, 1, 1);
}
ctx.submit_encoder(encoder);
let result = ctx.download_f32(&y_buf, m as usize);
for i in 0..m as usize {
let diff = (expected[i] - result[i]).abs();
assert!(
diff < 1e-3,
"GEMV mismatch at row {i}: cpu={}, gpu={}, diff={diff}",
expected[i],
result[i]
);
}
}
#[test]
fn test_gpu_gemv_f16() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m = 9u32;
let k = 64u32; let a: Vec<f32> = (0..m * k).map(|i| (i as f32 - 288.0) * 0.013).collect();
let x: Vec<f32> = (0..k).map(|i| (i as f32 - 31.0) * 0.05).collect();
let a_f16: Vec<f16> = a.iter().map(|&v| f16::from_f32(v)).collect();
let mut a_packed = Vec::with_capacity((m * k) as usize / 2);
for pair in a_f16.chunks(2) {
let lo = pair[0].to_bits() as u32;
let hi = pair[1].to_bits() as u32;
a_packed.push(lo | (hi << 16));
}
let mut expected = vec![0.0f32; m as usize];
for i in 0..m as usize {
for j in 0..k as usize {
expected[i] += a_f16[i * k as usize + j].to_f32() * x[j];
}
}
let a_buf = ctx.upload_storage(bytemuck::cast_slice(&a_packed), "A_f16");
let x_buf = ctx.upload_f32(&x, "x");
let y_buf = ctx.create_storage_rw((m as u64) * 4, "y");
let params = [m, k, 0u32, 0u32];
let params_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline_with_defines(
shaders::GEMV_F32,
"gemv_f32",
"gemv_f16",
&[("F16_A", "1")],
);
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("gemv_f16"),
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: params_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(m.div_ceil(8), 1, 1);
}
ctx.submit_encoder(enc);
let result = ctx.download_f32(&y_buf, m as usize);
for i in 0..m as usize {
let denom = expected[i].abs().max(1.0);
let rel = (expected[i] - result[i]).abs() / denom;
assert!(
rel < 2e-3,
"f16 GEMV mismatch at row {i}: cpu={}, gpu={}, rel={rel:.2e}",
expected[i],
result[i]
);
}
}
fn run_kernel(
ctx: &GpuContext,
pipeline: &wgpu::ComputePipeline,
bufs: &[&wgpu::Buffer],
workgroups: (u32, u32, u32),
) {
let entries: Vec<wgpu::BindGroupEntry> = bufs
.iter()
.enumerate()
.map(|(i, b)| wgpu::BindGroupEntry {
binding: i as u32,
resource: b.as_entire_binding(),
})
.collect();
let bind_group = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &entries,
});
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
{
let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: None,
timestamp_writes: None,
});
pass.set_pipeline(pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.dispatch_workgroups(workgroups.0, workgroups.1, workgroups.2);
}
ctx.submit_encoder(enc);
ctx.device.poll(wgpu::Maintain::Wait);
}
#[test]
fn test_gpu_layernorm_batch_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let rows = 5usize;
let dim = 320usize; let eps = 1e-6f32;
let src: Vec<f32> = (0..rows * dim)
.map(|i| ((i * 31 + 7) % 197) as f32 * 0.03 - 2.9)
.collect();
let weight: Vec<f32> = (0..dim).map(|i| 0.5 + (i % 7) as f32 * 0.1).collect();
let bias: Vec<f32> = (0..dim).map(|i| (i % 5) as f32 * 0.2 - 0.4).collect();
let mut expected = src.clone();
for r in 0..rows {
crate::backend::cpu::layer_norm_inplace(
&mut expected[r * dim..(r + 1) * dim],
&weight,
&bias,
eps,
);
}
let src_buf = ctx.upload_f32(&src, "src");
let dst_buf = ctx.create_storage_rw((rows * dim * 4) as u64, "dst");
let w_buf = ctx.upload_f32(&weight, "w");
let b_buf = ctx.upload_f32(&bias, "b");
let params = [dim as u32, eps.to_bits(), dim as u32, dim as u32];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(
shaders::LAYERNORM_BATCH,
"layernorm_batch",
"layernorm_batch",
);
run_kernel(
&ctx,
&pipeline,
&[&src_buf, &dst_buf, &w_buf, &b_buf, &p_buf],
(rows as u32, 1, 1),
);
let result = ctx.download_f32(&dst_buf, rows * dim);
for i in 0..rows * dim {
let diff = (expected[i] - result[i]).abs();
assert!(
diff < 1e-3,
"layernorm mismatch at {i}: cpu={}, gpu={}, diff={diff}",
expected[i],
result[i]
);
}
}
#[test]
fn test_gpu_gelu_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let n = 1000usize;
let x: Vec<f32> = (0..n).map(|i| (i as f32 - 500.0) * 0.05).collect();
let mut expected = x.clone();
crate::backend::cpu::gelu_inplace(&mut expected);
let x_buf = ctx.upload_f32(&x, "x");
let params = [n as u32, 0u32];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(shaders::GELU, "gelu_inplace", "gelu");
run_kernel(
&ctx,
&pipeline,
&[&x_buf, &p_buf],
(n.div_ceil(256) as u32, 1, 1),
);
let result = ctx.download_f32(&x_buf, n);
for i in 0..n {
assert!(
result[i].is_finite(),
"gelu produced non-finite at {i} (x={}): {} — tanh overflow?",
x[i],
result[i]
);
let diff = (expected[i] - result[i]).abs();
assert!(
diff < 1e-4,
"gelu mismatch at {i}: cpu={}, gpu={}, diff={diff}",
expected[i],
result[i]
);
}
}
#[test]
fn test_gpu_bias_add_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let rows = 7usize;
let dim = 130usize;
let x: Vec<f32> = (0..rows * dim).map(|i| (i as f32) * 0.01).collect();
let bias: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.05 - 3.0).collect();
let mut expected = x.clone();
for r in 0..rows {
for j in 0..dim {
expected[r * dim + j] += bias[j];
}
}
let x_buf = ctx.upload_f32(&x, "x");
let b_buf = ctx.upload_f32(&bias, "bias");
let total = (rows * dim) as u32;
let params = [total, dim as u32];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(shaders::BIAS_ADD, "bias_add", "bias_add");
run_kernel(
&ctx,
&pipeline,
&[&x_buf, &b_buf, &p_buf],
((total as usize).div_ceil(256) as u32, 1, 1),
);
let result = ctx.download_f32(&x_buf, rows * dim);
for i in 0..rows * dim {
assert!(
(expected[i] - result[i]).abs() < 1e-5,
"bias_add mismatch at {i}: cpu={}, gpu={}",
expected[i],
result[i]
);
}
}
#[test]
fn test_gpu_vit_attention_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let tokens = 20usize;
let n_head = 3usize;
let head_dim = 8usize;
let dim = n_head * head_dim;
let scale = 1.0f32 / (head_dim as f32).sqrt();
let mk = |seed: usize| -> Vec<f32> {
(0..tokens * dim)
.map(|i| (((i + seed) * 37 + 11) % 101) as f32 * 0.02 - 1.0)
.collect()
};
let q = mk(1);
let k = mk(2);
let v = mk(3);
let mut expected = vec![0.0f32; tokens * dim];
for h in 0..n_head {
for qi in 0..tokens {
let q_off = qi * dim + h * head_dim;
let mut scores = vec![0.0f32; tokens];
for (ki, si) in scores.iter_mut().enumerate() {
let k_off = ki * dim + h * head_dim;
let mut s = 0.0f32;
for d in 0..head_dim {
s += q[q_off + d] * k[k_off + d];
}
*si = s * scale;
}
crate::backend::cpu::softmax_inplace(&mut scores);
for d in 0..head_dim {
let mut acc = 0.0f32;
for ki in 0..tokens {
acc += scores[ki] * v[ki * dim + h * head_dim + d];
}
expected[q_off + d] = acc;
}
}
}
let q_buf = ctx.upload_f32(&q, "q");
let k_buf = ctx.upload_f32(&k, "k");
let v_buf = ctx.upload_f32(&v, "v");
let out_buf = ctx.create_storage_rw((tokens * dim * 4) as u64, "out");
let params = [
tokens as u32,
n_head as u32,
head_dim as u32,
scale.to_bits(),
];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline =
ctx.create_pipeline(shaders::VIT_ATTENTION, "vit_attention", "vit_attention");
run_kernel(
&ctx,
&pipeline,
&[&q_buf, &k_buf, &v_buf, &out_buf, &p_buf],
(tokens as u32, n_head as u32, 1),
);
let result = ctx.download_f32(&out_buf, tokens * dim);
for i in 0..tokens * dim {
let diff = (expected[i] - result[i]).abs();
assert!(
diff < 1e-4,
"vit_attention mismatch at {i}: cpu={}, gpu={}, diff={diff}",
expected[i],
result[i]
);
}
}
fn quantize_q4_0_for_test(m: usize, k: usize, weights: &[f32]) -> Vec<u8> {
let mut q4_bytes: Vec<u8> = Vec::with_capacity(m * (k / 32) * 18);
for row in 0..m {
for b in 0..(k / 32) {
let start = row * k + b * 32;
let chunk = &weights[start..start + 32];
let max_abs = chunk.iter().map(|x| x.abs()).fold(0.0f32, f32::max);
let scale = max_abs / 7.0;
let inv = if scale > 0.0 { 1.0 / scale } else { 0.0 };
let d_bits = half::f16::from_f32(scale).to_bits();
q4_bytes.push((d_bits & 0xFF) as u8);
q4_bytes.push((d_bits >> 8) as u8);
for qi in 0..16 {
let lo = ((chunk[qi] * inv).round() + 8.0).clamp(0.0, 15.0) as u8;
let hi = ((chunk[qi + 16] * inv).round() + 8.0).clamp(0.0, 15.0) as u8;
q4_bytes.push(lo | (hi << 4));
}
}
}
q4_bytes
}
fn cpu_matmul_q4_0(m: usize, k: usize, n: usize, q4_bytes: &[u8], x_batch: &[f32]) -> Vec<f32> {
let mut expected = vec![0.0f32; n * m];
for t in 0..n {
let x_slice = &x_batch[t * k..(t + 1) * k];
for row in 0..m {
let mut acc = 0.0f32;
for b in 0..(k / 32) {
let block_off = (row * (k / 32) + b) * 18;
let d_bits = u16::from_le_bytes([q4_bytes[block_off], q4_bytes[block_off + 1]]);
let delta = half::f16::from_bits(d_bits).to_f32();
for qi in 0..16 {
let byte = q4_bytes[block_off + 2 + qi];
let lo = (byte & 0xF) as f32 - 8.0;
let hi = ((byte >> 4) & 0xF) as f32 - 8.0;
acc += lo * delta * x_slice[b * 32 + qi];
acc += hi * delta * x_slice[b * 32 + qi + 16];
}
}
expected[t * m + row] = acc;
}
}
expected
}
#[test]
fn test_gpu_mul_mat_tile_q4_0_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m: u32 = 32;
let k: u32 = 128;
let n: u32 = 16;
let weights_f32: Vec<f32> = (0..m * k)
.map(|i| ((i * 17 + 3) % 29) as f32 * 0.1 - 1.4)
.collect();
let q4_bytes = quantize_q4_0_for_test(m as usize, k as usize, &weights_f32);
let mut x_batch: Vec<f32> = Vec::with_capacity((n * k) as usize);
for t in 0..n {
for i in 0..k {
x_batch.push(((t as f32 + 1.0) * (i as f32 - 64.0)) * 0.05);
}
}
let expected = cpu_matmul_q4_0(m as usize, k as usize, n as usize, &q4_bytes, &x_batch);
let a_buf = ctx.upload_storage(&q4_bytes, "weights");
let x_buf = ctx.upload_f32(&x_batch, "x_batch");
let y_buf = ctx.create_storage_rw(((n * m) as u64) * 4, "y_batch");
let params: [u32; 5] = [m, k, n, k, m];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline_with_defines(
shaders::MUL_MAT_REG_TILE,
"main",
"mul_mat_q4_0_tile_test",
&[
("VEC", ""),
("SRC0_INNER_TYPE", "u32"),
("SRC1_INNER_TYPE", "f32"),
("INIT_SRC0_SHMEM_Q4_0", ""),
("INIT_SRC1_SHMEM_FLOAT", ""),
("WORKGROUP_SIZE_M", "8u"),
("WORKGROUP_SIZE_N", "8u"),
("TILE_M", "4u"),
("TILE_N", "4u"),
("TILE_K", "32u"),
],
);
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: None,
timestamp_writes: None,
});
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
let wg_m = m.div_ceil(8 * 4);
let wg_n = n.div_ceil(8 * 4);
pass.dispatch_workgroups(wg_m, wg_n, 1);
}
ctx.submit_encoder(enc);
ctx.device.poll(wgpu::Maintain::Wait);
let result = ctx.download_f32(&y_buf, (n * m) as usize);
for i in 0..(n * m) as usize {
let diff = (result[i] - expected[i]).abs();
assert!(
diff < 1e-2,
"mismatch at {}: {} vs {}",
i,
result[i],
expected[i]
);
}
}
#[test]
fn reg_tile_parity_tests_match_production_geometry() {
assert_eq!(
(
crate::model::gpu_lfm2::MUL_MAT_TILE_M,
crate::model::gpu_lfm2::MUL_MAT_TILE_N,
),
(8, 4),
"production reg-tile geometry changed; update the hardcoded TILE_M/TILE_N \
defines and dispatch divisors in the mul_mat reg-tile parity tests to match"
);
}
#[test]
fn test_gpu_mul_mat_tile_q4_0_vec_prod_geometry() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m: u32 = 40; let k: u32 = 128;
let n: u32 = 20;
let weights_f32: Vec<f32> = (0..m * k)
.map(|i| ((i * 17 + 3) % 29) as f32 * 0.1 - 1.4)
.collect();
let q4_bytes = quantize_q4_0_for_test(m as usize, k as usize, &weights_f32);
let mut x_batch: Vec<f32> = Vec::with_capacity((n * k) as usize);
for t in 0..n {
for i in 0..k {
x_batch.push(((t as f32 + 1.0) * (i as f32 - 64.0)) * 0.05);
}
}
let expected = cpu_matmul_q4_0(m as usize, k as usize, n as usize, &q4_bytes, &x_batch);
let a_buf = ctx.upload_storage(&q4_bytes, "weights");
let x_buf = ctx.upload_f32(&x_batch, "x_batch");
let y_buf = ctx.create_storage_rw(((n * m) as u64) * 4, "y_batch");
let params: [u32; 5] = [m, k, n, k, m];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline_with_defines(
shaders::MUL_MAT_REG_TILE,
"main",
"mul_mat_q4_0_prod_tile_test",
&[
("VEC", ""),
("SRC0_INNER_TYPE", "u32"),
("SRC1_INNER_TYPE", "f32"),
("INIT_SRC0_SHMEM_Q4_0", ""),
("INIT_SRC1_SHMEM_FLOAT", ""),
("WORKGROUP_SIZE_M", "8u"),
("WORKGROUP_SIZE_N", "32u"),
("TILE_M", "8u"),
("TILE_N", "4u"),
("TILE_K", "32u"),
],
);
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: None,
timestamp_writes: None,
});
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
let wg_m = m.div_ceil(8 * 8);
let wg_n = n.div_ceil(32 * 4);
pass.dispatch_workgroups(wg_m, wg_n, 1);
}
ctx.submit_encoder(enc);
ctx.device.poll(wgpu::Maintain::Wait);
let result = ctx.download_f32(&y_buf, (n * m) as usize);
for i in 0..(n * m) as usize {
let diff = (result[i] - expected[i]).abs();
assert!(
diff < 1e-2,
"mismatch at {}: {} vs {}",
i,
result[i],
expected[i]
);
}
}
#[test]
fn test_gpu_mul_mat_tile_scalar_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m: u32 = 30;
let k: u32 = 128;
let n: u32 = 16;
let x_stride: u32 = k + 3;
let y_stride: u32 = m + 5;
let weights_f32: Vec<f32> = (0..m * k)
.map(|i| ((i * 17 + 3) % 29) as f32 * 0.1 - 1.4)
.collect();
let mut x_batch: Vec<f32> = Vec::with_capacity((n * x_stride) as usize);
for t in 0..n {
for i in 0..k {
x_batch.push(((t as f32 + 1.0) * (i as f32 - 64.0)) * 0.05);
}
x_batch.resize(x_batch.len() + (x_stride - k) as usize, -999.0);
}
let mut expected = vec![0.0f32; (n * y_stride) as usize];
for t in 0..n as usize {
let x_slice = &x_batch[t * x_stride as usize..t * x_stride as usize + k as usize];
for row in 0..m as usize {
let mut acc = 0.0f32;
for col in 0..k as usize {
acc += weights_f32[row * k as usize + col] * x_slice[col];
}
expected[t * y_stride as usize + row] = acc;
}
}
let a_buf = ctx.upload_f32(&weights_f32, "weights");
let x_buf = ctx.upload_f32(&x_batch, "x_batch");
let y_buf = ctx.create_storage_rw(((n * y_stride) as u64) * 4, "y_batch");
let params: [u32; 5] = [m, k, n, x_stride, y_stride];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline_with_defines(
shaders::MUL_MAT_REG_TILE,
"main",
"mul_mat_tile_scalar_test",
&[
("SCALAR", ""),
("SRC0_INNER_TYPE", "f32"),
("SRC1_INNER_TYPE", "f32"),
("INIT_SRC0_SHMEM_FLOAT", ""),
("INIT_SRC1_SHMEM_FLOAT", ""),
("WORKGROUP_SIZE_M", "8u"),
("WORKGROUP_SIZE_N", "8u"),
("TILE_M", "4u"),
("TILE_N", "4u"),
("TILE_K", "32u"),
],
);
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: None,
timestamp_writes: None,
});
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
let wg_m = m.div_ceil(8 * 4);
let wg_n = n.div_ceil(8 * 4);
pass.dispatch_workgroups(wg_m, wg_n, 1);
}
ctx.submit_encoder(enc);
ctx.device.poll(wgpu::Maintain::Wait);
let result = ctx.download_f32(&y_buf, (n * y_stride) as usize);
for t in 0..n as usize {
for row in 0..m as usize {
let i = t * y_stride as usize + row;
let diff = (result[i] - expected[i]).abs();
assert!(
diff < 1e-4,
"mismatch at {}: {} vs {}",
i,
result[i],
expected[i]
);
}
}
}
#[test]
fn test_gpu_mul_mat_tile_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m: u32 = 32;
let k: u32 = 128;
let n: u32 = 16;
let weights_f32: Vec<f32> = (0..m * k)
.map(|i| ((i * 17 + 3) % 29) as f32 * 0.1 - 1.4)
.collect();
let mut x_batch: Vec<f32> = Vec::with_capacity((n * k) as usize);
for t in 0..n {
for i in 0..k {
x_batch.push(((t as f32 + 1.0) * (i as f32 - 64.0)) * 0.05);
}
}
let mut expected = vec![0.0f32; (n * m) as usize];
for t in 0..n as usize {
let x_slice = &x_batch[t * k as usize..(t + 1) * k as usize];
for row in 0..m as usize {
let mut acc = 0.0f32;
for col in 0..k as usize {
acc += weights_f32[row * k as usize + col] * x_slice[col];
}
expected[t * m as usize + row] = acc;
}
}
let a_buf = ctx.upload_f32(&weights_f32, "weights");
let x_buf = ctx.upload_f32(&x_batch, "x_batch");
let y_buf = ctx.create_storage_rw(((n * m) as u64) * 4, "y_batch");
let params: [u32; 5] = [m, k, n, k, m];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline_with_defines(
shaders::MUL_MAT_REG_TILE,
"main",
"mul_mat_tile_test",
&[
("VEC", ""),
("SRC0_INNER_TYPE", "f32"),
("SRC1_INNER_TYPE", "f32"),
("INIT_SRC0_SHMEM_FLOAT", ""),
("INIT_SRC1_SHMEM_FLOAT", ""),
("WORKGROUP_SIZE_M", "8u"),
("WORKGROUP_SIZE_N", "8u"),
("TILE_M", "4u"),
("TILE_N", "4u"),
("TILE_K", "32u"),
],
);
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: None,
timestamp_writes: None,
});
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
let wg_m = m.div_ceil(8 * 4);
let wg_n = n.div_ceil(8 * 4);
pass.dispatch_workgroups(wg_m, wg_n, 1);
}
ctx.submit_encoder(enc);
ctx.device.poll(wgpu::Maintain::Wait);
let result = ctx.download_f32(&y_buf, (n * m) as usize);
for i in 0..(n * m) as usize {
let diff = (result[i] - expected[i]).abs();
assert!(
diff < 1e-4,
"mismatch at {}: {} vs {}",
i,
result[i],
expected[i]
);
}
}
#[test]
fn test_gpu_mul_mat_tile_q4_0_scalar_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m: u32 = 30;
let k: u32 = 128;
let n: u32 = 16;
let weights_f32: Vec<f32> = (0..m * k)
.map(|i| ((i * 17 + 3) % 29) as f32 * 0.1 - 1.4)
.collect();
let q4_bytes = quantize_q4_0_for_test(m as usize, k as usize, &weights_f32);
let mut x_batch: Vec<f32> = Vec::with_capacity((n * k) as usize);
for t in 0..n {
for i in 0..k {
x_batch.push(((t as f32 + 1.0) * (i as f32 - 64.0)) * 0.05);
}
}
let expected = cpu_matmul_q4_0(m as usize, k as usize, n as usize, &q4_bytes, &x_batch);
let a_buf = ctx.upload_storage(&q4_bytes, "weights");
let x_buf = ctx.upload_f32(&x_batch, "x_batch");
let y_buf = ctx.create_storage_rw(((n * m) as u64) * 4, "y_batch");
let params: [u32; 5] = [m, k, n, k, m];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline_with_defines(
shaders::MUL_MAT_REG_TILE,
"main",
"mul_mat_q4_0_scalar_test",
&[
("SCALAR", ""),
("SRC0_INNER_TYPE", "u32"),
("SRC1_INNER_TYPE", "f32"),
("INIT_SRC0_SHMEM_Q4_0", ""),
("INIT_SRC1_SHMEM_FLOAT", ""),
("WORKGROUP_SIZE_M", "8u"),
("WORKGROUP_SIZE_N", "8u"),
("TILE_M", "4u"),
("TILE_N", "4u"),
("TILE_K", "32u"),
],
);
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: None,
timestamp_writes: None,
});
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
let wg_m = m.div_ceil(8 * 4);
let wg_n = n.div_ceil(8 * 4);
pass.dispatch_workgroups(wg_m, wg_n, 1);
}
ctx.submit_encoder(enc);
ctx.device.poll(wgpu::Maintain::Wait);
let result = ctx.download_f32(&y_buf, (n * m) as usize);
for i in 0..(n * m) as usize {
let diff = (result[i] - expected[i]).abs();
assert!(
diff < 1e-2,
"mismatch at {}: {} vs {}",
i,
result[i],
expected[i]
);
}
}
#[test]
fn test_gpu_gemv_f32_realistic() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m = 2816u32;
let k = 1024u32;
let a: Vec<f32> = (0..m * k)
.map(|i| ((i * 17 + 3) % 997) as f32 * 0.001 - 0.5)
.collect();
let x: Vec<f32> = (0..k)
.map(|i| ((i * 13 + 7) % 251) as f32 * 0.01 - 1.25)
.collect();
let mut expected = vec![0.0f32; m as usize];
for i in 0..m as usize {
for j in 0..k as usize {
expected[i] += a[i * k as usize + j] * x[j];
}
}
let a_buf = ctx.upload_f32(&a, "A");
let x_buf = ctx.upload_f32(&x, "x");
let y_buf = ctx.create_storage_rw((m as u64) * 4, "y");
let params = [m, k, 0u32, 0u32];
let params_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(shaders::GEMV_F32, "gemv_f32", "gemv_f32");
let bind_group = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: params_buf.as_entire_binding(),
},
],
});
let mut encoder = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
{
let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: None,
timestamp_writes: None,
});
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.dispatch_workgroups(m, 1, 1);
}
ctx.submit_encoder(encoder);
let result = ctx.download_f32(&y_buf, m as usize);
let mut max_diff = 0.0f32;
for i in 0..m as usize {
let diff = (expected[i] - result[i]).abs();
max_diff = max_diff.max(diff);
assert!(
diff < 0.1, "GEMV mismatch at row {i}: cpu={}, gpu={}, diff={diff}",
expected[i],
result[i]
);
}
println!(
"GPU GEMV 2816×1024: max_diff={max_diff:.6}, all {} rows match",
m
);
}
#[test]
#[ignore] fn bench_gpu_gemv_f32() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m = 2816u32;
let k = 1024u32;
let a: Vec<f32> = (0..m * k)
.map(|i| ((i * 17 + 3) % 997) as f32 * 0.001 - 0.5)
.collect();
let x: Vec<f32> = (0..k)
.map(|i| ((i * 13 + 7) % 251) as f32 * 0.01 - 1.25)
.collect();
let a_buf = ctx.upload_f32(&a, "A");
let x_buf = ctx.upload_f32(&x, "x");
let y_buf = ctx.create_storage_rw((m as u64) * 4, "y");
let params = [m, k, 0u32, 0u32];
let params_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(shaders::GEMV_F32, "gemv_f32", "gemv_f32");
let bind_group = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: params_buf.as_entire_binding(),
},
],
});
for _ in 0..5 {
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.dispatch_workgroups(m, 1, 1);
}
ctx.submit_encoder(enc);
}
ctx.device.poll(wgpu::Maintain::Wait);
let iters = 100;
let start = std::time::Instant::now();
for _ in 0..iters {
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.dispatch_workgroups(m, 1, 1);
}
ctx.submit_encoder(enc);
}
ctx.device.poll(wgpu::Maintain::Wait);
let elapsed = start.elapsed();
let us_per_gemv = elapsed.as_micros() as f64 / iters as f64;
let gflops = (2.0 * m as f64 * k as f64) / (us_per_gemv * 1e3);
let mut cpu_y = vec![0.0f32; m as usize];
let cpu_iters = 1000;
let cpu_start = std::time::Instant::now();
for _ in 0..cpu_iters {
for i in 0..m as usize {
let mut sum = 0.0f32;
for j in 0..k as usize {
sum += a[i * k as usize + j] * x[j];
}
cpu_y[i] = sum;
}
std::hint::black_box(&cpu_y);
}
let cpu_elapsed = cpu_start.elapsed();
let cpu_us = cpu_elapsed.as_micros() as f64 / cpu_iters as f64;
let cpu_gflops = (2.0 * m as f64 * k as f64) / (cpu_us * 1e3);
#[cfg(target_arch = "aarch64")]
let neon_q4_us = {
let nb = k as usize / 32;
let mut q4_bytes = Vec::new();
for row in 0..m as usize {
for b in 0..nb {
let start = row * k as usize + b * 32;
let block = &a[start..start + 32];
let amax = block.iter().map(|v| v.abs()).fold(0.0f32, f32::max);
let d = amax / 7.0;
let d_f16 = half::f16::from_f32(d);
q4_bytes.extend_from_slice(&d_f16.to_bits().to_le_bytes());
let id = if d != 0.0 { 1.0 / d } else { 0.0 };
let mut qs = [0u8; 16];
for i in 0..16 {
let lo = ((block[i] * id + 8.5) as u8).min(15);
let hi = ((block[16 + i] * id + 8.5) as u8).min(15);
qs[i] = lo | (hi << 4);
}
q4_bytes.extend_from_slice(&qs);
}
}
let mut q4_y = vec![0.0f32; m as usize];
let mut q8s = Vec::new();
let mut q8q = Vec::new();
let q4_iters = 1000;
let q4_start = std::time::Instant::now();
for _ in 0..q4_iters {
unsafe {
crate::backend::simd::neon::gemv_q4_0_f32_neon(
&q4_bytes, &x, &mut q4_y, m as usize, k as usize, &mut q8s, &mut q8q,
);
}
std::hint::black_box(&q4_y);
}
let q4_elapsed = q4_start.elapsed();
q4_elapsed.as_micros() as f64 / q4_iters as f64
};
#[cfg(not(target_arch = "aarch64"))]
let neon_q4_us = 0.0;
let neon_q4_gflops = if neon_q4_us > 0.0 {
(2.0 * m as f64 * k as f64) / (neon_q4_us * 1e3)
} else {
0.0
};
println!(
"GEMV {m}×{k}:\n GPU(f32 Metal) = {us_per_gemv:.0}µs ({gflops:.1} GFLOPS)\n CPU(scalar f32) = {cpu_us:.0}µs ({cpu_gflops:.1} GFLOPS)\n CPU(NEON Q4_0) = {neon_q4_us:.0}µs ({neon_q4_gflops:.1} GFLOPS)\n GPU vs scalar: {:.1}x\n GPU vs NEON Q4_0: {:.1}x",
cpu_us / us_per_gemv,
if neon_q4_us > 0.0 {
neon_q4_us / us_per_gemv
} else {
0.0
},
);
}
#[test]
fn test_gpu_elementwise_add() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let n = 1024u32;
let a: Vec<f32> = (0..n).map(|i| i as f32 * 0.1).collect();
let b: Vec<f32> = (0..n).map(|i| (n - i) as f32 * 0.05).collect();
let expected: Vec<f32> = a.iter().zip(b.iter()).map(|(x, y)| x + y).collect();
let a_buf = ctx.create_storage_rw((n as u64) * 4, "a");
ctx.queue.write_buffer(&a_buf, 0, bytemuck::cast_slice(&a));
let b_buf = ctx.upload_f32(&b, "b");
let params = [n, 0u32];
let params_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(shaders::ELEMENTWISE, "add_inplace", "add_inplace");
let bind_group = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: b_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: params_buf.as_entire_binding(),
},
],
});
let mut encoder = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
{
let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: None,
timestamp_writes: None,
});
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.dispatch_workgroups(n.div_ceil(256), 1, 1);
}
ctx.submit_encoder(encoder);
let result = ctx.download_f32(&a_buf, n as usize);
for i in 0..n as usize {
let diff = (expected[i] - result[i]).abs();
assert!(diff < 1e-5, "add mismatch at {i}: {diff}");
}
}
#[test]
fn test_gpu_silu_mul() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let n = 512u32;
let gate: Vec<f32> = (0..n).map(|i| (i as f32 - 256.0) * 0.02).collect();
let up: Vec<f32> = (0..n).map(|i| (i as f32 + 1.0) * 0.1).collect();
let expected: Vec<f32> = gate
.iter()
.zip(up.iter())
.map(|(&g, &u)| (g / (1.0 + (-g).exp())) * u)
.collect();
let gate_buf = ctx.create_storage_rw((n as u64) * 4, "gate");
ctx.queue
.write_buffer(&gate_buf, 0, bytemuck::cast_slice(&gate));
let up_buf = ctx.upload_f32(&up, "up");
let params = [n, 0u32];
let params_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(shaders::ELEMENTWISE, "silu_mul_inplace", "silu_mul");
let bind_group = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: gate_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: up_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: params_buf.as_entire_binding(),
},
],
});
let mut encoder = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
{
let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: None,
timestamp_writes: None,
});
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.dispatch_workgroups(n.div_ceil(256), 1, 1);
}
ctx.submit_encoder(encoder);
let result = ctx.download_f32(&gate_buf, n as usize);
for i in 0..n as usize {
let diff = (expected[i] - result[i]).abs();
assert!(
diff < 1e-4,
"silu_mul mismatch at {i}: cpu={}, gpu={}, diff={diff}",
expected[i],
result[i]
);
}
}
#[test]
fn test_gpu_rmsnorm() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let n = 1024u32;
let eps = 1e-5f32;
let x: Vec<f32> = (0..n).map(|i| (i as f32 - 512.0) * 0.01).collect();
let weight: Vec<f32> = (0..n).map(|i| 0.8 + (i as f32 % 7.0) * 0.05).collect();
let mut expected = x.clone();
crate::backend::cpu::rmsnorm(&mut expected, &weight, eps);
let x_buf = ctx.create_storage_rw((n as u64) * 4, "x");
ctx.queue.write_buffer(&x_buf, 0, bytemuck::cast_slice(&x));
let w_buf = ctx.upload_f32(&weight, "w");
let params = [n, eps.to_bits(), 0u32, 0u32];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(shaders::RMSNORM, "rmsnorm", "rmsnorm");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: w_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: p_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(1, 1, 1);
}
ctx.submit_encoder(enc);
let result = ctx.download_f32(&x_buf, n as usize);
for i in 0..n as usize {
let diff = (expected[i] - result[i]).abs();
assert!(
diff < 1e-3,
"rmsnorm mismatch at {i}: cpu={}, gpu={}",
expected[i],
result[i]
);
}
}
#[test]
fn test_gpu_softmax() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let n = 128u32;
let x: Vec<f32> = (0..n).map(|i| (i as f32 - 64.0) * 0.1).collect();
let mut expected = x.clone();
crate::backend::cpu::softmax_inplace(&mut expected);
let x_buf = ctx.create_storage_rw((n as u64) * 4, "x");
ctx.queue.write_buffer(&x_buf, 0, bytemuck::cast_slice(&x));
let params = [n, 0u32];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(shaders::SOFTMAX, "softmax", "softmax");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: p_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(1, 1, 1);
}
ctx.submit_encoder(enc);
let result = ctx.download_f32(&x_buf, n as usize);
let sum: f32 = result.iter().sum();
assert!(
(sum - 1.0).abs() < 1e-4,
"softmax sum should be 1.0, got {sum}"
);
for i in 0..n as usize {
let diff = (expected[i] - result[i]).abs();
assert!(
diff < 1e-5,
"softmax mismatch at {i}: cpu={}, gpu={}",
expected[i],
result[i]
);
}
}
#[test]
fn test_gpu_gemv_q4_0() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m = 8u32;
let k = 64u32;
let nb = k / 32;
let weights_f32: Vec<f32> = (0..m * k)
.map(|i| ((i * 17 + 3) % 29) as f32 * 0.1 - 1.4)
.collect();
let mut q4_bytes: Vec<u8> = Vec::new();
for row in 0..m as usize {
for b in 0..nb as usize {
let start = row * k as usize + b * 32;
let block = &weights_f32[start..start + 32];
let amax = block.iter().map(|v| v.abs()).fold(0.0f32, f32::max);
let d = amax / 7.0;
let d_f16 = half::f16::from_f32(d);
q4_bytes.extend_from_slice(&d_f16.to_bits().to_le_bytes());
let id = if d != 0.0 { 1.0 / d } else { 0.0 };
for qi in 0..16 {
let lo = ((block[qi] * id + 8.5) as u8).min(15);
let hi = ((block[16 + qi] * id + 8.5) as u8).min(15);
q4_bytes.push(lo | (hi << 4));
}
}
}
let x: Vec<f32> = (0..k).map(|i| (i as f32 - 32.0) * 0.05).collect();
let mut expected = vec![0.0f32; m as usize];
for (row, exp) in expected.iter_mut().enumerate() {
for b in 0..nb as usize {
let block_off = (row * nb as usize + b) * 18;
let d_bits = u16::from_le_bytes([q4_bytes[block_off], q4_bytes[block_off + 1]]);
let delta = half::f16::from_bits(d_bits).to_f32();
for qi in 0..16 {
let byte = q4_bytes[block_off + 2 + qi];
let lo = (byte & 0xF) as f32 - 8.0;
let hi = ((byte >> 4) & 0xF) as f32 - 8.0;
*exp += lo * delta * x[b * 32 + qi];
*exp += hi * delta * x[b * 32 + qi + 16];
}
}
}
let a_buf = ctx.upload_storage(&q4_bytes, "A_q4");
let x_buf = ctx.upload_f32(&x, "x");
let y_buf = ctx.create_storage_rw((m as u64) * 4, "y");
let params = [m, k, 0u32, 0u32];
let params_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(shaders::GEMV_Q4_0, "gemv_q4_0", "gemv_q4_0");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: params_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(m.div_ceil(8), 1, 1);
}
ctx.submit_encoder(enc);
let result = ctx.download_f32(&y_buf, m as usize);
for i in 0..m as usize {
let diff = (expected[i] - result[i]).abs();
assert!(
diff < 0.5,
"Q4_0 GEMV mismatch at row {i}: cpu={}, gpu={}, diff={diff}",
expected[i],
result[i]
);
}
println!("Q4_0 GEMV {m}×{k}: all rows match");
}
#[test]
fn test_gpu_gemv_q8_0() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m = 9u32;
let k = 64u32;
let nb = k / 32;
let weights_f32: Vec<f32> = (0..m * k)
.map(|i| ((i * 13 + 5) % 41) as f32 * 0.07 - 1.3)
.collect();
let mut q8_bytes: Vec<u8> = Vec::new();
let mut expected = vec![0.0f32; m as usize];
let x: Vec<f32> = (0..k).map(|i| (i as f32 - 17.0) * 0.03125).collect();
for (row, exp) in expected.iter_mut().enumerate() {
for b in 0..nb as usize {
let start = row * k as usize + b * 32;
let block = &weights_f32[start..start + 32];
let amax = block.iter().map(|v| v.abs()).fold(0.0f32, f32::max);
let d = if amax != 0.0 { amax / 127.0 } else { 0.0 };
let d_f16 = half::f16::from_f32(d);
q8_bytes.extend_from_slice(&d_f16.to_bits().to_le_bytes());
let id = if d != 0.0 { 1.0 / d } else { 0.0 };
for (qi, &value) in block.iter().enumerate() {
let quant = (value * id).round().clamp(-127.0, 127.0) as i8;
q8_bytes.push(quant as u8);
*exp += f32::from(quant) * d_f16.to_f32() * x[b * 32 + qi];
}
}
}
let a_buf = ctx.upload_storage(&q8_bytes, "A_q8");
let x_buf = ctx.upload_f32(&x, "x");
let y_buf = ctx.create_storage_rw((m as u64) * 4, "y");
let params = [m, k, 0u32, 0u32];
let params_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(shaders::GEMV_Q8_0, "gemv_q8_0", "gemv_q8_0");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: params_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(m.div_ceil(8), 1, 1);
}
ctx.submit_encoder(enc);
let result = ctx.download_f32(&y_buf, m as usize);
for i in 0..m as usize {
let diff = (expected[i] - result[i]).abs();
assert!(
diff < 1e-3,
"Q8_0 GEMV mismatch at row {i}: cpu={}, gpu={}, diff={diff}",
expected[i],
result[i]
);
}
}
#[test]
fn test_gpu_gemv_q6_k() {
use crate::quant::{BlockQ6K, dequantize_q6_k_block};
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m = 65u32;
let k = 512u32; let qk_k = 256usize;
let nb = k as usize / qk_k;
let mut raw = Vec::with_capacity(m as usize * nb * 210);
let mut expected_f32 = vec![0.0f32; m as usize * k as usize];
for row in 0..m as usize {
for b in 0..nb {
let mut blk = BlockQ6K {
ql: [0u8; 128],
qh: [0u8; 64],
scales: [0i8; 16],
d: half::f16::from_f32(0.01 + (row as f32 * 0.003).sin() * 0.002).to_bits(),
};
for (i, v) in blk.ql.iter_mut().enumerate() {
*v = ((row * 37 + b * 13 + i) & 0xFF) as u8;
}
for (i, v) in blk.qh.iter_mut().enumerate() {
*v = ((row * 11 + b * 7 + i) & 0xFF) as u8;
}
for (i, v) in blk.scales.iter_mut().enumerate() {
*v = (((row * 3 + b * 5 + i) as i32 & 0x7F) - 32) as i8;
}
let dq = dequantize_q6_k_block(&blk);
let row_off = row * k as usize + b * qk_k;
expected_f32[row_off..row_off + qk_k].copy_from_slice(&dq);
raw.extend_from_slice(&blk.ql);
raw.extend_from_slice(&blk.qh);
raw.extend_from_slice(bytemuck::cast_slice(&blk.scales));
raw.extend_from_slice(&blk.d.to_le_bytes());
}
}
let x: Vec<f32> = (0..k).map(|i| (i as f32 * 0.013).sin()).collect();
let mut expected = vec![0.0f32; m as usize];
for (row, exp) in expected.iter_mut().enumerate() {
let mut s = 0.0f32;
for i in 0..k as usize {
s += expected_f32[row * k as usize + i] * x[i];
}
*exp = s;
}
let a_buf = ctx.upload_storage(&raw, "A_q6k");
let x_buf = ctx.upload_f32(&x, "x");
let y_buf = ctx.create_storage_rw((m as u64) * 4, "y");
let params = [m, k, 0u32, 0u32];
let params_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(shaders::GEMV_Q6_K, "gemv_q6_k", "gemv_q6_k");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: params_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(m.div_ceil(2), 1, 1);
}
ctx.submit_encoder(enc);
let result = ctx.download_f32(&y_buf, m as usize);
for i in 0..m as usize {
let denom = expected[i].abs().max(1.0);
let rel = (expected[i] - result[i]).abs() / denom;
assert!(
rel < 5e-3,
"Q6_K GEMV mismatch at row {i}: cpu={}, gpu={}, rel={rel:.2e}",
expected[i],
result[i]
);
}
}
#[test]
fn test_gpu_gemv_q4_k() {
use crate::quant::{BlockQ4KM, dequantize_q4_k_m_block};
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m = 65u32;
let k = 512u32; let qk_k = 256usize;
let nb = k as usize / qk_k;
let mut raw = Vec::with_capacity(m as usize * nb * 144);
let mut expected_f32 = vec![0.0f32; m as usize * k as usize];
for row in 0..m as usize {
for b in 0..nb {
let mut blk = BlockQ4KM {
d: half::f16::from_f32(0.02 + (row as f32 * 0.004).sin() * 0.003).to_bits(),
dmin: half::f16::from_f32(0.01 + (b as f32 * 0.002)).to_bits(),
scales: [0u8; 12],
qs: [0u8; 128],
};
for (i, v) in blk.scales.iter_mut().enumerate() {
*v = ((row * 5 + b * 7 + i * 3) & 0xFF) as u8;
}
for (i, v) in blk.qs.iter_mut().enumerate() {
*v = ((row * 37 + b * 13 + i) & 0xFF) as u8;
}
let dq = dequantize_q4_k_m_block(&blk);
let row_off = row * k as usize + b * qk_k;
expected_f32[row_off..row_off + qk_k].copy_from_slice(&dq);
raw.extend_from_slice(&blk.d.to_le_bytes());
raw.extend_from_slice(&blk.dmin.to_le_bytes());
raw.extend_from_slice(&blk.scales);
raw.extend_from_slice(&blk.qs);
}
}
let x: Vec<f32> = (0..k).map(|i| (i as f32 * 0.017).cos()).collect();
let mut expected = vec![0.0f32; m as usize];
for (row, exp) in expected.iter_mut().enumerate() {
let mut s = 0.0f32;
for i in 0..k as usize {
s += expected_f32[row * k as usize + i] * x[i];
}
*exp = s;
}
let a_buf = ctx.upload_storage(&raw, "A_q4k");
let x_buf = ctx.upload_f32(&x, "x");
let y_buf = ctx.create_storage_rw((m as u64) * 4, "y");
let params = [m, k, 0u32, 0u32];
let params_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(shaders::GEMV_Q4_K, "gemv_q4_k", "gemv_q4_k");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: params_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(m.div_ceil(2), 1, 1);
}
ctx.submit_encoder(enc);
let result = ctx.download_f32(&y_buf, m as usize);
for i in 0..m as usize {
let denom = expected[i].abs().max(1.0);
let rel = (expected[i] - result[i]).abs() / denom;
assert!(
rel < 5e-3,
"Q4_K GEMV mismatch at row {i}: cpu={}, gpu={}, rel={rel:.2e}",
expected[i],
result[i]
);
}
}
#[test]
fn test_gpu_gemv_q5_k() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m = 5u32;
let k = 512u32;
let nb = (k / 256) as usize;
let bpb = 176usize;
let n_blocks = m as usize * nb;
let mut a_bytes = vec![0u8; n_blocks * bpb];
for (bi, chunk) in a_bytes.chunks_mut(bpb).enumerate() {
let d = half::f16::from_f32(0.015 + (bi % 5) as f32 * 0.004);
let dmin = half::f16::from_f32(0.008 + (bi % 3) as f32 * 0.003);
chunk[0..2].copy_from_slice(&d.to_bits().to_le_bytes());
chunk[2..4].copy_from_slice(&dmin.to_bits().to_le_bytes());
for (i, b) in chunk[4..16].iter_mut().enumerate() {
*b = ((bi * 7 + i * 13 + 1) % 256) as u8; }
for (i, b) in chunk[16..48].iter_mut().enumerate() {
*b = ((bi * 29 + i * 7) % 256) as u8; }
for (i, b) in chunk[48..176].iter_mut().enumerate() {
*b = ((bi * 17 + i * 5) % 256) as u8; }
}
let x: Vec<f32> = (0..k).map(|i| (i as f32 - 100.0) * 0.01).collect();
let mut expected = vec![0.0f32; m as usize];
for (row, exp) in expected.iter_mut().enumerate() {
let row_bytes = &a_bytes[row * nb * bpb..(row + 1) * nb * bpb];
let mut deq = vec![0.0f32; k as usize];
crate::quant::dequantize_q5_k_row(row_bytes, &mut deq);
*exp = deq.iter().zip(x.iter()).map(|(a, b)| a * b).sum();
}
let a_buf = ctx.upload_storage(&a_bytes, "A_q5k");
let x_buf = ctx.upload_f32(&x, "x");
let y_buf = ctx.create_storage_rw((m as u64) * 4, "y");
let params = [m, k, 0u32, 0u32];
let params_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let pipeline = ctx.create_pipeline(shaders::GEMV_Q5_K, "gemv_q5_k", "gemv_q5_k");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: params_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(m.div_ceil(2), 1, 1); }
ctx.submit_encoder(enc);
let result = ctx.download_f32(&y_buf, m as usize);
for i in 0..m as usize {
let diff = (expected[i] - result[i]).abs();
let tol = 1e-3 * expected[i].abs().max(1.0);
assert!(
diff <= tol,
"Q5_K GEMV mismatch at row {i}: cpu={}, gpu={}, diff={diff}",
expected[i],
result[i]
);
}
}
#[test]
fn test_gpu_flash_attention_matches_cpu() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let flash = ctx.create_pipeline(
shaders::FLASH_ATTENTION,
"flash_attention",
"flash_attention",
);
let run = |pipeline: &wgpu::ComputePipeline,
bindings: &[&wgpu::Buffer],
out: &wgpu::Buffer,
out_len: usize,
n_heads: u32|
-> Vec<f32> {
let entries: Vec<wgpu::BindGroupEntry> = bindings
.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: &pipeline.get_bind_group_layout(0),
entries: &entries,
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(n_heads, 1, 1);
}
ctx.submit_encoder(enc);
ctx.download_f32(out, out_len)
};
let configs = [
(4u32, 4u32, 64u32, 1u32),
(4, 4, 64, 7),
(8, 2, 64, 300),
(4, 4, 128, 256),
(6, 3, 128, 1000),
];
for (n_heads, n_kv_heads, head_dim, seq_len) in configs {
let kv_dim = n_kv_heads * head_dim;
let scale = 1.0f32 / (head_dim as f32).sqrt();
let q: Vec<f32> = (0..n_heads * head_dim)
.map(|i| ((i * 7 + 3) % 17) as f32 * 0.1 - 0.8)
.collect();
let k: Vec<f32> = (0..seq_len * kv_dim)
.map(|i| ((i * 13 + 5) % 23) as f32 * 0.05 - 0.55)
.collect();
let v: Vec<f32> = (0..seq_len * kv_dim)
.map(|i| ((i * 11 + 1) % 19) as f32 * 0.05 - 0.45)
.collect();
let params: [u32; 8] = [
n_heads,
n_kv_heads,
head_dim,
kv_dim,
seq_len,
scale.to_bits(),
0,
0,
];
let q_buf = ctx.upload_f32(&q, "q");
let k_buf = ctx.upload_f32(&k, "k");
let v_buf = ctx.upload_f32(&v, "v");
let params_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let out_len = (n_heads * head_dim) as usize;
let out_f = ctx.create_storage_rw(out_len as u64 * 4, "out_flash");
let flash_out = run(
&flash,
&[&q_buf, &k_buf, &v_buf, &out_f, ¶ms_buf],
&out_f,
out_len,
n_heads,
);
let gs = (n_heads / n_kv_heads) as usize;
let hd = head_dim as usize;
let kvd = kv_dim as usize;
let sl = seq_len as usize;
let mut cpu_ref = vec![0.0f32; out_len];
for h in 0..n_heads as usize {
let kvo = (h / gs) * hd;
let mut scores = vec![0.0f32; sl];
let mut mx = f32::NEG_INFINITY;
for (t, s) in scores.iter_mut().enumerate() {
let mut dot = 0.0f32;
for d in 0..hd {
dot += q[h * hd + d] * k[t * kvd + kvo + d];
}
*s = dot * scale;
mx = mx.max(*s);
}
let mut sum = 0.0f32;
for s in scores.iter_mut() {
*s = (*s - mx).exp();
sum += *s;
}
for d in 0..hd {
let mut a = 0.0f32;
for (t, s) in scores.iter().enumerate() {
a += s * v[t * kvd + kvo + d];
}
cpu_ref[h * hd + d] = a / sum;
}
}
for i in 0..out_len {
let diff = (flash_out[i] - cpu_ref[i]).abs();
let tol = 1e-3 + 1e-3 * cpu_ref[i].abs();
assert!(
diff <= tol,
"flash≠cpu at cfg (h={n_heads},kv={n_kv_heads},hd={head_dim},\
seq={seq_len}) idx {i}: cpu={}, flash={}, diff={diff}",
cpu_ref[i],
flash_out[i],
);
}
}
}
#[test]
fn test_gpu_mul_mat_q8_0_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m = 11u32;
let k = 64u32;
let n = 3u32;
let x_stride = k + 4;
let y_stride = m + 5;
let nb = k / 32;
let weights_f32: Vec<f32> = (0..m * k)
.map(|i| ((i * 17 + 3) % 53) as f32 * 0.045 - 1.1)
.collect();
let mut q8_bytes: Vec<u8> = Vec::new();
for row in 0..m as usize {
for b in 0..nb as usize {
let start = row * k as usize + b * 32;
let block = &weights_f32[start..start + 32];
let amax = block.iter().map(|v| v.abs()).fold(0.0f32, f32::max);
let d = if amax != 0.0 { amax / 127.0 } else { 0.0 };
let d_f16 = half::f16::from_f32(d);
q8_bytes.extend_from_slice(&d_f16.to_bits().to_le_bytes());
let id = if d != 0.0 { 1.0 / d } else { 0.0 };
for &value in block {
let quant = (value * id).round().clamp(-127.0, 127.0) as i8;
q8_bytes.push(quant as u8);
}
}
}
let mut x_batch = vec![0.0f32; (n * x_stride) as usize];
for t in 0..n as usize {
for i in 0..k as usize {
x_batch[t * x_stride as usize + i] = ((t as f32 + 1.0) * (i as f32 - 19.0)) * 0.021;
}
}
let mut expected = vec![0.0f32; (n * y_stride) as usize];
for t in 0..n as usize {
let x_slice = &x_batch[t * x_stride as usize..t * x_stride as usize + k as usize];
for row in 0..m as usize {
let mut acc = 0.0f32;
for b in 0..nb as usize {
let block_off = (row * nb as usize + b) * 34;
let d_bits = u16::from_le_bytes([q8_bytes[block_off], q8_bytes[block_off + 1]]);
let d = half::f16::from_bits(d_bits).to_f32();
for qi in 0..32 {
let quant = q8_bytes[block_off + 2 + qi] as i8;
acc += f32::from(quant) * d * x_slice[b * 32 + qi];
}
}
expected[t * y_stride as usize + row] = acc;
}
}
let a_buf = ctx.upload_storage(&q8_bytes, "mm_q8_weights");
let x_buf = ctx.upload_f32(&x_batch, "mm_q8_x");
let y_buf = ctx.create_storage_rw(((n * y_stride) as u64) * 4, "mm_q8_y");
let params: [u32; 5] = [m, k, n, x_stride, y_stride];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "mm_q8_params");
let pipeline = ctx.create_pipeline_with_defines(
shaders::MUL_MAT_REG_TILE,
"main",
"mul_mat_q8_0_test",
&[
("SCALAR", ""),
("SRC0_INNER_TYPE", "u32"),
("SRC1_INNER_TYPE", "f32"),
("INIT_SRC0_SHMEM_Q8_0", ""),
("INIT_SRC1_SHMEM_FLOAT", ""),
("WORKGROUP_SIZE_M", "8u"),
("WORKGROUP_SIZE_N", "32u"),
("TILE_M", "8u"),
("TILE_N", "4u"),
("TILE_K", "32u"),
],
);
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(m.div_ceil(64), n.div_ceil(128), 1);
}
ctx.submit_encoder(enc);
let got = ctx.download_f32(&y_buf, (n * y_stride) as usize);
for t in 0..n as usize {
for row in 0..m as usize {
let idx = t * y_stride as usize + row;
let diff = (expected[idx] - got[idx]).abs();
assert!(
diff < 1e-3,
"Q8_0 GEMM mismatch at token {t}, row {row}: cpu={}, gpu={}, diff={diff}",
expected[idx],
got[idx]
);
}
}
}
#[test]
fn test_gpu_mul_mat_q4_k_parity() {
use crate::quant::{BlockQ4KM, dequantize_q4_k_m_block};
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m = 37u32; let k = 512u32;
let n = 5u32;
let x_stride = k + 4;
let y_stride = m + 5;
let qk_k = 256usize;
let nb = k as usize / qk_k;
let mut raw = Vec::with_capacity(m as usize * nb * 144);
let mut w_f32 = vec![0.0f32; m as usize * k as usize];
for row in 0..m as usize {
for b in 0..nb {
let mut blk = BlockQ4KM {
d: half::f16::from_f32(0.02 + (row as f32 * 0.004).sin() * 0.003).to_bits(),
dmin: half::f16::from_f32(0.01 + (b as f32 * 0.002)).to_bits(),
scales: [0u8; 12],
qs: [0u8; 128],
};
for (i, v) in blk.scales.iter_mut().enumerate() {
*v = ((row * 5 + b * 7 + i * 3) & 0xFF) as u8;
}
for (i, v) in blk.qs.iter_mut().enumerate() {
*v = ((row * 37 + b * 13 + i) & 0xFF) as u8;
}
let dq = dequantize_q4_k_m_block(&blk);
let off = row * k as usize + b * qk_k;
w_f32[off..off + qk_k].copy_from_slice(&dq);
raw.extend_from_slice(&blk.d.to_le_bytes());
raw.extend_from_slice(&blk.dmin.to_le_bytes());
raw.extend_from_slice(&blk.scales);
raw.extend_from_slice(&blk.qs);
}
}
let mut x_batch = vec![0.0f32; (n * x_stride) as usize];
for t in 0..n as usize {
for i in 0..k as usize {
x_batch[t * x_stride as usize + i] =
((t as f32 + 1.0) * (i as f32 - 200.0)) * 0.0007;
}
}
let mut expected = vec![0.0f32; (n * y_stride) as usize];
for t in 0..n as usize {
let x_slice = &x_batch[t * x_stride as usize..t * x_stride as usize + k as usize];
for row in 0..m as usize {
let mut acc = 0.0f32;
for i in 0..k as usize {
acc += w_f32[row * k as usize + i] * x_slice[i];
}
expected[t * y_stride as usize + row] = acc;
}
}
let a_buf = ctx.upload_storage(&raw, "mm_q4k_weights");
let x_buf = ctx.upload_f32(&x_batch, "mm_q4k_x");
let y_buf = ctx.create_storage_rw(((n * y_stride) as u64) * 4, "mm_q4k_y");
let params: [u32; 5] = [m, k, n, x_stride, y_stride];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "mm_q4k_params");
let pipeline = ctx.create_pipeline_with_defines(
shaders::MUL_MAT_REG_TILE,
"main",
"mul_mat_q4_k_test",
&[
("SCALAR", ""),
("SRC0_INNER_TYPE", "u32"),
("SRC1_INNER_TYPE", "f32"),
("INIT_SRC0_SHMEM_Q4_K", ""),
("INIT_SRC1_SHMEM_FLOAT", ""),
("WORKGROUP_SIZE_M", "8u"),
("WORKGROUP_SIZE_N", "32u"),
("TILE_M", "8u"),
("TILE_N", "4u"),
("TILE_K", "32u"),
],
);
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(m.div_ceil(64), n.div_ceil(128), 1);
}
ctx.submit_encoder(enc);
let got = ctx.download_f32(&y_buf, (n * y_stride) as usize);
for t in 0..n as usize {
for row in 0..m as usize {
let idx = t * y_stride as usize + row;
let denom = expected[idx].abs().max(1.0);
let rel = (expected[idx] - got[idx]).abs() / denom;
assert!(
rel < 5e-3,
"Q4_K reg-tile mismatch at token {t}, row {row}: cpu={}, gpu={}, rel={rel:.2e}",
expected[idx],
got[idx]
);
}
}
}
#[test]
fn test_gpu_gemm_q6_k_parity() {
use crate::quant::{BlockQ6K, dequantize_q6_k_block};
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m = 11u32;
let k = 512u32; let n = 3u32;
let x_stride = k + 4;
let y_stride = m + 5;
let qk_k = 256usize;
let nb = k as usize / qk_k;
let mut raw = Vec::with_capacity(m as usize * nb * 210);
let mut w_f32 = vec![0.0f32; m as usize * k as usize];
for row in 0..m as usize {
for b in 0..nb {
let mut blk = BlockQ6K {
ql: [0u8; 128],
qh: [0u8; 64],
scales: [0i8; 16],
d: half::f16::from_f32(0.02 + (row as f32 * 0.004).sin() * 0.003).to_bits(),
};
for (i, v) in blk.ql.iter_mut().enumerate() {
*v = ((row * 37 + b * 13 + i) & 0xFF) as u8;
}
for (i, v) in blk.qh.iter_mut().enumerate() {
*v = ((row * 17 + b * 29 + i * 5) & 0xFF) as u8;
}
for (i, v) in blk.scales.iter_mut().enumerate() {
*v = (((row * 5 + b * 7 + i * 3) % 97) as i32 - 48) as i8;
}
let dq = dequantize_q6_k_block(&blk);
let off = row * k as usize + b * qk_k;
w_f32[off..off + qk_k].copy_from_slice(&dq);
raw.extend_from_slice(&blk.ql);
raw.extend_from_slice(&blk.qh);
raw.extend_from_slice(&blk.scales.map(|s| s as u8));
raw.extend_from_slice(&blk.d.to_le_bytes());
}
}
let mut x_batch = vec![0.0f32; (n * x_stride) as usize];
for t in 0..n as usize {
for i in 0..k as usize {
x_batch[t * x_stride as usize + i] =
((t as f32 + 1.0) * (i as f32 - 200.0)) * 0.0007;
}
}
let mut expected = vec![0.0f32; (n * y_stride) as usize];
for t in 0..n as usize {
let x_slice = &x_batch[t * x_stride as usize..t * x_stride as usize + k as usize];
for row in 0..m as usize {
let mut acc = 0.0f32;
for i in 0..k as usize {
acc += w_f32[row * k as usize + i] * x_slice[i];
}
expected[t * y_stride as usize + row] = acc;
}
}
let a_buf = ctx.upload_storage(&raw, "gemm_q6k_weights");
let x_buf = ctx.upload_f32(&x_batch, "gemm_q6k_x");
let y_buf = ctx.create_storage_rw(((n * y_stride) as u64) * 4, "gemm_q6k_y");
let params: [u32; 5] = [m, k, n, x_stride, y_stride];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "gemm_q6k_params");
let pipeline = ctx.create_pipeline_with_defines(
shaders::MUL_MAT_REG_TILE,
"main",
"mul_mat_q6_k_test",
&[
("SCALAR", ""),
("SRC0_INNER_TYPE", "u32"),
("SRC1_INNER_TYPE", "f32"),
("INIT_SRC0_SHMEM_Q6_K", ""),
("INIT_SRC1_SHMEM_FLOAT", ""),
("WORKGROUP_SIZE_M", "8u"),
("WORKGROUP_SIZE_N", "32u"),
("TILE_M", "8u"),
("TILE_N", "4u"),
("TILE_K", "32u"),
],
);
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(m.div_ceil(64), n.div_ceil(128), 1);
}
ctx.queue.submit(Some(enc.finish()));
let got = ctx.download_f32(&y_buf, (n * y_stride) as usize);
for t in 0..n as usize {
for row in 0..m as usize {
let idx = t * y_stride as usize + row;
let denom = expected[idx].abs().max(1.0);
let rel = (expected[idx] - got[idx]).abs() / denom;
assert!(
rel < 5e-3,
"Q6_K GEMM mismatch at token {t}, row {row}: cpu={}, gpu={}, rel={rel:.2e}",
expected[idx],
got[idx]
);
}
}
}
#[test]
fn test_gpu_gemv_quant_dispatch_smoke() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let m = 9u32;
let k = 64u32;
let nb = k / 32;
let weights_f32: Vec<f32> = (0..m * k)
.map(|i| ((i * 11 + 7) % 37) as f32 * 0.08 - 1.2)
.collect();
let x: Vec<f32> = (0..k).map(|i| (i as f32 - 23.0) * 0.025).collect();
let mut q4_bytes = Vec::new();
let mut q4_expected = vec![0.0f32; m as usize];
let mut q8_bytes = Vec::new();
let mut q8_expected = vec![0.0f32; m as usize];
let mut f32_expected = vec![0.0f32; m as usize];
for row in 0..m as usize {
for col in 0..k as usize {
f32_expected[row] += weights_f32[row * k as usize + col] * x[col];
}
for b in 0..nb as usize {
let start = row * k as usize + b * 32;
let block = &weights_f32[start..start + 32];
let amax = block.iter().map(|v| v.abs()).fold(0.0f32, f32::max);
let d4 = if amax != 0.0 { amax / 7.0 } else { 0.0 };
let d4_f16 = half::f16::from_f32(d4);
q4_bytes.extend_from_slice(&d4_f16.to_bits().to_le_bytes());
let id4 = if d4 != 0.0 { 1.0 / d4 } else { 0.0 };
for qi in 0..16 {
let lo = ((block[qi] * id4 + 8.5) as u8).min(15);
let hi = ((block[16 + qi] * id4 + 8.5) as u8).min(15);
q4_bytes.push(lo | (hi << 4));
q4_expected[row] += (f32::from(lo) - 8.0) * d4_f16.to_f32() * x[b * 32 + qi];
q4_expected[row] +=
(f32::from(hi) - 8.0) * d4_f16.to_f32() * x[b * 32 + 16 + qi];
}
let d8 = if amax != 0.0 { amax / 127.0 } else { 0.0 };
let d8_f16 = half::f16::from_f32(d8);
q8_bytes.extend_from_slice(&d8_f16.to_bits().to_le_bytes());
let id8 = if d8 != 0.0 { 1.0 / d8 } else { 0.0 };
for (qi, &value) in block.iter().enumerate() {
let quant = (value * id8).round().clamp(-127.0, 127.0) as i8;
q8_bytes.push(quant as u8);
q8_expected[row] += f32::from(quant) * d8_f16.to_f32() * x[b * 32 + qi];
}
}
}
struct Case<'a> {
name: &'static str,
dtype: DType,
shader: &'static str,
entry: &'static str,
rows_per_wg: u32,
weight_bytes: &'a [u8],
expected: &'a [f32],
tolerance: f32,
}
let f32_weight_bytes = bytemuck::cast_slice(&weights_f32);
let cases = [
Case {
name: "f32",
dtype: DType::F32,
shader: shaders::GEMV_F32,
entry: "gemv_f32",
rows_per_wg: 8,
weight_bytes: f32_weight_bytes,
expected: &f32_expected,
tolerance: 1e-4,
},
Case {
name: "q4_0",
dtype: DType::Q4_0,
shader: shaders::GEMV_Q4_0_FAST,
entry: "gemv_q4_0_fast",
rows_per_wg: 4,
weight_bytes: &q4_bytes,
expected: &q4_expected,
tolerance: 5e-2,
},
Case {
name: "q8_0",
dtype: DType::Q8_0,
shader: shaders::GEMV_Q8_0,
entry: "gemv_q8_0",
rows_per_wg: 8,
weight_bytes: &q8_bytes,
expected: &q8_expected,
tolerance: 1e-3,
},
];
let x_buf = ctx.upload_f32(&x, "gemv_dispatch_x");
let params = [m, k, 0u32, 0u32];
let params_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "gemv_dispatch_params");
for case in cases {
let pipeline = ctx.create_pipeline(case.shader, case.entry, case.name);
let a_buf = ctx.upload_storage(case.weight_bytes, case.name);
let y_buf = ctx.create_storage_rw((m as u64) * 4, "gemv_dispatch_y");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some(case.name),
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: params_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(m.div_ceil(case.rows_per_wg), 1, 1);
}
ctx.submit_encoder(enc);
let result = ctx.download_f32(&y_buf, m as usize);
for (i, (&exp, &got)) in case.expected.iter().zip(&result).enumerate() {
let diff = (exp - got).abs();
assert!(
diff < case.tolerance,
"{:?} {} GEMV mismatch at row {i}: cpu={exp}, gpu={got}, diff={diff}",
case.dtype,
case.name,
);
}
}
}
#[test]
fn test_gpu_argmax_f32() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let pipeline = ctx.create_pipeline(shaders::ARGMAX_F32, "argmax_f32", "argmax_f32");
let cases: &[(usize, usize)] = &[
(32, 17), (256, 200), (2048, 1733), (50000, 12345), ];
for &(n, plant_idx) in cases {
let mut x: Vec<f32> = (0..n)
.map(|i| ((i as i32 * 31 + 7) % 211) as f32 / 211.0)
.collect();
x[plant_idx] = 99.0;
let x_buf = ctx.upload_f32(&x, "argmax_in");
let out_buf = ctx.create_storage_rw(4, "argmax_out");
let params =
ctx.upload_storage(bytemuck::cast_slice(&[n as u32, 0u32]), "argmax_params");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: out_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: params.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(1, 1, 1);
}
ctx.submit_encoder(enc);
let out = ctx.download_u32(&out_buf, 1);
assert_eq!(
out[0] as usize, plant_idx,
"argmax(n={n}) returned {}, expected {plant_idx}",
out[0]
);
}
let n: usize = 1024;
let mut x = vec![0.0f32; n];
x[100] = 5.0;
x[700] = 5.0; let x_buf = ctx.upload_f32(&x, "argmax_tie_in");
let out_buf = ctx.create_storage_rw(4, "argmax_tie_out");
let params =
ctx.upload_storage(bytemuck::cast_slice(&[n as u32, 0u32]), "argmax_tie_params");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: out_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: params.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(1, 1, 1);
}
ctx.submit_encoder(enc);
let out = ctx.download_u32(&out_buf, 1);
assert_eq!(
out[0], 100,
"tie-break: lower index must win, got {}",
out[0]
);
}
#[test]
fn spike_wgsl_override_constants() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let src = r#"
override HEAD_DIM: u32 = 1u;
@group(0) @binding(0) var<storage, read_write> out: array<u32>;
@compute @workgroup_size(1)
fn main() { out[0] = HEAD_DIM; }
"#;
let module = ctx
.device
.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some("override_spike"),
source: wgpu::ShaderSource::Wgsl(src.into()),
});
let make_pipeline = |head_dim: u32| {
let mut consts: std::collections::HashMap<String, f64> =
std::collections::HashMap::new();
consts.insert("HEAD_DIM".to_string(), head_dim as f64);
ctx.device
.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some(&format!("override_spike_hd{head_dim}")),
layout: None,
module: &module,
entry_point: Some("main"),
compilation_options: wgpu::PipelineCompilationOptions {
constants: &consts,
zero_initialize_workgroup_memory: true,
},
cache: None,
})
};
let dispatch_and_read = |pipeline: &wgpu::ComputePipeline| -> u32 {
let buf = ctx.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("override_spike_out"),
size: 4,
usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_SRC
| wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let bind_group = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[wgpu::BindGroupEntry {
binding: 0,
resource: buf.as_entire_binding(),
}],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.dispatch_workgroups(1, 1, 1);
}
ctx.submit_encoder(enc);
let staging = ctx.device.create_buffer(&wgpu::BufferDescriptor {
label: None,
size: 4,
usage: wgpu::BufferUsages::COPY_DST | wgpu::BufferUsages::MAP_READ,
mapped_at_creation: false,
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
enc.copy_buffer_to_buffer(&buf, 0, &staging, 0, 4);
ctx.submit_encoder(enc);
io_stats::record_readback(4);
let slice = staging.slice(..);
let (tx, rx) = std::sync::mpsc::channel();
slice.map_async(wgpu::MapMode::Read, move |r| {
tx.send(r).ok();
});
ctx.device.poll(wgpu::Maintain::Wait);
rx.recv().unwrap().unwrap();
let data = slice.get_mapped_range();
let v = u32::from_le_bytes([data[0], data[1], data[2], data[3]]);
drop(data);
staging.unmap();
v
};
let pipeline_64 = make_pipeline(64);
let pipeline_128 = make_pipeline(128);
let v64 = dispatch_and_read(&pipeline_64);
let v128 = dispatch_and_read(&pipeline_128);
assert_eq!(v64, 64, "override HEAD_DIM=64 not honored");
assert_eq!(v128, 128, "override HEAD_DIM=128 not honored");
println!("WGSL override spike OK: same module → HEAD_DIM={v64} and {v128}");
}
#[test]
fn test_gpu_rmsnorm_batch_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let n: u32 = 1024; let batch: u32 = 7; let eps = 1e-5f32;
let mut src: Vec<f32> = Vec::with_capacity((n * batch) as usize);
for b in 0..batch {
for i in 0..n {
src.push(((b as f32 + 1.0) * (i as f32 - 512.0)) * 0.001);
}
}
let weight: Vec<f32> = (0..n).map(|i| 0.8 + (i as f32 % 7.0) * 0.05).collect();
let pipeline_per = ctx.create_pipeline(shaders::RMSNORM, "rmsnorm", "rmsnorm_ref");
let w_buf = ctx.upload_f32(&weight, "w");
let params_per = [n, eps.to_bits(), 0u32, 0u32];
let p_buf_per = ctx.upload_storage(bytemuck::cast_slice(¶ms_per), "params_per");
let mut reference = vec![0.0f32; (n * batch) as usize];
for b in 0..batch {
let row_start = (b * n) as usize;
let row_end = row_start + n as usize;
let scratch = ctx.create_storage_rw((n as u64) * 4, "rmsnorm_ref_scratch");
ctx.queue
.write_buffer(&scratch, 0, bytemuck::cast_slice(&src[row_start..row_end]));
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline_per.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: scratch.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: w_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: p_buf_per.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline_per);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(1, 1, 1);
}
ctx.submit_encoder(enc);
let out = ctx.download_f32(&scratch, n as usize);
reference[row_start..row_end].copy_from_slice(&out);
}
let pipeline_batch =
ctx.create_pipeline(shaders::RMSNORM_BATCH, "rmsnorm_batch", "rmsnorm_batch");
let src_buf = ctx.create_storage_rw((src.len() as u64) * 4, "src");
ctx.queue
.write_buffer(&src_buf, 0, bytemuck::cast_slice(&src));
let dst_buf = ctx.create_storage_rw((src.len() as u64) * 4, "dst");
let params_batch = [n, eps.to_bits(), n, n, 1.0f32.to_bits()];
let p_buf_batch = ctx.upload_storage(bytemuck::cast_slice(¶ms_batch), "params_batch");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline_batch.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: src_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: dst_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: w_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p_buf_batch.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline_batch);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(batch, 1, 1);
}
ctx.submit_encoder(enc);
let batched = ctx.download_f32(&dst_buf, (n * batch) as usize);
for i in 0..(n * batch) as usize {
let diff = (reference[i] - batched[i]).abs();
assert!(
diff < 1e-3,
"rmsnorm_batch mismatch at idx {i} (token {}, dim {}): \
ref={}, batched={}, diff={diff}",
i / n as usize,
i % n as usize,
reference[i],
batched[i]
);
}
}
#[test]
fn test_gpu_qk_norm_rope_batch_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let n_heads: u32 = 4;
let n_kv_heads: u32 = 2;
let head_dim: u32 = 64;
let n_tokens: u32 = 3;
let start_pos: u32 = 5;
let eps = 1e-5f32;
let freq_base = 10000.0f32;
let rope_type: u32 = 0;
let q_stride = n_heads * head_dim;
let k_stride = n_kv_heads * head_dim;
let mut q_batch: Vec<f32> = Vec::with_capacity((n_tokens * q_stride) as usize);
let mut k_batch: Vec<f32> = Vec::with_capacity((n_tokens * k_stride) as usize);
for t in 0..n_tokens {
for i in 0..q_stride {
q_batch.push(((t as f32 + 1.0) * (i as f32 - 32.0)) * 0.01);
}
for i in 0..k_stride {
k_batch.push(((t as f32 + 2.0) * (i as f32 - 16.0)) * 0.013);
}
}
let q_norm_w: Vec<f32> = (0..head_dim)
.map(|i| 0.9 + (i as f32 % 5.0) * 0.04)
.collect();
let k_norm_w: Vec<f32> = (0..head_dim)
.map(|i| 1.1 - (i as f32 % 5.0) * 0.03)
.collect();
let mut ref_q = q_batch.clone();
let mut ref_k = k_batch.clone();
for t in 0..n_tokens {
let q_off = (t * q_stride) as usize;
let k_off = (t * k_stride) as usize;
for h in 0..n_heads as usize {
let head_start = q_off + h * head_dim as usize;
let head_end = head_start + head_dim as usize;
crate::backend::cpu::rmsnorm(&mut ref_q[head_start..head_end], &q_norm_w, eps);
}
for h in 0..n_kv_heads as usize {
let head_start = k_off + h * head_dim as usize;
let head_end = head_start + head_dim as usize;
crate::backend::cpu::rmsnorm(&mut ref_k[head_start..head_end], &k_norm_w, eps);
}
let q_end = q_off + (n_heads * head_dim) as usize;
let k_end = k_off + (n_kv_heads * head_dim) as usize;
crate::backend::cpu::rope(
&mut ref_q[q_off..q_end],
&mut ref_k[k_off..k_end],
(start_pos + t) as usize,
n_heads as usize,
n_kv_heads as usize,
head_dim as usize,
freq_base,
);
}
let pipeline = ctx.create_pipeline(
shaders::QK_NORM_ROPE_BATCH,
"qk_norm_rope_batch",
"qk_norm_rope_batch",
);
let q_buf = ctx.create_storage_rw((q_batch.len() as u64) * 4, "q");
ctx.queue
.write_buffer(&q_buf, 0, bytemuck::cast_slice(&q_batch));
let k_buf = ctx.create_storage_rw((k_batch.len() as u64) * 4, "k");
ctx.queue
.write_buffer(&k_buf, 0, bytemuck::cast_slice(&k_batch));
let qw_buf = ctx.upload_f32(&q_norm_w, "q_norm_w");
let kw_buf = ctx.upload_f32(&k_norm_w, "k_norm_w");
let params = [
start_pos,
n_tokens,
n_heads,
n_kv_heads,
head_dim,
eps.to_bits(),
freq_base.to_bits(),
rope_type,
q_stride,
k_stride,
0, 1, ];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let ff_buf = ctx.upload_f32(&[1.0f32], "freq_factors_dummy");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: q_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: k_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: qw_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: kw_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: p_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 5,
resource: ff_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(n_tokens * (n_heads + n_kv_heads), 1, 1);
}
ctx.submit_encoder(enc);
let got_q = ctx.download_f32(&q_buf, q_batch.len());
let got_k = ctx.download_f32(&k_buf, k_batch.len());
let tol = 2e-3f32;
for i in 0..ref_q.len() {
let diff = (ref_q[i] - got_q[i]).abs();
assert!(
diff < tol,
"Q mismatch at idx {i} (token {}, dim {}): cpu={}, gpu={}, diff={diff}",
i / q_stride as usize,
i % q_stride as usize,
ref_q[i],
got_q[i]
);
}
for i in 0..ref_k.len() {
let diff = (ref_k[i] - got_k[i]).abs();
assert!(
diff < tol,
"K mismatch at idx {i} (token {}, dim {}): cpu={}, gpu={}, diff={diff}",
i / k_stride as usize,
i % k_stride as usize,
ref_k[i],
got_k[i]
);
}
}
#[test]
fn test_gpu_qk_norm_rope_batch_rope_only_and_freq_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let n_heads: u32 = 4;
let n_kv_heads: u32 = 2;
let head_dim: u32 = 64;
let n_tokens: u32 = 3;
let start_pos: u32 = 5;
let eps = 1e-5f32; let freq_base = 10000.0f32;
let q_stride = n_heads * head_dim;
let k_stride = n_kv_heads * head_dim;
let freqs: Vec<f32> = (0..head_dim / 2).map(|i| 1.0 + (i as f32) * 0.05).collect();
let cases: [(&str, u32, bool); 2] = [("neox_rope_only", 0, false), ("norm_freq", 1, true)];
for (label, rope_type, has_freq_factors) in cases {
let mut q_batch: Vec<f32> = Vec::with_capacity((n_tokens * q_stride) as usize);
let mut k_batch: Vec<f32> = Vec::with_capacity((n_tokens * k_stride) as usize);
for t in 0..n_tokens {
for i in 0..q_stride {
q_batch.push(((t as f32 + 1.0) * (i as f32 - 32.0)) * 0.01);
}
for i in 0..k_stride {
k_batch.push(((t as f32 + 2.0) * (i as f32 - 16.0)) * 0.013);
}
}
let ff = if has_freq_factors {
Some(freqs.as_slice())
} else {
None
};
let mut ref_q = q_batch.clone();
let mut ref_k = k_batch.clone();
for t in 0..n_tokens {
let q_off = (t * q_stride) as usize;
let k_off = (t * k_stride) as usize;
let q_end = q_off + q_stride as usize;
let k_end = k_off + k_stride as usize;
let pos = (start_pos + t) as usize;
if rope_type == 0 {
crate::backend::cpu::rope(
&mut ref_q[q_off..q_end],
&mut ref_k[k_off..k_end],
pos,
n_heads as usize,
n_kv_heads as usize,
head_dim as usize,
freq_base,
);
} else {
crate::backend::cpu::rope_norm(
&mut ref_q[q_off..q_end],
&mut ref_k[k_off..k_end],
pos,
n_heads as usize,
n_kv_heads as usize,
head_dim as usize,
freq_base,
ff,
);
}
}
let pipeline = ctx.create_pipeline(
shaders::QK_NORM_ROPE_BATCH,
"qk_norm_rope_batch",
"qk_norm_rope_batch",
);
let q_buf = ctx.create_storage_rw((q_batch.len() as u64) * 4, "q");
ctx.queue
.write_buffer(&q_buf, 0, bytemuck::cast_slice(&q_batch));
let k_buf = ctx.create_storage_rw((k_batch.len() as u64) * 4, "k");
ctx.queue
.write_buffer(&k_buf, 0, bytemuck::cast_slice(&k_batch));
let dummy = ctx.upload_f32(&[1.0f32], "qk_norm_dummy");
let ff_buf = if has_freq_factors {
ctx.upload_f32(&freqs, "freq_factors")
} else {
ctx.upload_f32(&[1.0f32], "freq_factors_dummy")
};
let params = [
start_pos,
n_tokens,
n_heads,
n_kv_heads,
head_dim,
eps.to_bits(),
freq_base.to_bits(),
rope_type,
q_stride,
k_stride,
has_freq_factors as u32,
0, ];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: q_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: k_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: dummy.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: dummy.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: p_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 5,
resource: ff_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(n_tokens * (n_heads + n_kv_heads), 1, 1);
}
ctx.submit_encoder(enc);
let got_q = ctx.download_f32(&q_buf, q_batch.len());
let got_k = ctx.download_f32(&k_buf, k_batch.len());
let tol = 2e-3f32;
for i in 0..ref_q.len() {
let diff = (ref_q[i] - got_q[i]).abs();
assert!(
diff < tol,
"[{label}] Q mismatch at idx {i}: cpu={}, gpu={}, diff={diff}",
ref_q[i],
got_q[i]
);
}
for i in 0..ref_k.len() {
let diff = (ref_k[i] - got_k[i]).abs();
assert!(
diff < tol,
"[{label}] K mismatch at idx {i}: cpu={}, gpu={}, diff={diff}",
ref_k[i],
got_k[i]
);
}
}
}
#[test]
fn test_gpu_conv1d_fused_batch_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let hs: usize = 64;
let kernel_size: usize = 4;
let d_conv: usize = kernel_size - 1; let n_tokens: usize = 5;
let proj_stride = 3 * hs;
let out_stride = hs;
let mut proj: Vec<f32> = Vec::with_capacity(n_tokens * proj_stride);
for t in 0..n_tokens {
for i in 0..hs {
proj.push(((t as f32 + 1.0) * (i as f32 - 32.0)) * 0.011);
}
for i in 0..hs {
proj.push(((t as f32 + 2.0) * (i as f32 + 5.0)) * 0.007);
}
for i in 0..hs {
proj.push(((t as f32 + 3.0) * (i as f32 - 16.0)) * 0.013);
}
}
let mut rb_initial: Vec<f32> = Vec::with_capacity(d_conv * hs);
for k in 0..d_conv {
for i in 0..hs {
rb_initial.push(((k as f32 + 1.0) * (i as f32 - 8.0)) * 0.005);
}
}
let mut weight: Vec<f32> = Vec::with_capacity(hs * kernel_size);
for ch in 0..hs {
for k in 0..kernel_size {
weight.push(0.1 + (ch as f32 % 7.0) * 0.02 - (k as f32) * 0.03);
}
}
let mut ref_out = vec![0.0f32; n_tokens * out_stride];
let mut ref_rb = rb_initial.clone();
for t in 0..n_tokens {
let base = t * proj_stride;
for ch in 0..hs {
let x = proj[base + ch];
let c = proj[base + hs + ch];
let b = proj[base + 2 * hs + ch];
let bx = x * b;
let mut sum = 0.0f32;
for k in 0..d_conv {
sum += ref_rb[k * hs + ch] * weight[ch * kernel_size + k];
}
sum += bx * weight[ch * kernel_size + d_conv];
if d_conv > 1 {
for k in 0..d_conv - 1 {
ref_rb[k * hs + ch] = ref_rb[(k + 1) * hs + ch];
}
}
if d_conv > 0 {
ref_rb[(d_conv - 1) * hs + ch] = bx;
}
ref_out[t * out_stride + ch] = c * sum;
}
}
let pipeline = ctx.create_pipeline(
shaders::CONV1D_FUSED_BATCH,
"conv1d_fused_batch",
"conv1d_fused_batch",
);
let proj_buf = ctx.upload_f32(&proj, "proj");
let rb_buf = ctx.create_storage_rw((rb_initial.len() as u64) * 4, "rb");
ctx.queue
.write_buffer(&rb_buf, 0, bytemuck::cast_slice(&rb_initial));
let weight_buf = ctx.upload_f32(&weight, "weight");
let out_buf = ctx.create_storage_rw((ref_out.len() as u64) * 4, "out");
let params: [u32; 6] = [
hs as u32,
kernel_size as u32,
d_conv as u32,
n_tokens as u32,
proj_stride as u32,
out_stride as u32,
];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: proj_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: rb_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: weight_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: out_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: p_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(hs.div_ceil(256) as u32, 1, 1);
}
ctx.submit_encoder(enc);
let got_out = ctx.download_f32(&out_buf, ref_out.len());
let got_rb = ctx.download_f32(&rb_buf, ref_rb.len());
let tol = 1e-4f32;
for i in 0..ref_out.len() {
let diff = (ref_out[i] - got_out[i]).abs();
assert!(
diff < tol,
"out mismatch at idx {i} (token {}, ch {}): cpu={}, gpu={}, diff={diff}",
i / hs,
i % hs,
ref_out[i],
got_out[i]
);
}
for i in 0..ref_rb.len() {
let diff = (ref_rb[i] - got_rb[i]).abs();
assert!(
diff < tol,
"rolling-buffer mismatch at idx {i}: cpu={}, gpu={}, diff={diff}",
ref_rb[i],
got_rb[i]
);
}
}
#[test]
fn test_gpu_conv1d_fused_decode_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let hs: usize = 1024;
let kernel_size: usize = 4;
let d_conv: usize = kernel_size - 1;
let mut proj: Vec<f32> = Vec::with_capacity(3 * hs);
for i in 0..hs {
proj.push((i as f32 - 32.0) * 0.011);
}
for i in 0..hs {
proj.push((i as f32 + 5.0) * 0.007);
}
for i in 0..hs {
proj.push((i as f32 - 16.0) * 0.013);
}
let mut rb_initial: Vec<f32> = Vec::with_capacity(d_conv * hs);
for k in 0..d_conv {
for i in 0..hs {
rb_initial.push(((k as f32 + 1.0) * (i as f32 - 8.0)) * 0.005);
}
}
let mut weight: Vec<f32> = Vec::with_capacity(hs * kernel_size);
for ch in 0..hs {
for k in 0..kernel_size {
weight.push(0.1 + (ch as f32 % 7.0) * 0.02 - (k as f32) * 0.03);
}
}
let mut ref_out = vec![0.0f32; hs];
let mut ref_rb = rb_initial.clone();
for ch in 0..hs {
let x = proj[ch];
let c = proj[hs + ch];
let b = proj[2 * hs + ch];
let bx = x * b;
let mut sum = 0.0f32;
for k in 0..d_conv {
sum += ref_rb[k * hs + ch] * weight[ch * kernel_size + k];
}
sum += bx * weight[ch * kernel_size + d_conv];
if d_conv > 1 {
for k in 0..d_conv - 1 {
ref_rb[k * hs + ch] = ref_rb[(k + 1) * hs + ch];
}
}
if d_conv > 0 {
ref_rb[(d_conv - 1) * hs + ch] = bx;
}
ref_out[ch] = c * sum;
}
let pipeline = ctx.create_pipeline(shaders::CONV1D_FUSED, "conv1d_fused", "conv1d_fused");
let proj_buf = ctx.upload_f32(&proj, "proj");
let rb_buf = ctx.create_storage_rw((rb_initial.len() as u64) * 4, "rb");
ctx.queue
.write_buffer(&rb_buf, 0, bytemuck::cast_slice(&rb_initial));
let weight_buf = ctx.upload_f32(&weight, "weight");
let out_buf = ctx.create_storage_rw((ref_out.len() as u64) * 4, "out");
let params: [u32; 4] = [hs as u32, kernel_size as u32, d_conv as u32, 0];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: proj_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: rb_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: weight_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: out_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: p_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(hs.div_ceil(256) as u32, 1, 1);
}
ctx.submit_encoder(enc);
let got_out = ctx.download_f32(&out_buf, ref_out.len());
let got_rb = ctx.download_f32(&rb_buf, ref_rb.len());
let tol = 1e-4f32;
for i in 0..ref_out.len() {
let diff = (ref_out[i] - got_out[i]).abs();
assert!(
diff < tol,
"out mismatch at ch {i}: cpu={}, gpu={}, diff={diff}",
ref_out[i],
got_out[i]
);
}
for i in 0..ref_rb.len() {
let diff = (ref_rb[i] - got_rb[i]).abs();
assert!(
diff < tol,
"rolling-buffer mismatch at idx {i}: cpu={}, gpu={}, diff={diff}",
ref_rb[i],
got_rb[i]
);
}
}
#[test]
fn test_gpu_add_rmsnorm_batch_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let n: u32 = 1024;
let batch: u32 = 5;
let eps = 1e-5f32;
let mut src: Vec<f32> = Vec::with_capacity((n * batch) as usize);
let mut residual: Vec<f32> = Vec::with_capacity((n * batch) as usize);
for b in 0..batch {
for i in 0..n {
src.push(((b + 1) as f32 * (i as f32 - 512.0)) * 0.001);
residual.push(((b + 2) as f32 * ((i as f32 + 17.0) % 13.0)) * 0.002);
}
}
let weight: Vec<f32> = (0..n).map(|i| 0.8 + (i as f32 % 7.0) * 0.05).collect();
let res_scale = 0.7f32;
let mut reference = src.clone();
for i in 0..reference.len() {
reference[i] += res_scale * residual[i];
}
for b in 0..batch {
let row_start = (b * n) as usize;
let row_end = row_start + n as usize;
crate::backend::cpu::rmsnorm(&mut reference[row_start..row_end], &weight, eps);
}
let pipeline = ctx.create_pipeline(
shaders::RMSNORM_BATCH,
"add_rmsnorm_batch",
"add_rmsnorm_batch",
);
let src_buf = ctx.create_storage_rw((src.len() as u64) * 4, "src");
ctx.queue
.write_buffer(&src_buf, 0, bytemuck::cast_slice(&src));
let dst_buf = ctx.create_storage_rw((src.len() as u64) * 4, "dst");
let res_buf = ctx.upload_f32(&residual, "residual");
let w_buf = ctx.upload_f32(&weight, "w");
let params = [n, eps.to_bits(), n, n, res_scale.to_bits()];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: src_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: dst_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: w_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: res_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(batch, 1, 1);
}
ctx.submit_encoder(enc);
let batched = ctx.download_f32(&dst_buf, (n * batch) as usize);
for i in 0..(n * batch) as usize {
let diff = (reference[i] - batched[i]).abs();
assert!(
diff < 1e-3,
"add_rmsnorm_batch mismatch at idx {i} (token {}, dim {}): \
ref={}, batched={}, diff={diff}",
i / n as usize,
i % n as usize,
reference[i],
batched[i]
);
}
}
struct AttnPrefillFixture {
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
kv_dim: u32,
n_queries: u32,
start_pos: u32,
max_seq: u32,
scale: f32,
q_stride: u32,
out_stride: u32,
q_batch: Vec<f32>,
k_cache: Vec<f32>,
v_cache: Vec<f32>,
ref_out: Vec<f32>,
}
fn attn_prefill_fixture() -> AttnPrefillFixture {
build_attn_prefill_fixture(4, 2, 32, 5, 3)
}
fn build_attn_prefill_fixture(
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
n_queries: u32,
start_pos: u32,
) -> AttnPrefillFixture {
let kv_dim = n_kv_heads * head_dim;
let max_seq = start_pos + n_queries;
let scale = 1.0f32 / (head_dim as f32).sqrt();
let q_stride = n_heads * head_dim;
let out_stride = n_heads * head_dim;
let group_size = n_heads / n_kv_heads;
let mut q_batch = vec![0.0f32; (n_queries * q_stride) as usize];
for q in 0..n_queries {
for h in 0..n_heads {
for d in 0..head_dim {
let val = if d == 0 {
1.0
} else {
(((q + 1) * (h + 1) + d) as f32 * 0.05).sin()
};
q_batch[(q * q_stride + h * head_dim + d) as usize] = val;
}
}
}
let mut k_cache = vec![0.0f32; (max_seq * kv_dim) as usize];
let mut v_cache = vec![0.0f32; (max_seq * kv_dim) as usize];
for t in 0..max_seq {
for kh in 0..n_kv_heads {
for d in 0..head_dim {
let kv = if d == 0 {
t as f32 * 0.02 } else {
(((t + 1) * (kh + 1) + d) as f32 * 0.03).cos()
};
let va = ((t + 1) * (kh + 2) + d) as f32 * 0.02;
k_cache[(t * kv_dim + kh * head_dim + d) as usize] = kv;
v_cache[(t * kv_dim + kh * head_dim + d) as usize] = va.sin();
}
}
}
let mut ref_out = vec![0.0f32; (n_queries * out_stride) as usize];
for q in 0..n_queries as usize {
let seq_len = start_pos as usize + q + 1;
for h in 0..n_heads as usize {
let kv_h_off = (h / group_size as usize) * head_dim as usize;
let q_off = q * q_stride as usize + h * head_dim as usize;
let mut scores = vec![0.0f32; seq_len];
for (t, s) in scores.iter_mut().enumerate() {
let mut dot = 0.0f32;
for d in 0..head_dim as usize {
dot += q_batch[q_off + d] * k_cache[t * kv_dim as usize + kv_h_off + d];
}
*s = dot * scale;
}
let max_s = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
for s in scores.iter_mut() {
*s = (*s - max_s).exp();
sum += *s;
}
let inv = 1.0f32 / sum;
let out_off = q * out_stride as usize + h * head_dim as usize;
for d in 0..head_dim as usize {
let mut val = 0.0f32;
for (t, s) in scores.iter().enumerate() {
val += (*s * inv) * v_cache[t * kv_dim as usize + kv_h_off + d];
}
ref_out[out_off + d] = val;
}
}
}
AttnPrefillFixture {
n_heads,
n_kv_heads,
head_dim,
kv_dim,
n_queries,
start_pos,
max_seq,
scale,
q_stride,
out_stride,
q_batch,
k_cache,
v_cache,
ref_out,
}
}
fn run_gpu_attention_prefill_tiled(
ctx: &GpuContext,
f: &AttnPrefillFixture,
tile: u32,
) -> Vec<f32> {
assert!(tile > 0, "tile must be > 0 (0 would never advance q_base)");
let pipeline = ctx.create_pipeline(
shaders::ATTENTION_PREFILL,
"attention_prefill",
"attention_prefill",
);
let q_buf = ctx.upload_f32(&f.q_batch, "q");
let k_buf = ctx.upload_f32(&f.k_cache, "k");
let v_buf = ctx.upload_f32(&f.v_cache, "v");
let out_buf = ctx.create_storage_rw((f.ref_out.len() as u64) * 4, "out");
let mut q_base = 0u32;
while q_base < f.n_queries {
let n_sub = (f.n_queries - q_base).min(tile);
let params: [u32; 12] = [
f.n_heads,
f.n_kv_heads,
f.head_dim,
f.kv_dim,
f.max_seq,
f.scale.to_bits(),
f.start_pos,
n_sub,
f.q_stride,
f.out_stride,
q_base,
0,
];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: q_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: k_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: v_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: out_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: p_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(f.n_heads, n_sub, 1);
}
ctx.submit_encoder(enc);
q_base += n_sub;
}
ctx.download_f32(&out_buf, f.ref_out.len())
}
fn run_gpu_attention_prefill(ctx: &GpuContext, f: &AttnPrefillFixture) -> Vec<f32> {
run_gpu_attention_prefill_tiled(ctx, f, f.n_queries)
}
fn assert_attn_prefill_matches(f: &AttnPrefillFixture, got: &[f32], label: &str) {
for (i, &g) in got.iter().enumerate() {
let r = f.ref_out[i];
let tol = 1e-3f32 + 1e-3f32 * r.abs();
let diff = (r - g).abs();
assert!(
diff <= tol,
"{label} mismatch at idx {i} (token {}, head {}, dim {}): \
cpu={r}, gpu={g}, diff={diff}",
i / f.out_stride as usize,
(i % f.out_stride as usize) / f.head_dim as usize,
i % f.head_dim as usize,
);
}
}
#[test]
fn test_gpu_attention_prefill_parity() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let f = attn_prefill_fixture();
let got = run_gpu_attention_prefill(&ctx, &f);
assert_attn_prefill_matches(&f, &got, "attention_prefill");
}
#[test]
fn test_gpu_attention_prefill_long_kv() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let f = build_attn_prefill_fixture(4, 2, 32, 200, 900);
assert!(
f.max_seq > 4 * 256,
"long-KV fixture must span >4 key tiles; got max_seq={}",
f.max_seq
);
let got = run_gpu_attention_prefill(&ctx, &f);
assert_attn_prefill_matches(&f, &got, "long-kv attention_prefill");
}
#[test]
fn test_gpu_attention_prefill_seq1() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let f = build_attn_prefill_fixture(2, 1, 16, 1, 0);
assert_eq!(f.max_seq, 1);
let got = run_gpu_attention_prefill(&ctx, &f);
assert_attn_prefill_matches(&f, &got, "seq1 attention_prefill");
}
#[test]
fn test_gpu_attention_prefill_qbase_offset() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let f = attn_prefill_fixture(); let got = run_gpu_attention_prefill_tiled(&ctx, &f, 2);
assert_attn_prefill_matches(&f, &got, "q_base-offset attention_prefill");
}
#[test]
fn test_gpu_attention_prefill_seq_len_zero() {
let ctx = match GpuContext::new() {
Ok(ctx) => ctx,
Err(_) => return,
};
let n_heads = 2u32;
let n_kv_heads = 1u32;
let head_dim = 16u32;
let kv_dim = n_kv_heads * head_dim;
let n = 1u32;
let q_stride = n_heads * head_dim;
let out_stride = q_stride;
let out_len = (n * out_stride) as usize;
let pipeline = ctx.create_pipeline(
shaders::ATTENTION_PREFILL,
"attention_prefill",
"attention_prefill",
);
let q_buf = ctx.upload_f32(&vec![0.5f32; (n * q_stride) as usize], "q");
let k_buf = ctx.upload_f32(&vec![0.1f32; kv_dim as usize], "k");
let v_buf = ctx.upload_f32(&vec![0.2f32; kv_dim as usize], "v");
let out_buf = ctx.upload_f32(&vec![7.0f32; out_len], "out");
let params: [u32; 12] = [
n_heads,
n_kv_heads,
head_dim,
kv_dim,
0, 1.0f32.to_bits(),
0, n,
q_stride,
out_stride,
0, 0,
];
let p_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: q_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: k_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: v_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: out_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: p_buf.as_entire_binding(),
},
],
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(n_heads, n, 1);
}
ctx.submit_encoder(enc);
let got = ctx.download_f32(&out_buf, out_len);
for (i, &g) in got.iter().enumerate() {
assert_eq!(g, 0.0, "seq_len==0 must zero output at idx {i}, got {g}");
}
}
}