use core::marker::PhantomData;
use std::sync::OnceLock;
use cubecl::prelude::*;
use cubecl::server::Handle;
use crate::{ModelError, Result};
pub const Q4_0_BLOCK: usize = 32;
pub const Q4_0_BLOCK_BYTES: usize = 18;
pub fn repack_q4_0(data: &[u8]) -> Result<(Vec<u32>, Vec<f32>)> {
if data.is_empty() || data.len() % Q4_0_BLOCK_BYTES != 0 {
return Err(ModelError::BadShape {
tensor: "q4_0 block stream".into(),
expected: vec![Q4_0_BLOCK_BYTES],
got: vec![data.len()],
});
}
let n_blocks = data.len() / Q4_0_BLOCK_BYTES;
let mut qs = Vec::with_capacity(n_blocks * 4);
let mut d = Vec::with_capacity(n_blocks);
for block in data.chunks_exact(Q4_0_BLOCK_BYTES) {
d.push(burn::tensor::f16::from_le_bytes([block[0], block[1]]).to_f32());
for w in 0..4 {
let o = 2 + 4 * w;
qs.push(u32::from_le_bytes([
block[o],
block[o + 1],
block[o + 2],
block[o + 3],
]));
}
}
Ok((qs, d))
}
#[cube(launch_unchecked)]
fn q4_0_dequant_kernel(qs: &Array<u32>, d: &Array<f32>, out: &mut Array<f32>, n: usize) {
if ABSOLUTE_POS < n {
let block = ABSOLUTE_POS / 32;
let j = ABSOLUTE_POS % 32;
let byte_idx = j % 16;
let word = qs[block * 4 + byte_idx / 4];
let byte = (word >> (u32::cast_from(byte_idx % 4) * 8)) & 0xFF;
let mut nib = byte & 0xF;
if j >= 16 {
nib = byte >> 4;
}
out[ABSOLUTE_POS] = f32::cast_from(i32::cast_from(nib) - 8) * d[block];
}
}
#[cube(launch_unchecked)]
fn q4_0_matmul_kernel(
x: &Array<f32>,
qs: &Array<u32>,
d: &Array<f32>,
out: &mut Array<f32>,
m: usize,
k: usize,
n_out: usize,
) {
if ABSOLUTE_POS < m * n_out {
let row = ABSOLUTE_POS / n_out;
let col = ABSOLUTE_POS % n_out;
let blocks_per_row = k / 32;
let mut acc = 0.0f32;
for kb in 0..blocks_per_row {
let block = col * blocks_per_row + kb;
let x_base = row * k + kb * 32;
let mut block_acc = 0.0f32;
for w in 0..4usize {
let word = qs[block * 4 + w];
for b in 0..4usize {
let byte = (word >> (u32::cast_from(b) * 8)) & 0xFF;
let jj = w * 4 + b;
let lo = f32::cast_from(i32::cast_from(byte & 0xF) - 8);
let hi = f32::cast_from(i32::cast_from(byte >> 4) - 8);
block_acc += lo * x[x_base + jj];
block_acc += hi * x[x_base + 16 + jj];
}
}
acc += d[block] * block_acc;
}
out[row * n_out + col] = acc;
}
}
const CUBE_DIM: u32 = 256;
const MAX_CUBES_PER_DIM: u32 = 65535;
fn cube_count_1d(total: u32) -> CubeCount {
CubeCount::Static(total.div_ceil(CUBE_DIM).max(1), 1, 1)
}
fn cube_count_capped(total: u32) -> CubeCount {
let cubes = total.div_ceil(CUBE_DIM).max(1);
if cubes <= MAX_CUBES_PER_DIM {
CubeCount::Static(cubes, 1, 1)
} else {
let y = cubes.div_ceil(MAX_CUBES_PER_DIM);
CubeCount::Static(MAX_CUBES_PER_DIM, y, 1)
}
}
fn cube_count_tiled(n_out: u32, m: u32) -> CubeCount {
CubeCount::Static(n_out.div_ceil(CUBE_DIM).max(1), m.max(1), 1)
}
fn tiled_enabled() -> bool {
static DISABLED: OnceLock<bool> = OnceLock::new();
!*DISABLED.get_or_init(|| {
matches!(std::env::var("COMBS_NO_TILED_MATMUL").as_deref(), Ok("1"))
})
}
pub fn dequantize_q4_0_gpu<R: Runtime>(client: &ComputeClient<R>, data: &[u8]) -> Result<Vec<f32>> {
let (qs, d) = repack_q4_0(data)?;
let n = d.len() * Q4_0_BLOCK;
let qs_h = client.create_from_slice(u32::as_bytes(&qs));
let d_h = client.create_from_slice(f32::as_bytes(&d));
let out_h = client.empty(n * core::mem::size_of::<f32>());
unsafe {
q4_0_dequant_kernel::launch_unchecked::<R>(
client,
cube_count_1d(n as u32),
CubeDim::new_1d(CUBE_DIM),
ArrayArg::from_raw_parts(qs_h, qs.len()),
ArrayArg::from_raw_parts(d_h, d.len()),
ArrayArg::from_raw_parts(out_h.clone(), n),
n,
);
}
let bytes = client.read_one_unchecked(out_h);
Ok(f32::from_bytes(&bytes).to_vec())
}
pub struct Q40Weight<R: Runtime> {
qs: Handle,
d: Handle,
n_out: usize,
k: usize,
_runtime: PhantomData<R>,
}
impl<R: Runtime> Q40Weight<R> {
pub fn from_gguf_bytes(
client: &ComputeClient<R>,
data: &[u8],
n_out: usize,
k: usize,
) -> Result<Self> {
if k == 0 || k % Q4_0_BLOCK != 0 || data.len() != n_out * k / Q4_0_BLOCK * Q4_0_BLOCK_BYTES
{
return Err(ModelError::BadShape {
tensor: "q4_0 weight".into(),
expected: vec![n_out, k / Q4_0_BLOCK.max(1) * Q4_0_BLOCK_BYTES],
got: vec![data.len()],
});
}
let (qs, d) = repack_q4_0(data)?;
Ok(Q40Weight {
qs: client.create_from_slice(u32::as_bytes(&qs)),
d: client.create_from_slice(f32::as_bytes(&d)),
n_out,
k,
_runtime: PhantomData,
})
}
pub fn n_out(&self) -> usize {
self.n_out
}
pub fn k(&self) -> usize {
self.k
}
pub fn vram_bytes(&self) -> usize {
let n_blocks = self.n_out * self.k / Q4_0_BLOCK;
n_blocks * (16 + core::mem::size_of::<f32>())
}
pub fn matmul_device(&self, client: &ComputeClient<R>, x: Handle, m: usize) -> Handle {
let out_len = m * self.n_out;
let out_h = client.empty(out_len * core::mem::size_of::<f32>());
let n_blocks = self.n_out * self.k / Q4_0_BLOCK;
unsafe {
q4_0_matmul_kernel::launch_unchecked::<R>(
client,
cube_count_capped(out_len as u32),
CubeDim::new_1d(CUBE_DIM),
ArrayArg::from_raw_parts(x, m * self.k),
ArrayArg::from_raw_parts(self.qs.clone(), n_blocks * 4),
ArrayArg::from_raw_parts(self.d.clone(), n_blocks),
ArrayArg::from_raw_parts(out_h.clone(), out_len),
m,
self.k,
self.n_out,
);
}
out_h
}
pub fn matmul_host(&self, client: &ComputeClient<R>, x: &[f32], m: usize) -> Result<Vec<f32>> {
if m == 0 || x.len() != m * self.k {
return Err(ModelError::BadShape {
tensor: "q4_0 matmul input".into(),
expected: vec![m, self.k],
got: vec![x.len()],
});
}
let x_h = client.create_from_slice(f32::as_bytes(x));
let out_h = self.matmul_device(client, x_h, m);
let bytes = client.read_one_unchecked(out_h);
Ok(f32::from_bytes(&bytes).to_vec())
}
}
pub const Q5_0_BLOCK_BYTES: usize = 22;
pub const Q8_0_BLOCK_BYTES: usize = 34;
pub fn repack_q5_0(data: &[u8]) -> Result<(Vec<u32>, Vec<u32>, Vec<f32>)> {
if data.is_empty() || data.len() % Q5_0_BLOCK_BYTES != 0 {
return Err(ModelError::BadShape {
tensor: "q5_0 block stream".into(),
expected: vec![Q5_0_BLOCK_BYTES],
got: vec![data.len()],
});
}
let n_blocks = data.len() / Q5_0_BLOCK_BYTES;
let mut qs = Vec::with_capacity(n_blocks * 4);
let mut qh = Vec::with_capacity(n_blocks);
let mut d = Vec::with_capacity(n_blocks);
for block in data.chunks_exact(Q5_0_BLOCK_BYTES) {
d.push(burn::tensor::f16::from_le_bytes([block[0], block[1]]).to_f32());
qh.push(u32::from_le_bytes([block[2], block[3], block[4], block[5]]));
for w in 0..4 {
let o = 6 + 4 * w;
qs.push(u32::from_le_bytes([
block[o],
block[o + 1],
block[o + 2],
block[o + 3],
]));
}
}
Ok((qs, qh, d))
}
pub fn repack_q8_0(data: &[u8]) -> Result<(Vec<u32>, Vec<f32>)> {
if data.is_empty() || data.len() % Q8_0_BLOCK_BYTES != 0 {
return Err(ModelError::BadShape {
tensor: "q8_0 block stream".into(),
expected: vec![Q8_0_BLOCK_BYTES],
got: vec![data.len()],
});
}
let n_blocks = data.len() / Q8_0_BLOCK_BYTES;
let mut qs = Vec::with_capacity(n_blocks * 8);
let mut d = Vec::with_capacity(n_blocks);
for block in data.chunks_exact(Q8_0_BLOCK_BYTES) {
d.push(burn::tensor::f16::from_le_bytes([block[0], block[1]]).to_f32());
for w in 0..8 {
let o = 2 + 4 * w;
qs.push(u32::from_le_bytes([
block[o],
block[o + 1],
block[o + 2],
block[o + 3],
]));
}
}
Ok((qs, d))
}
#[cube(launch_unchecked)]
fn q5_0_dequant_kernel(
qs: &Array<u32>,
qh: &Array<u32>,
d: &Array<f32>,
out: &mut Array<f32>,
n: usize,
) {
if ABSOLUTE_POS < n {
let block = ABSOLUTE_POS / 32;
let j = ABSOLUTE_POS % 32;
let byte_idx = j % 16;
let word = qs[block * 4 + byte_idx / 4];
let byte = (word >> (u32::cast_from(byte_idx % 4) * 8)) & 0xFF;
let mut nib = byte & 0xF;
if j >= 16 {
nib = byte >> 4;
}
let hi_bit = (qh[block] >> u32::cast_from(j)) & 1;
let q = i32::cast_from(nib | (hi_bit << 4)) - 16;
out[ABSOLUTE_POS] = f32::cast_from(q) * d[block];
}
}
#[cube(launch_unchecked)]
fn q5_0_matmul_kernel(
x: &Array<f32>,
qs: &Array<u32>,
qh: &Array<u32>,
d: &Array<f32>,
out: &mut Array<f32>,
m: usize,
k: usize,
n_out: usize,
) {
if ABSOLUTE_POS < m * n_out {
let row = ABSOLUTE_POS / n_out;
let col = ABSOLUTE_POS % n_out;
let blocks_per_row = k / 32;
let mut acc = 0.0f32;
for kb in 0..blocks_per_row {
let block = col * blocks_per_row + kb;
let x_base = row * k + kb * 32;
let bits = qh[block];
let mut block_acc = 0.0f32;
for w in 0..4usize {
let word = qs[block * 4 + w];
for b in 0..4usize {
let byte = (word >> (u32::cast_from(b) * 8)) & 0xFF;
let jj = w * 4 + b;
let lo_bit = (bits >> u32::cast_from(jj)) & 1;
let hi_bit = (bits >> u32::cast_from(jj + 16)) & 1;
let lo = f32::cast_from(i32::cast_from((byte & 0xF) | (lo_bit << 4)) - 16);
let hi = f32::cast_from(i32::cast_from((byte >> 4) | (hi_bit << 4)) - 16);
block_acc += lo * x[x_base + jj];
block_acc += hi * x[x_base + 16 + jj];
}
}
acc += d[block] * block_acc;
}
out[row * n_out + col] = acc;
}
}
#[cube(launch_unchecked)]
fn q8_0_dequant_kernel(qs: &Array<u32>, d: &Array<f32>, out: &mut Array<f32>, n: usize) {
if ABSOLUTE_POS < n {
let block = ABSOLUTE_POS / 32;
let j = ABSOLUTE_POS % 32;
let word = qs[block * 8 + j / 4];
let byte = (word >> (u32::cast_from(j % 4) * 8)) & 0xFF;
let q = (i32::cast_from(byte) << 24) >> 24;
out[ABSOLUTE_POS] = f32::cast_from(q) * d[block];
}
}
#[cube(launch_unchecked)]
fn q8_0_matmul_kernel(
x: &Array<f32>,
qs: &Array<u32>,
d: &Array<f32>,
out: &mut Array<f32>,
m: usize,
k: usize,
n_out: usize,
) {
if ABSOLUTE_POS < m * n_out {
let row = ABSOLUTE_POS / n_out;
let col = ABSOLUTE_POS % n_out;
let blocks_per_row = k / 32;
let mut acc = 0.0f32;
for kb in 0..blocks_per_row {
let block = col * blocks_per_row + kb;
let x_base = row * k + kb * 32;
let mut block_acc = 0.0f32;
for w in 0..8usize {
let word = qs[block * 8 + w];
for b in 0..4usize {
let byte = (word >> (u32::cast_from(b) * 8)) & 0xFF;
let q = (i32::cast_from(byte) << 24) >> 24;
block_acc += f32::cast_from(q) * x[x_base + w * 4 + b];
}
}
acc += d[block] * block_acc;
}
out[row * n_out + col] = acc;
}
}
#[cube(launch_unchecked)]
fn q8_0_matmul_tiled_kernel(
x: &Array<f32>,
qs: &Array<u32>,
d: &Array<f32>,
out: &mut Array<f32>,
k: usize,
n_out: usize,
) {
let mut staged = SharedMemory::<f32>::new(256usize);
let unit = UNIT_POS as usize;
let row = CUBE_POS_Y as usize;
let col = (CUBE_POS_X * CUBE_DIM + UNIT_POS) as usize;
let blocks_per_row = k / 32;
let n_tiles = (k + 255) / 256;
let mut acc = 0.0f32;
for t in 0..n_tiles {
let k0 = t * 256;
if k0 + unit < k {
staged[unit] = x[row * k + k0 + unit];
}
sync_cube();
if col < n_out {
let mut kb_end = (k0 + 256) / 32;
if blocks_per_row < kb_end {
kb_end = blocks_per_row;
}
for kb in (k0 / 32)..kb_end {
let block = col * blocks_per_row + kb;
let s_base = kb * 32 - k0;
let mut block_acc = 0.0f32;
for w in 0..8usize {
let word = qs[block * 8 + w];
for b in 0..4usize {
let byte = (word >> (u32::cast_from(b) * 8)) & 0xFF;
let q = (i32::cast_from(byte) << 24) >> 24;
block_acc += f32::cast_from(q) * staged[s_base + w * 4 + b];
}
}
acc += d[block] * block_acc;
}
}
sync_cube();
}
if col < n_out {
out[row * n_out + col] = acc;
}
}
pub fn dequantize_q5_0_gpu<R: Runtime>(client: &ComputeClient<R>, data: &[u8]) -> Result<Vec<f32>> {
let (qs, qh, d) = repack_q5_0(data)?;
let n = d.len() * Q4_0_BLOCK;
let qs_h = client.create_from_slice(u32::as_bytes(&qs));
let qh_h = client.create_from_slice(u32::as_bytes(&qh));
let d_h = client.create_from_slice(f32::as_bytes(&d));
let out_h = client.empty(n * core::mem::size_of::<f32>());
unsafe {
q5_0_dequant_kernel::launch_unchecked::<R>(
client,
cube_count_1d(n as u32),
CubeDim::new_1d(CUBE_DIM),
ArrayArg::from_raw_parts(qs_h, qs.len()),
ArrayArg::from_raw_parts(qh_h, qh.len()),
ArrayArg::from_raw_parts(d_h, d.len()),
ArrayArg::from_raw_parts(out_h.clone(), n),
n,
);
}
let bytes = client.read_one_unchecked(out_h);
Ok(f32::from_bytes(&bytes).to_vec())
}
pub fn dequantize_q8_0_gpu<R: Runtime>(client: &ComputeClient<R>, data: &[u8]) -> Result<Vec<f32>> {
let (qs, d) = repack_q8_0(data)?;
let n = d.len() * Q4_0_BLOCK;
let qs_h = client.create_from_slice(u32::as_bytes(&qs));
let d_h = client.create_from_slice(f32::as_bytes(&d));
let out_h = client.empty(n * core::mem::size_of::<f32>());
unsafe {
q8_0_dequant_kernel::launch_unchecked::<R>(
client,
cube_count_1d(n as u32),
CubeDim::new_1d(CUBE_DIM),
ArrayArg::from_raw_parts(qs_h, qs.len()),
ArrayArg::from_raw_parts(d_h, d.len()),
ArrayArg::from_raw_parts(out_h.clone(), n),
n,
);
}
let bytes = client.read_one_unchecked(out_h);
Ok(f32::from_bytes(&bytes).to_vec())
}
pub struct Q50Weight<R: Runtime> {
qs: Handle,
qh: Handle,
d: Handle,
n_out: usize,
k: usize,
_runtime: PhantomData<R>,
}
impl<R: Runtime> Q50Weight<R> {
pub fn from_gguf_bytes(
client: &ComputeClient<R>,
data: &[u8],
n_out: usize,
k: usize,
) -> Result<Self> {
if k == 0 || k % Q4_0_BLOCK != 0 || data.len() != n_out * k / Q4_0_BLOCK * Q5_0_BLOCK_BYTES
{
return Err(ModelError::BadShape {
tensor: "q5_0 weight".into(),
expected: vec![n_out, k],
got: vec![data.len()],
});
}
let (qs, qh, d) = repack_q5_0(data)?;
Ok(Q50Weight {
qs: client.create_from_slice(u32::as_bytes(&qs)),
qh: client.create_from_slice(u32::as_bytes(&qh)),
d: client.create_from_slice(f32::as_bytes(&d)),
n_out,
k,
_runtime: PhantomData,
})
}
pub fn vram_bytes(&self) -> usize {
(self.n_out * self.k / Q4_0_BLOCK) * 24
}
pub fn matmul_device(&self, client: &ComputeClient<R>, x: Handle, m: usize) -> Handle {
let out_len = m * self.n_out;
let out_h = client.empty(out_len * core::mem::size_of::<f32>());
let n_blocks = self.n_out * self.k / Q4_0_BLOCK;
unsafe {
q5_0_matmul_kernel::launch_unchecked::<R>(
client,
cube_count_capped(out_len as u32),
CubeDim::new_1d(CUBE_DIM),
ArrayArg::from_raw_parts(x, m * self.k),
ArrayArg::from_raw_parts(self.qs.clone(), n_blocks * 4),
ArrayArg::from_raw_parts(self.qh.clone(), n_blocks),
ArrayArg::from_raw_parts(self.d.clone(), n_blocks),
ArrayArg::from_raw_parts(out_h.clone(), out_len),
m,
self.k,
self.n_out,
);
}
out_h
}
pub fn matmul_host(&self, client: &ComputeClient<R>, x: &[f32], m: usize) -> Result<Vec<f32>> {
if m == 0 || x.len() != m * self.k {
return Err(ModelError::BadShape {
tensor: "q5_0 matmul input".into(),
expected: vec![m, self.k],
got: vec![x.len()],
});
}
let x_h = client.create_from_slice(f32::as_bytes(x));
let out_h = self.matmul_device(client, x_h, m);
let bytes = client.read_one_unchecked(out_h);
Ok(f32::from_bytes(&bytes).to_vec())
}
}
pub struct Q80Weight<R: Runtime> {
qs: Handle,
d: Handle,
n_out: usize,
k: usize,
_runtime: PhantomData<R>,
}
impl<R: Runtime> Q80Weight<R> {
pub fn from_gguf_bytes(
client: &ComputeClient<R>,
data: &[u8],
n_out: usize,
k: usize,
) -> Result<Self> {
if k == 0 || k % Q4_0_BLOCK != 0 || data.len() != n_out * k / Q4_0_BLOCK * Q8_0_BLOCK_BYTES
{
return Err(ModelError::BadShape {
tensor: "q8_0 weight".into(),
expected: vec![n_out, k],
got: vec![data.len()],
});
}
let (qs, d) = repack_q8_0(data)?;
Ok(Q80Weight {
qs: client.create_from_slice(u32::as_bytes(&qs)),
d: client.create_from_slice(f32::as_bytes(&d)),
n_out,
k,
_runtime: PhantomData,
})
}
pub fn vram_bytes(&self) -> usize {
(self.n_out * self.k / Q4_0_BLOCK) * 36
}
pub fn matmul_device(&self, client: &ComputeClient<R>, x: Handle, m: usize) -> Handle {
self.matmul_device_with(client, x, m, m > 1 && tiled_enabled())
}
pub(crate) fn matmul_device_with(
&self,
client: &ComputeClient<R>,
x: Handle,
m: usize,
tiled: bool,
) -> Handle {
let out_len = m * self.n_out;
let out_h = client.empty(out_len * core::mem::size_of::<f32>());
let n_blocks = self.n_out * self.k / Q4_0_BLOCK;
if tiled {
unsafe {
q8_0_matmul_tiled_kernel::launch_unchecked::<R>(
client,
cube_count_tiled(self.n_out as u32, m as u32),
CubeDim::new_1d(CUBE_DIM),
ArrayArg::from_raw_parts(x, m * self.k),
ArrayArg::from_raw_parts(self.qs.clone(), n_blocks * 8),
ArrayArg::from_raw_parts(self.d.clone(), n_blocks),
ArrayArg::from_raw_parts(out_h.clone(), out_len),
self.k,
self.n_out,
);
}
} else {
unsafe {
q8_0_matmul_kernel::launch_unchecked::<R>(
client,
cube_count_capped(out_len as u32),
CubeDim::new_1d(CUBE_DIM),
ArrayArg::from_raw_parts(x, m * self.k),
ArrayArg::from_raw_parts(self.qs.clone(), n_blocks * 8),
ArrayArg::from_raw_parts(self.d.clone(), n_blocks),
ArrayArg::from_raw_parts(out_h.clone(), out_len),
m,
self.k,
self.n_out,
);
}
}
out_h
}
pub fn matmul_host(&self, client: &ComputeClient<R>, x: &[f32], m: usize) -> Result<Vec<f32>> {
if m == 0 || x.len() != m * self.k {
return Err(ModelError::BadShape {
tensor: "q8_0 matmul input".into(),
expected: vec![m, self.k],
got: vec![x.len()],
});
}
let x_h = client.create_from_slice(f32::as_bytes(x));
let out_h = self.matmul_device(client, x_h, m);
let bytes = client.read_one_unchecked(out_h);
Ok(f32::from_bytes(&bytes).to_vec())
}
}
pub const K_SUPERBLOCK: usize = 256;
pub const Q4_K_BLOCK_BYTES: usize = 144;
pub const Q5_K_BLOCK_BYTES: usize = 176;
pub const Q6_K_BLOCK_BYTES: usize = 210;
#[cube]
fn byte_at(words: &Array<u32>, idx: usize) -> u32 {
(words[idx / 4] >> (u32::cast_from(idx % 4) * 8)) & 0xFF
}
#[cube]
fn i8_at(words: &Array<u32>, idx: usize) -> i32 {
(i32::cast_from(byte_at(words, idx)) << 24) >> 24
}
#[cube]
fn k4_scale(scales: &Array<u32>, base: usize, j: usize) -> u32 {
let mut v = 0u32;
if j < 4 {
v = byte_at(scales, base + j) & 63;
} else {
v = (byte_at(scales, base + j + 4) & 0xF) | ((byte_at(scales, base + j - 4) >> 6) << 4);
}
v
}
#[cube]
fn k4_min(scales: &Array<u32>, base: usize, j: usize) -> u32 {
let mut v = 0u32;
if j < 4 {
v = byte_at(scales, base + j + 4) & 63;
} else {
v = (byte_at(scales, base + j + 4) >> 4) | ((byte_at(scales, base + j) >> 6) << 4);
}
v
}
pub fn repack_q4_k(data: &[u8]) -> Result<(Vec<u32>, Vec<f32>, Vec<u32>)> {
if data.is_empty() || data.len() % Q4_K_BLOCK_BYTES != 0 {
return Err(ModelError::BadShape {
tensor: "q4_k superblock stream".into(),
expected: vec![Q4_K_BLOCK_BYTES],
got: vec![data.len()],
});
}
let n_sb = data.len() / Q4_K_BLOCK_BYTES;
let mut qs = Vec::with_capacity(n_sb * 32);
let mut dd = Vec::with_capacity(n_sb * 2);
let mut scales = Vec::with_capacity(n_sb * 3);
for sb in data.chunks_exact(Q4_K_BLOCK_BYTES) {
dd.push(burn::tensor::f16::from_le_bytes([sb[0], sb[1]]).to_f32());
dd.push(burn::tensor::f16::from_le_bytes([sb[2], sb[3]]).to_f32());
for w in 0..3 {
let o = 4 + 4 * w;
scales.push(u32::from_le_bytes([sb[o], sb[o + 1], sb[o + 2], sb[o + 3]]));
}
for w in 0..32 {
let o = 16 + 4 * w;
qs.push(u32::from_le_bytes([sb[o], sb[o + 1], sb[o + 2], sb[o + 3]]));
}
}
Ok((qs, dd, scales))
}
#[cube(launch_unchecked)]
fn q4_k_dequant_kernel(
qs: &Array<u32>,
dd: &Array<f32>,
scales: &Array<u32>,
out: &mut Array<f32>,
n: usize,
) {
if ABSOLUTE_POS < n {
let sb = ABSOLUTE_POS / 256;
let r = ABSOLUTE_POS % 256;
let j = r / 64; let t = (r % 64) / 32; let l = r % 32;
let byte = byte_at(qs, sb * 128 + j * 32 + l);
let mut q = byte & 0xF;
if t == 1 {
q = byte >> 4;
}
let sidx = 2 * j + t;
let sc = k4_scale(scales, sb * 12, sidx);
let mn = k4_min(scales, sb * 12, sidx);
let d1 = dd[sb * 2] * f32::cast_from(sc);
let fmin = dd[sb * 2 + 1] * f32::cast_from(mn);
out[ABSOLUTE_POS] = d1 * f32::cast_from(q) - fmin;
}
}
#[cube(launch_unchecked)]
fn q4_k_matmul_kernel(
x: &Array<f32>,
qs: &Array<u32>,
dd: &Array<f32>,
scales: &Array<u32>,
out: &mut Array<f32>,
m: usize,
k: usize,
n_out: usize,
) {
if ABSOLUTE_POS < m * n_out {
let row = ABSOLUTE_POS / n_out;
let col = ABSOLUTE_POS % n_out;
let sb_per_row = k / 256;
let mut acc = 0.0f32;
for sbi in 0..sb_per_row {
let sb = col * sb_per_row + sbi;
let d = dd[sb * 2];
let dmin = dd[sb * 2 + 1];
let s_base = sb * 12;
let x_base = row * k + sbi * 256;
for j in 0..4usize {
let mut sum_lo = 0.0f32;
let mut sum_hi = 0.0f32;
let mut xs_lo = 0.0f32;
let mut xs_hi = 0.0f32;
for w in 0..8usize {
let word = qs[sb * 32 + j * 8 + w];
for b in 0..4usize {
let byte = (word >> (u32::cast_from(b) * 8)) & 0xFF;
let l = 4 * w + b;
let x1 = x[x_base + 64 * j + l];
let x2 = x[x_base + 64 * j + 32 + l];
sum_lo += f32::cast_from(byte & 0xF) * x1;
sum_hi += f32::cast_from(byte >> 4) * x2;
xs_lo += x1;
xs_hi += x2;
}
}
let sc1 = f32::cast_from(k4_scale(scales, s_base, 2 * j));
let mn1 = f32::cast_from(k4_min(scales, s_base, 2 * j));
let sc2 = f32::cast_from(k4_scale(scales, s_base, 2 * j + 1));
let mn2 = f32::cast_from(k4_min(scales, s_base, 2 * j + 1));
acc += d * sc1 * sum_lo - dmin * mn1 * xs_lo;
acc += d * sc2 * sum_hi - dmin * mn2 * xs_hi;
}
}
out[row * n_out + col] = acc;
}
}
#[cube(launch_unchecked)]
fn q4_k_matmul_tiled_kernel(
x: &Array<f32>,
qs: &Array<u32>,
dd: &Array<f32>,
scales: &Array<u32>,
out: &mut Array<f32>,
k: usize,
n_out: usize,
) {
let mut staged = SharedMemory::<f32>::new(256usize);
let unit = UNIT_POS as usize;
let row = CUBE_POS_Y as usize;
let col = (CUBE_POS_X * CUBE_DIM + UNIT_POS) as usize;
let sb_per_row = k / 256;
let mut acc = 0.0f32;
for sbi in 0..sb_per_row {
staged[unit] = x[row * k + sbi * 256 + unit];
sync_cube();
if col < n_out {
let sb = col * sb_per_row + sbi;
let d = dd[sb * 2];
let dmin = dd[sb * 2 + 1];
let s_base = sb * 12;
for j in 0..4usize {
let mut sum_lo = 0.0f32;
let mut sum_hi = 0.0f32;
let mut xs_lo = 0.0f32;
let mut xs_hi = 0.0f32;
for w in 0..8usize {
let word = qs[sb * 32 + j * 8 + w];
for b in 0..4usize {
let byte = (word >> (u32::cast_from(b) * 8)) & 0xFF;
let l = 4 * w + b;
let x1 = staged[64 * j + l];
let x2 = staged[64 * j + 32 + l];
sum_lo += f32::cast_from(byte & 0xF) * x1;
sum_hi += f32::cast_from(byte >> 4) * x2;
xs_lo += x1;
xs_hi += x2;
}
}
let sc1 = f32::cast_from(k4_scale(scales, s_base, 2 * j));
let mn1 = f32::cast_from(k4_min(scales, s_base, 2 * j));
let sc2 = f32::cast_from(k4_scale(scales, s_base, 2 * j + 1));
let mn2 = f32::cast_from(k4_min(scales, s_base, 2 * j + 1));
acc += d * sc1 * sum_lo - dmin * mn1 * xs_lo;
acc += d * sc2 * sum_hi - dmin * mn2 * xs_hi;
}
}
sync_cube();
}
if col < n_out {
out[row * n_out + col] = acc;
}
}
pub fn repack_q5_k(data: &[u8]) -> Result<(Vec<u32>, Vec<u32>, Vec<f32>, Vec<u32>)> {
if data.is_empty() || data.len() % Q5_K_BLOCK_BYTES != 0 {
return Err(ModelError::BadShape {
tensor: "q5_k superblock stream".into(),
expected: vec![Q5_K_BLOCK_BYTES],
got: vec![data.len()],
});
}
let n_sb = data.len() / Q5_K_BLOCK_BYTES;
let word = |sb: &[u8], o: usize| u32::from_le_bytes([sb[o], sb[o + 1], sb[o + 2], sb[o + 3]]);
let mut qs = Vec::with_capacity(n_sb * 32);
let mut qh = Vec::with_capacity(n_sb * 8);
let mut dd = Vec::with_capacity(n_sb * 2);
let mut scales = Vec::with_capacity(n_sb * 3);
for sb in data.chunks_exact(Q5_K_BLOCK_BYTES) {
dd.push(burn::tensor::f16::from_le_bytes([sb[0], sb[1]]).to_f32());
dd.push(burn::tensor::f16::from_le_bytes([sb[2], sb[3]]).to_f32());
for w in 0..3 {
scales.push(word(sb, 4 + 4 * w));
}
for w in 0..8 {
qh.push(word(sb, 16 + 4 * w));
}
for w in 0..32 {
qs.push(word(sb, 48 + 4 * w));
}
}
Ok((qs, qh, dd, scales))
}
#[cube(launch_unchecked)]
fn q5_k_dequant_kernel(
qs: &Array<u32>,
qh: &Array<u32>,
dd: &Array<f32>,
scales: &Array<u32>,
out: &mut Array<f32>,
n: usize,
) {
if ABSOLUTE_POS < n {
let sb = ABSOLUTE_POS / 256;
let r = ABSOLUTE_POS % 256;
let j = r / 64; let t = (r % 64) / 32; let l = r % 32;
let byte = byte_at(qs, sb * 128 + j * 32 + l);
let mut nib = byte & 0xF;
if t == 1 {
nib = byte >> 4;
}
let hi = (byte_at(qh, sb * 32 + l) >> u32::cast_from(2 * j + t)) & 1;
let sidx = 2 * j + t;
let sc = k4_scale(scales, sb * 12, sidx);
let mn = k4_min(scales, sb * 12, sidx);
let d1 = dd[sb * 2] * f32::cast_from(sc);
let fmin = dd[sb * 2 + 1] * f32::cast_from(mn);
out[ABSOLUTE_POS] = d1 * f32::cast_from(nib | (hi << 4)) - fmin;
}
}
#[cube(launch_unchecked)]
fn q5_k_matmul_kernel(
x: &Array<f32>,
qs: &Array<u32>,
qh: &Array<u32>,
dd: &Array<f32>,
scales: &Array<u32>,
out: &mut Array<f32>,
m: usize,
k: usize,
n_out: usize,
) {
if ABSOLUTE_POS < m * n_out {
let row = ABSOLUTE_POS / n_out;
let col = ABSOLUTE_POS % n_out;
let sb_per_row = k / 256;
let mut acc = 0.0f32;
for sbi in 0..sb_per_row {
let sb = col * sb_per_row + sbi;
let d = dd[sb * 2];
let dmin = dd[sb * 2 + 1];
let s_base = sb * 12;
let x_base = row * k + sbi * 256;
for j in 0..4usize {
let mut sum_lo = 0.0f32;
let mut sum_hi = 0.0f32;
let mut xs_lo = 0.0f32;
let mut xs_hi = 0.0f32;
for w in 0..8usize {
let word = qs[sb * 32 + j * 8 + w];
for b in 0..4usize {
let byte = (word >> (u32::cast_from(b) * 8)) & 0xFF;
let l = 4 * w + b;
let hb = byte_at(qh, sb * 32 + l);
let hi_lo = (hb >> u32::cast_from(2 * j)) & 1;
let hi_hi = (hb >> u32::cast_from(2 * j + 1)) & 1;
let x1 = x[x_base + 64 * j + l];
let x2 = x[x_base + 64 * j + 32 + l];
sum_lo += f32::cast_from((byte & 0xF) | (hi_lo << 4)) * x1;
sum_hi += f32::cast_from((byte >> 4) | (hi_hi << 4)) * x2;
xs_lo += x1;
xs_hi += x2;
}
}
let sc1 = f32::cast_from(k4_scale(scales, s_base, 2 * j));
let mn1 = f32::cast_from(k4_min(scales, s_base, 2 * j));
let sc2 = f32::cast_from(k4_scale(scales, s_base, 2 * j + 1));
let mn2 = f32::cast_from(k4_min(scales, s_base, 2 * j + 1));
acc += d * sc1 * sum_lo - dmin * mn1 * xs_lo;
acc += d * sc2 * sum_hi - dmin * mn2 * xs_hi;
}
}
out[row * n_out + col] = acc;
}
}
pub fn repack_q6_k(data: &[u8]) -> Result<(Vec<u32>, Vec<u32>, Vec<u32>, Vec<f32>)> {
if data.is_empty() || data.len() % Q6_K_BLOCK_BYTES != 0 {
return Err(ModelError::BadShape {
tensor: "q6_k superblock stream".into(),
expected: vec![Q6_K_BLOCK_BYTES],
got: vec![data.len()],
});
}
let n_sb = data.len() / Q6_K_BLOCK_BYTES;
let word = |sb: &[u8], o: usize| u32::from_le_bytes([sb[o], sb[o + 1], sb[o + 2], sb[o + 3]]);
let mut ql = Vec::with_capacity(n_sb * 32);
let mut qh = Vec::with_capacity(n_sb * 16);
let mut sc = Vec::with_capacity(n_sb * 4);
let mut d = Vec::with_capacity(n_sb);
for sb in data.chunks_exact(Q6_K_BLOCK_BYTES) {
for w in 0..32 {
ql.push(word(sb, 4 * w));
}
for w in 0..16 {
qh.push(word(sb, 128 + 4 * w));
}
for w in 0..4 {
sc.push(word(sb, 192 + 4 * w));
}
d.push(burn::tensor::f16::from_le_bytes([sb[208], sb[209]]).to_f32());
}
Ok((ql, qh, sc, d))
}
#[cube(launch_unchecked)]
fn q6_k_dequant_kernel(
ql: &Array<u32>,
qh: &Array<u32>,
sc: &Array<u32>,
d: &Array<f32>,
out: &mut Array<f32>,
n: usize,
) {
if ABSOLUTE_POS < n {
let sb = ABSOLUTE_POS / 256;
let r = ABSOLUTE_POS % 256;
let half = r / 128;
let t = (r % 128) / 32; let l = r % 32;
let ql_byte = byte_at(ql, sb * 128 + half * 64 + (t % 2) * 32 + l);
let mut nib = ql_byte & 0xF;
if t >= 2 {
nib = ql_byte >> 4;
}
let hi = (byte_at(qh, sb * 64 + half * 32 + l) >> (u32::cast_from(t) * 2)) & 3;
let q = i32::cast_from(nib | (hi << 4)) - 32;
let scale = i8_at(sc, sb * 16 + half * 8 + l / 16 + 2 * t);
out[ABSOLUTE_POS] = d[sb] * f32::cast_from(scale) * f32::cast_from(q);
}
}
#[cube(launch_unchecked)]
fn q6_k_matmul_kernel(
x: &Array<f32>,
ql: &Array<u32>,
qh: &Array<u32>,
sc: &Array<u32>,
d: &Array<f32>,
out: &mut Array<f32>,
m: usize,
k: usize,
n_out: usize,
) {
if ABSOLUTE_POS < m * n_out {
let row = ABSOLUTE_POS / n_out;
let col = ABSOLUTE_POS % n_out;
let sb_per_row = k / 256;
let mut acc = 0.0f32;
for sbi in 0..sb_per_row {
let sb = col * sb_per_row + sbi;
let dsb = d[sb];
let x_base = row * k + sbi * 256;
for half in 0..2usize {
for t in 0..4usize {
for g in 0..2usize {
let mut sum = 0.0f32;
for l0 in 0..16usize {
let l = g * 16 + l0;
let ql_byte = byte_at(ql, sb * 128 + half * 64 + (t % 2) * 32 + l);
let mut nib = ql_byte & 0xF;
if t >= 2 {
nib = ql_byte >> 4;
}
let hi =
(byte_at(qh, sb * 64 + half * 32 + l) >> (u32::cast_from(t) * 2))
& 3;
let q = i32::cast_from(nib | (hi << 4)) - 32;
sum += f32::cast_from(q) * x[x_base + half * 128 + t * 32 + l];
}
let scale = i8_at(sc, sb * 16 + half * 8 + g + 2 * t);
acc += dsb * f32::cast_from(scale) * sum;
}
}
}
}
out[row * n_out + col] = acc;
}
}
pub fn dequantize_q4_k_gpu<R: Runtime>(client: &ComputeClient<R>, data: &[u8]) -> Result<Vec<f32>> {
let (qs, dd, scales) = repack_q4_k(data)?;
let n = (dd.len() / 2) * K_SUPERBLOCK;
let qs_h = client.create_from_slice(u32::as_bytes(&qs));
let dd_h = client.create_from_slice(f32::as_bytes(&dd));
let sc_h = client.create_from_slice(u32::as_bytes(&scales));
let out_h = client.empty(n * core::mem::size_of::<f32>());
unsafe {
q4_k_dequant_kernel::launch_unchecked::<R>(
client,
cube_count_1d(n as u32),
CubeDim::new_1d(CUBE_DIM),
ArrayArg::from_raw_parts(qs_h, qs.len()),
ArrayArg::from_raw_parts(dd_h, dd.len()),
ArrayArg::from_raw_parts(sc_h, scales.len()),
ArrayArg::from_raw_parts(out_h.clone(), n),
n,
);
}
let bytes = client.read_one_unchecked(out_h);
Ok(f32::from_bytes(&bytes).to_vec())
}
pub fn dequantize_q5_k_gpu<R: Runtime>(client: &ComputeClient<R>, data: &[u8]) -> Result<Vec<f32>> {
let (qs, qh, dd, scales) = repack_q5_k(data)?;
let n = (dd.len() / 2) * K_SUPERBLOCK;
let qs_h = client.create_from_slice(u32::as_bytes(&qs));
let qh_h = client.create_from_slice(u32::as_bytes(&qh));
let dd_h = client.create_from_slice(f32::as_bytes(&dd));
let sc_h = client.create_from_slice(u32::as_bytes(&scales));
let out_h = client.empty(n * core::mem::size_of::<f32>());
unsafe {
q5_k_dequant_kernel::launch_unchecked::<R>(
client,
cube_count_1d(n as u32),
CubeDim::new_1d(CUBE_DIM),
ArrayArg::from_raw_parts(qs_h, qs.len()),
ArrayArg::from_raw_parts(qh_h, qh.len()),
ArrayArg::from_raw_parts(dd_h, dd.len()),
ArrayArg::from_raw_parts(sc_h, scales.len()),
ArrayArg::from_raw_parts(out_h.clone(), n),
n,
);
}
let bytes = client.read_one_unchecked(out_h);
Ok(f32::from_bytes(&bytes).to_vec())
}
pub fn dequantize_q6_k_gpu<R: Runtime>(client: &ComputeClient<R>, data: &[u8]) -> Result<Vec<f32>> {
let (ql, qh, sc, d) = repack_q6_k(data)?;
let n = d.len() * K_SUPERBLOCK;
let ql_h = client.create_from_slice(u32::as_bytes(&ql));
let qh_h = client.create_from_slice(u32::as_bytes(&qh));
let sc_h = client.create_from_slice(u32::as_bytes(&sc));
let d_h = client.create_from_slice(f32::as_bytes(&d));
let out_h = client.empty(n * core::mem::size_of::<f32>());
unsafe {
q6_k_dequant_kernel::launch_unchecked::<R>(
client,
cube_count_1d(n as u32),
CubeDim::new_1d(CUBE_DIM),
ArrayArg::from_raw_parts(ql_h, ql.len()),
ArrayArg::from_raw_parts(qh_h, qh.len()),
ArrayArg::from_raw_parts(sc_h, sc.len()),
ArrayArg::from_raw_parts(d_h, d.len()),
ArrayArg::from_raw_parts(out_h.clone(), n),
n,
);
}
let bytes = client.read_one_unchecked(out_h);
Ok(f32::from_bytes(&bytes).to_vec())
}
pub struct Q4KWeight<R: Runtime> {
qs: Handle,
dd: Handle,
scales: Handle,
n_out: usize,
k: usize,
_runtime: PhantomData<R>,
}
impl<R: Runtime> Q4KWeight<R> {
pub fn from_gguf_bytes(
client: &ComputeClient<R>,
data: &[u8],
n_out: usize,
k: usize,
) -> Result<Self> {
if k == 0
|| k % K_SUPERBLOCK != 0
|| data.len() != n_out * k / K_SUPERBLOCK * Q4_K_BLOCK_BYTES
{
return Err(ModelError::BadShape {
tensor: "q4_k weight".into(),
expected: vec![n_out, k],
got: vec![data.len()],
});
}
let (qs, dd, scales) = repack_q4_k(data)?;
Ok(Q4KWeight {
qs: client.create_from_slice(u32::as_bytes(&qs)),
dd: client.create_from_slice(f32::as_bytes(&dd)),
scales: client.create_from_slice(u32::as_bytes(&scales)),
n_out,
k,
_runtime: PhantomData,
})
}
pub fn vram_bytes(&self) -> usize {
(self.n_out * self.k / K_SUPERBLOCK) * (128 + 12 + 8)
}
pub fn matmul_device(&self, client: &ComputeClient<R>, x: Handle, m: usize) -> Handle {
self.matmul_device_with(client, x, m, m > 1 && tiled_enabled())
}
pub(crate) fn matmul_device_with(
&self,
client: &ComputeClient<R>,
x: Handle,
m: usize,
tiled: bool,
) -> Handle {
let out_len = m * self.n_out;
let out_h = client.empty(out_len * core::mem::size_of::<f32>());
let n_sb = self.n_out * self.k / K_SUPERBLOCK;
if tiled {
unsafe {
q4_k_matmul_tiled_kernel::launch_unchecked::<R>(
client,
cube_count_tiled(self.n_out as u32, m as u32),
CubeDim::new_1d(CUBE_DIM),
ArrayArg::from_raw_parts(x, m * self.k),
ArrayArg::from_raw_parts(self.qs.clone(), n_sb * 32),
ArrayArg::from_raw_parts(self.dd.clone(), n_sb * 2),
ArrayArg::from_raw_parts(self.scales.clone(), n_sb * 3),
ArrayArg::from_raw_parts(out_h.clone(), out_len),
self.k,
self.n_out,
);
}
} else {
unsafe {
q4_k_matmul_kernel::launch_unchecked::<R>(
client,
cube_count_capped(out_len as u32),
CubeDim::new_1d(CUBE_DIM),
ArrayArg::from_raw_parts(x, m * self.k),
ArrayArg::from_raw_parts(self.qs.clone(), n_sb * 32),
ArrayArg::from_raw_parts(self.dd.clone(), n_sb * 2),
ArrayArg::from_raw_parts(self.scales.clone(), n_sb * 3),
ArrayArg::from_raw_parts(out_h.clone(), out_len),
m,
self.k,
self.n_out,
);
}
}
out_h
}
pub fn matmul_host(&self, client: &ComputeClient<R>, x: &[f32], m: usize) -> Result<Vec<f32>> {
if m == 0 || x.len() != m * self.k {
return Err(ModelError::BadShape {
tensor: "q4_k matmul input".into(),
expected: vec![m, self.k],
got: vec![x.len()],
});
}
let x_h = client.create_from_slice(f32::as_bytes(x));
let out_h = self.matmul_device(client, x_h, m);
let bytes = client.read_one_unchecked(out_h);
Ok(f32::from_bytes(&bytes).to_vec())
}
}
pub struct Q5KWeight<R: Runtime> {
qs: Handle,
qh: Handle,
dd: Handle,
scales: Handle,
n_out: usize,
k: usize,
_runtime: PhantomData<R>,
}
impl<R: Runtime> Q5KWeight<R> {
pub fn from_gguf_bytes(
client: &ComputeClient<R>,
data: &[u8],
n_out: usize,
k: usize,
) -> Result<Self> {
if k == 0
|| k % K_SUPERBLOCK != 0
|| data.len() != n_out * k / K_SUPERBLOCK * Q5_K_BLOCK_BYTES
{
return Err(ModelError::BadShape {
tensor: "q5_k weight".into(),
expected: vec![n_out, k],
got: vec![data.len()],
});
}
let (qs, qh, dd, scales) = repack_q5_k(data)?;
Ok(Q5KWeight {
qs: client.create_from_slice(u32::as_bytes(&qs)),
qh: client.create_from_slice(u32::as_bytes(&qh)),
dd: client.create_from_slice(f32::as_bytes(&dd)),
scales: client.create_from_slice(u32::as_bytes(&scales)),
n_out,
k,
_runtime: PhantomData,
})
}
pub fn vram_bytes(&self) -> usize {
(self.n_out * self.k / K_SUPERBLOCK) * (128 + 32 + 12 + 8)
}
pub fn matmul_device(&self, client: &ComputeClient<R>, x: Handle, m: usize) -> Handle {
let out_len = m * self.n_out;
let out_h = client.empty(out_len * core::mem::size_of::<f32>());
let n_sb = self.n_out * self.k / K_SUPERBLOCK;
unsafe {
q5_k_matmul_kernel::launch_unchecked::<R>(
client,
cube_count_capped(out_len as u32),
CubeDim::new_1d(CUBE_DIM),
ArrayArg::from_raw_parts(x, m * self.k),
ArrayArg::from_raw_parts(self.qs.clone(), n_sb * 32),
ArrayArg::from_raw_parts(self.qh.clone(), n_sb * 8),
ArrayArg::from_raw_parts(self.dd.clone(), n_sb * 2),
ArrayArg::from_raw_parts(self.scales.clone(), n_sb * 3),
ArrayArg::from_raw_parts(out_h.clone(), out_len),
m,
self.k,
self.n_out,
);
}
out_h
}
pub fn matmul_host(&self, client: &ComputeClient<R>, x: &[f32], m: usize) -> Result<Vec<f32>> {
if m == 0 || x.len() != m * self.k {
return Err(ModelError::BadShape {
tensor: "q5_k matmul input".into(),
expected: vec![m, self.k],
got: vec![x.len()],
});
}
let x_h = client.create_from_slice(f32::as_bytes(x));
let out_h = self.matmul_device(client, x_h, m);
let bytes = client.read_one_unchecked(out_h);
Ok(f32::from_bytes(&bytes).to_vec())
}
}
pub struct Q6KWeight<R: Runtime> {
ql: Handle,
qh: Handle,
sc: Handle,
d: Handle,
n_out: usize,
k: usize,
_runtime: PhantomData<R>,
}
impl<R: Runtime> Q6KWeight<R> {
pub fn from_gguf_bytes(
client: &ComputeClient<R>,
data: &[u8],
n_out: usize,
k: usize,
) -> Result<Self> {
if k == 0
|| k % K_SUPERBLOCK != 0
|| data.len() != n_out * k / K_SUPERBLOCK * Q6_K_BLOCK_BYTES
{
return Err(ModelError::BadShape {
tensor: "q6_k weight".into(),
expected: vec![n_out, k],
got: vec![data.len()],
});
}
let (ql, qh, sc, d) = repack_q6_k(data)?;
Ok(Q6KWeight {
ql: client.create_from_slice(u32::as_bytes(&ql)),
qh: client.create_from_slice(u32::as_bytes(&qh)),
sc: client.create_from_slice(u32::as_bytes(&sc)),
d: client.create_from_slice(f32::as_bytes(&d)),
n_out,
k,
_runtime: PhantomData,
})
}
pub fn vram_bytes(&self) -> usize {
(self.n_out * self.k / K_SUPERBLOCK) * (128 + 64 + 16 + 4)
}
pub fn matmul_device(&self, client: &ComputeClient<R>, x: Handle, m: usize) -> Handle {
let out_len = m * self.n_out;
let out_h = client.empty(out_len * core::mem::size_of::<f32>());
let n_sb = self.n_out * self.k / K_SUPERBLOCK;
unsafe {
q6_k_matmul_kernel::launch_unchecked::<R>(
client,
cube_count_capped(out_len as u32),
CubeDim::new_1d(CUBE_DIM),
ArrayArg::from_raw_parts(x, m * self.k),
ArrayArg::from_raw_parts(self.ql.clone(), n_sb * 32),
ArrayArg::from_raw_parts(self.qh.clone(), n_sb * 16),
ArrayArg::from_raw_parts(self.sc.clone(), n_sb * 4),
ArrayArg::from_raw_parts(self.d.clone(), n_sb),
ArrayArg::from_raw_parts(out_h.clone(), out_len),
m,
self.k,
self.n_out,
);
}
out_h
}
pub fn matmul_host(&self, client: &ComputeClient<R>, x: &[f32], m: usize) -> Result<Vec<f32>> {
if m == 0 || x.len() != m * self.k {
return Err(ModelError::BadShape {
tensor: "q6_k matmul input".into(),
expected: vec![m, self.k],
got: vec![x.len()],
});
}
let x_h = client.create_from_slice(f32::as_bytes(x));
let out_h = self.matmul_device(client, x_h, m);
let bytes = client.read_one_unchecked(out_h);
Ok(f32::from_bytes(&bytes).to_vec())
}
}
pub enum QuantWeight {
Q40(Q40Weight<cubecl::wgpu::WgpuRuntime>),
Q50(Q50Weight<cubecl::wgpu::WgpuRuntime>),
Q80(Q80Weight<cubecl::wgpu::WgpuRuntime>),
Q4K(Q4KWeight<cubecl::wgpu::WgpuRuntime>),
Q5K(Q5KWeight<cubecl::wgpu::WgpuRuntime>),
Q6K(Q6KWeight<cubecl::wgpu::WgpuRuntime>),
}
impl QuantWeight {
pub fn from_quant_tensor(
client: &ComputeClient<cubecl::wgpu::WgpuRuntime>,
format: combs_formats::QuantFormat,
data: &[u8],
n_out: usize,
k: usize,
) -> Result<Self> {
use combs_formats::QuantFormat;
Ok(match format {
QuantFormat::Q4_0 => QuantWeight::Q40(Q40Weight::from_gguf_bytes(client, data, n_out, k)?),
QuantFormat::Q5_0 => QuantWeight::Q50(Q50Weight::from_gguf_bytes(client, data, n_out, k)?),
QuantFormat::Q8_0 => QuantWeight::Q80(Q80Weight::from_gguf_bytes(client, data, n_out, k)?),
QuantFormat::Q4K => QuantWeight::Q4K(Q4KWeight::from_gguf_bytes(client, data, n_out, k)?),
QuantFormat::Q5K => QuantWeight::Q5K(Q5KWeight::from_gguf_bytes(client, data, n_out, k)?),
QuantFormat::Q6K => QuantWeight::Q6K(Q6KWeight::from_gguf_bytes(client, data, n_out, k)?),
})
}
pub fn n_out(&self) -> usize {
match self {
QuantWeight::Q40(w) => w.n_out,
QuantWeight::Q50(w) => w.n_out,
QuantWeight::Q80(w) => w.n_out,
QuantWeight::Q4K(w) => w.n_out,
QuantWeight::Q5K(w) => w.n_out,
QuantWeight::Q6K(w) => w.n_out,
}
}
pub fn k(&self) -> usize {
match self {
QuantWeight::Q40(w) => w.k,
QuantWeight::Q50(w) => w.k,
QuantWeight::Q80(w) => w.k,
QuantWeight::Q4K(w) => w.k,
QuantWeight::Q5K(w) => w.k,
QuantWeight::Q6K(w) => w.k,
}
}
pub fn vram_bytes(&self) -> usize {
match self {
QuantWeight::Q40(w) => w.vram_bytes(),
QuantWeight::Q50(w) => w.vram_bytes(),
QuantWeight::Q80(w) => w.vram_bytes(),
QuantWeight::Q4K(w) => w.vram_bytes(),
QuantWeight::Q5K(w) => w.vram_bytes(),
QuantWeight::Q6K(w) => w.vram_bytes(),
}
}
pub fn matmul_device(
&self,
client: &ComputeClient<cubecl::wgpu::WgpuRuntime>,
x: Handle,
m: usize,
) -> Handle {
match self {
QuantWeight::Q40(w) => w.matmul_device(client, x, m),
QuantWeight::Q50(w) => w.matmul_device(client, x, m),
QuantWeight::Q80(w) => w.matmul_device(client, x, m),
QuantWeight::Q4K(w) => w.matmul_device(client, x, m),
QuantWeight::Q5K(w) => w.matmul_device(client, x, m),
QuantWeight::Q6K(w) => w.matmul_device(client, x, m),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use cubecl::wgpu::WgpuRuntime;
fn synth_q4_0(n_blocks: usize) -> Vec<u8> {
let mut out = Vec::with_capacity(n_blocks * Q4_0_BLOCK_BYTES);
let mut s = 0x12345678u32;
for b in 0..n_blocks {
let scale = burn::tensor::f16::from_f32(0.003 * ((b % 11) as f32 + 1.0));
out.extend_from_slice(&scale.to_le_bytes());
for _ in 0..16 {
s = s.wrapping_mul(1664525).wrapping_add(1013904223);
out.push((s >> 24) as u8);
}
}
out
}
fn ref_matmul(x: &[f32], w: &[f32], m: usize, k: usize, n_out: usize) -> Vec<f32> {
let mut out = vec![0f32; m * n_out];
for r in 0..m {
for c in 0..n_out {
let mut acc = 0f32;
for i in 0..k {
acc += x[r * k + i] * w[c * k + i];
}
out[r * n_out + c] = acc;
}
}
out
}
#[test]
fn dequant_kernel_is_bit_exact_vs_cpu_reference() {
if crate::skip_no_gpu() {
return;
}
let n_blocks = 33; let data = synth_q4_0(n_blocks);
let n = n_blocks * Q4_0_BLOCK;
let expect = combs_formats::quants::dequantize_q4_0(&data, n).unwrap();
let device = Default::default();
let client = WgpuRuntime::client(&device);
let got = dequantize_q4_0_gpu::<WgpuRuntime>(&client, &data).unwrap();
assert_eq!(got, expect, "GPU dequant must be bit-exact vs gguf.rs");
}
#[test]
fn fused_matmul_matches_reference() {
if crate::skip_no_gpu() {
return;
}
let (n_out, k) = (67, 128); let n_blocks = n_out * k / Q4_0_BLOCK;
let data = synth_q4_0(n_blocks);
let w = combs_formats::quants::dequantize_q4_0(&data, n_out * k).unwrap();
let device = Default::default();
let client = WgpuRuntime::client(&device);
let weight = Q40Weight::<WgpuRuntime>::from_gguf_bytes(&client, &data, n_out, k).unwrap();
assert_eq!(weight.vram_bytes(), n_blocks * 20);
for m in [1usize, 3] {
let x: Vec<f32> = (0..m * k)
.map(|i| ((i * 7 % 13) as f32 - 6.0) / 8.0)
.collect();
let expect = ref_matmul(&x, &w, m, k, n_out);
let got = weight.matmul_host(&client, &x, m).unwrap();
assert_eq!(got.len(), expect.len());
for (i, (g, e)) in got.iter().zip(expect.iter()).enumerate() {
let tol = 1e-4 * e.abs().max(1.0);
assert!(
(g - e).abs() <= tol,
"m={m} out[{i}]: got {g}, expect {e}"
);
}
}
}
fn lcg_bytes(n: usize, seed: u32) -> Vec<u8> {
let mut s = seed;
(0..n)
.map(|_| {
s = s.wrapping_mul(1664525).wrapping_add(1013904223);
(s >> 24) as u8
})
.collect()
}
fn synth_q4_k(n_sb: usize) -> Vec<u8> {
let mut out = Vec::with_capacity(n_sb * Q4_K_BLOCK_BYTES);
for b in 0..n_sb {
let d = burn::tensor::f16::from_f32(0.002 * ((b % 9) as f32 + 1.0));
let dmin = burn::tensor::f16::from_f32(0.001 * ((b % 5) as f32 + 1.0));
out.extend_from_slice(&d.to_le_bytes());
out.extend_from_slice(&dmin.to_le_bytes());
out.extend_from_slice(&lcg_bytes(140, 0xC0FFEE ^ b as u32));
}
out
}
fn synth_q5_k(n_sb: usize) -> Vec<u8> {
let mut out = Vec::with_capacity(n_sb * Q5_K_BLOCK_BYTES);
for b in 0..n_sb {
let d = burn::tensor::f16::from_f32(0.003 * ((b % 7) as f32 + 1.0));
let dmin = burn::tensor::f16::from_f32(0.001 * ((b % 5) as f32 + 1.0));
out.extend_from_slice(&d.to_le_bytes());
out.extend_from_slice(&dmin.to_le_bytes());
out.extend_from_slice(&lcg_bytes(172, 0x5EED ^ b as u32));
}
out
}
fn synth_q6_k(n_sb: usize) -> Vec<u8> {
let mut out = Vec::with_capacity(n_sb * Q6_K_BLOCK_BYTES);
for b in 0..n_sb {
out.extend_from_slice(&lcg_bytes(208, 0xBEE5 ^ b as u32));
let d = burn::tensor::f16::from_f32(0.002 * ((b % 9) as f32 + 1.0));
out.extend_from_slice(&d.to_le_bytes());
}
out
}
fn assert_close(got: &[f32], expect: &[f32], rel: f32, what: &str) {
assert_eq!(got.len(), expect.len(), "{what}: length");
for (i, (g, e)) in got.iter().zip(expect.iter()).enumerate() {
let tol = rel * e.abs().max(1.0);
assert!((g - e).abs() <= tol, "{what}[{i}]: got {g}, expect {e}");
}
}
#[test]
fn q4_k_dequant_matches_cpu_reference() {
if crate::skip_no_gpu() {
return;
}
let n_sb = 9;
let data = synth_q4_k(n_sb);
let n = n_sb * K_SUPERBLOCK;
let expect = combs_formats::quants::dequantize_q4_k(&data, n).unwrap();
let device = Default::default();
let client = WgpuRuntime::client(&device);
let got = dequantize_q4_k_gpu::<WgpuRuntime>(&client, &data).unwrap();
assert_close(&got, &expect, 1e-6, "q4_k dequant");
}
#[test]
fn q5_k_dequant_matches_cpu_reference() {
if crate::skip_no_gpu() {
return;
}
let n_sb = 9;
let data = synth_q5_k(n_sb);
let n = n_sb * K_SUPERBLOCK;
let expect = combs_formats::quants::dequantize_q5_k(&data, n).unwrap();
let device = Default::default();
let client = WgpuRuntime::client(&device);
let got = dequantize_q5_k_gpu::<WgpuRuntime>(&client, &data).unwrap();
assert_close(&got, &expect, 1e-6, "q5_k dequant");
}
#[test]
fn q6_k_dequant_matches_cpu_reference() {
if crate::skip_no_gpu() {
return;
}
let n_sb = 9;
let data = synth_q6_k(n_sb);
let n = n_sb * K_SUPERBLOCK;
let expect = combs_formats::quants::dequantize_q6_k(&data, n).unwrap();
let device = Default::default();
let client = WgpuRuntime::client(&device);
let got = dequantize_q6_k_gpu::<WgpuRuntime>(&client, &data).unwrap();
assert_close(&got, &expect, 1e-6, "q6_k dequant");
}
#[test]
fn q5_k_fused_matmul_matches_reference() {
if crate::skip_no_gpu() {
return;
}
let (n_out, k) = (35, 512); let n_sb = n_out * k / K_SUPERBLOCK;
let data = synth_q5_k(n_sb);
let w = combs_formats::quants::dequantize_q5_k(&data, n_out * k).unwrap();
let device = Default::default();
let client = WgpuRuntime::client(&device);
let weight = Q5KWeight::<WgpuRuntime>::from_gguf_bytes(&client, &data, n_out, k).unwrap();
assert_eq!(weight.vram_bytes(), n_sb * 180);
for m in [1usize, 3] {
let x: Vec<f32> = (0..m * k)
.map(|i| ((i * 7 % 13) as f32 - 6.0) / 8.0)
.collect();
let expect = ref_matmul(&x, &w, m, k, n_out);
let got = weight.matmul_host(&client, &x, m).unwrap();
assert_close(&got, &expect, 1e-3, &format!("q5_k matmul m={m}"));
}
}
#[test]
fn q4_k_fused_matmul_matches_reference() {
if crate::skip_no_gpu() {
return;
}
let (n_out, k) = (35, 512); let n_sb = n_out * k / K_SUPERBLOCK;
let data = synth_q4_k(n_sb);
let w = combs_formats::quants::dequantize_q4_k(&data, n_out * k).unwrap();
let device = Default::default();
let client = WgpuRuntime::client(&device);
let weight = Q4KWeight::<WgpuRuntime>::from_gguf_bytes(&client, &data, n_out, k).unwrap();
assert_eq!(weight.vram_bytes(), n_sb * 148);
for m in [1usize, 3] {
let x: Vec<f32> = (0..m * k)
.map(|i| ((i * 7 % 13) as f32 - 6.0) / 8.0)
.collect();
let expect = ref_matmul(&x, &w, m, k, n_out);
let got = weight.matmul_host(&client, &x, m).unwrap();
assert_close(&got, &expect, 1e-3, &format!("q4_k matmul m={m}"));
}
}
#[test]
fn q6_k_fused_matmul_matches_reference() {
if crate::skip_no_gpu() {
return;
}
let (n_out, k) = (35, 512);
let n_sb = n_out * k / K_SUPERBLOCK;
let data = synth_q6_k(n_sb);
let w = combs_formats::quants::dequantize_q6_k(&data, n_out * k).unwrap();
let device = Default::default();
let client = WgpuRuntime::client(&device);
let weight = Q6KWeight::<WgpuRuntime>::from_gguf_bytes(&client, &data, n_out, k).unwrap();
assert_eq!(weight.vram_bytes(), n_sb * 212);
for m in [1usize, 3] {
let x: Vec<f32> = (0..m * k)
.map(|i| ((i * 7 % 13) as f32 - 6.0) / 8.0)
.collect();
let expect = ref_matmul(&x, &w, m, k, n_out);
let got = weight.matmul_host(&client, &x, m).unwrap();
assert_close(&got, &expect, 1e-3, &format!("q6_k matmul m={m}"));
}
}
fn synth_q5_0(n_blocks: usize) -> Vec<u8> {
let mut out = Vec::with_capacity(n_blocks * Q5_0_BLOCK_BYTES);
for b in 0..n_blocks {
let scale = burn::tensor::f16::from_f32(0.003 * ((b % 11) as f32 + 1.0));
out.extend_from_slice(&scale.to_le_bytes());
out.extend_from_slice(&lcg_bytes(20, 0x51D0 ^ b as u32));
}
out
}
fn synth_q8_0(n_blocks: usize) -> Vec<u8> {
let mut out = Vec::with_capacity(n_blocks * Q8_0_BLOCK_BYTES);
for b in 0..n_blocks {
let scale = burn::tensor::f16::from_f32(0.003 * ((b % 11) as f32 + 1.0));
out.extend_from_slice(&scale.to_le_bytes());
out.extend_from_slice(&lcg_bytes(32, 0x80C0 ^ b as u32));
}
out
}
#[test]
fn q5_0_and_q8_0_dequant_are_bit_exact() {
if crate::skip_no_gpu() {
return;
}
let device = Default::default();
let client = WgpuRuntime::client(&device);
let data = synth_q5_0(33);
let n = 33 * Q4_0_BLOCK;
let expect = combs_formats::quants::dequantize_q5_0(&data, n).unwrap();
let got = dequantize_q5_0_gpu::<WgpuRuntime>(&client, &data).unwrap();
assert_eq!(got, expect, "q5_0 GPU dequant must be bit-exact");
let data = synth_q8_0(33);
let expect = combs_formats::quants::dequantize_q8_0(&data, n).unwrap();
let got = dequantize_q8_0_gpu::<WgpuRuntime>(&client, &data).unwrap();
assert_eq!(got, expect, "q8_0 GPU dequant must be bit-exact");
}
#[test]
fn q5_0_and_q8_0_fused_matmul_match_reference() {
if crate::skip_no_gpu() {
return;
}
let device = Default::default();
let client = WgpuRuntime::client(&device);
let (n_out, k) = (67, 128);
let n_blocks = n_out * k / Q4_0_BLOCK;
let data5 = synth_q5_0(n_blocks);
let w5 = combs_formats::quants::dequantize_q5_0(&data5, n_out * k).unwrap();
let q5 = Q50Weight::<WgpuRuntime>::from_gguf_bytes(&client, &data5, n_out, k).unwrap();
assert_eq!(q5.vram_bytes(), n_blocks * 24);
let data8 = synth_q8_0(n_blocks);
let w8 = combs_formats::quants::dequantize_q8_0(&data8, n_out * k).unwrap();
let q8 = Q80Weight::<WgpuRuntime>::from_gguf_bytes(&client, &data8, n_out, k).unwrap();
assert_eq!(q8.vram_bytes(), n_blocks * 36);
for m in [1usize, 3] {
let x: Vec<f32> = (0..m * k)
.map(|i| ((i * 7 % 13) as f32 - 6.0) / 8.0)
.collect();
let got5 = q5.matmul_host(&client, &x, m).unwrap();
assert_close(&got5, &ref_matmul(&x, &w5, m, k, n_out), 1e-3, &format!("q5_0 m={m}"));
let got8 = q8.matmul_host(&client, &x, m).unwrap();
assert_close(&got8, &ref_matmul(&x, &w8, m, k, n_out), 1e-3, &format!("q8_0 m={m}"));
}
}
#[test]
fn q5_q8_shape_validation() {
if crate::skip_no_gpu() {
return;
}
let device = Default::default();
let client = WgpuRuntime::client(&device);
assert!(repack_q5_0(&[0u8; 21]).is_err());
assert!(repack_q8_0(&[0u8; 33]).is_err());
assert!(Q50Weight::<WgpuRuntime>::from_gguf_bytes(&client, &synth_q5_0(2), 2, 31).is_err());
assert!(Q80Weight::<WgpuRuntime>::from_gguf_bytes(&client, &synth_q8_0(2), 2, 64).is_err());
}
#[test]
fn k_quant_shape_validation() {
if crate::skip_no_gpu() {
return;
}
let device = Default::default();
let client = WgpuRuntime::client(&device);
assert!(repack_q4_k(&[0u8; 143]).is_err());
assert!(repack_q6_k(&[0u8; 209]).is_err());
assert!(
Q4KWeight::<WgpuRuntime>::from_gguf_bytes(&client, &synth_q4_k(1), 1, 128).is_err()
);
assert!(
Q6KWeight::<WgpuRuntime>::from_gguf_bytes(&client, &synth_q6_k(1), 1, 128).is_err()
);
}
fn assert_tiled_bit_identical(
untiled: &[f32],
tiled: &[f32],
label: &str,
) {
assert_eq!(untiled.len(), tiled.len(), "{label}: length mismatch");
for (i, (u, t)) in untiled.iter().zip(tiled.iter()).enumerate() {
assert_eq!(
u.to_bits(),
t.to_bits(),
"{label} out[{i}]: untiled {u} vs tiled {t} — accumulation order drifted"
);
}
}
#[test]
fn tiled_q8_0_matmul_is_bit_identical() {
if crate::skip_no_gpu() {
return;
}
let (n_out, k) = (300, 320);
let n_blocks = n_out * k / Q4_0_BLOCK;
let data = synth_q8_0(n_blocks);
let device = Default::default();
let client = WgpuRuntime::client(&device);
let w = Q80Weight::<WgpuRuntime>::from_gguf_bytes(&client, &data, n_out, k).unwrap();
for m in [2usize, 3, 17, 256, 1024] {
let x: Vec<f32> = (0..m * k)
.map(|i| ((i * 11 % 29) as f32 - 14.0) / 16.0)
.collect();
let x_h = client.create_from_slice(f32::as_bytes(&x));
let un_h = w.matmul_device_with(&client, x_h.clone(), m, false);
let ti_h = w.matmul_device_with(&client, x_h, m, true);
let un = f32::from_bytes(&client.read_one_unchecked(un_h)).to_vec();
let ti = f32::from_bytes(&client.read_one_unchecked(ti_h)).to_vec();
assert_tiled_bit_identical(&un, &ti, &format!("q8_0 m={m}"));
}
}
#[test]
fn tiled_q4_k_matmul_is_bit_identical() {
if crate::skip_no_gpu() {
return;
}
let (n_out, k) = (300, 512);
let n_sb = n_out * k / K_SUPERBLOCK;
let data = synth_q4_k(n_sb);
let device = Default::default();
let client = WgpuRuntime::client(&device);
let w = Q4KWeight::<WgpuRuntime>::from_gguf_bytes(&client, &data, n_out, k).unwrap();
for m in [2usize, 3, 17, 256, 1024] {
let x: Vec<f32> = (0..m * k)
.map(|i| ((i * 13 % 31) as f32 - 15.0) / 16.0)
.collect();
let x_h = client.create_from_slice(f32::as_bytes(&x));
let un_h = w.matmul_device_with(&client, x_h.clone(), m, false);
let ti_h = w.matmul_device_with(&client, x_h, m, true);
let un = f32::from_bytes(&client.read_one_unchecked(un_h)).to_vec();
let ti = f32::from_bytes(&client.read_one_unchecked(ti_h)).to_vec();
assert_tiled_bit_identical(&un, &ti, &format!("q4_k m={m}"));
}
}
#[test]
fn shape_validation() {
if crate::skip_no_gpu() {
return;
}
let device = Default::default();
let client = WgpuRuntime::client(&device);
assert!(repack_q4_0(&[0u8; 17]).is_err());
assert!(Q40Weight::<WgpuRuntime>::from_gguf_bytes(&client, &synth_q4_0(2), 2, 31).is_err());
assert!(Q40Weight::<WgpuRuntime>::from_gguf_bytes(&client, &synth_q4_0(2), 2, 64).is_err());
let w = Q40Weight::<WgpuRuntime>::from_gguf_bytes(&client, &synth_q4_0(2), 2, 32).unwrap();
assert!(w.matmul_host(&client, &[0f32; 31], 1).is_err());
}
}