use crate::prelude::*;
use crate::tune::{Config, Evaluator, Evolution, Evolved, Space, Tuned, Tuner, Verdict};
use cubecl::server::Handle;
pub const QK8_0: usize = 32;
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn matvec_q8<F: Float>(
wd: &Array<F>,
wq: &Array<i32>,
x: &Array<F>,
out: &mut Array<F>,
#[comptime] k: usize,
) {
let row = ABSOLUTE_POS;
if row < out.len() {
let nb = k / 32;
let wbase = row * k;
let dbase = row * nb;
let mut acc = F::new(0.0);
for i in 0..k {
let d = wd[dbase + i / 32];
let q = F::cast_from(wq[wbase + i]);
acc += d * q * x[i];
}
out[row] = acc;
}
}
pub fn matvec_q8_run<R: Runtime>(
client: &ComputeClient<R>,
wd: &[f32],
wq: &[i32],
x: &[f32],
rows: usize,
k: usize,
) -> Vec<f32> {
let wdh = client.create_from_slice(f32::as_bytes(wd));
let wqh = client.create_from_slice(i32::as_bytes(wq));
let xh = client.create_from_slice(f32::as_bytes(x));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; rows]));
let block = 64u32;
let grid = (rows as u32).div_ceil(block);
unsafe {
matvec_q8::launch_unchecked::<f32, R>(
client,
Grid::Static(grid, 1, 1),
Block::new_1d(block),
ArrayArg::from_raw_parts(wdh.clone(), wd.len()),
ArrayArg::from_raw_parts(wqh.clone(), wq.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(oh.clone(), rows),
k,
);
}
let bytes = client.read_one_unchecked(oh);
f32::from_bytes(&bytes).to_vec()
}
pub fn matvec_q8_bench<R: Runtime>(
client: &ComputeClient<R>,
wd: &[f32],
wq: &[i32],
x: &[f32],
rows: usize,
k: usize,
iters: usize,
) -> f64 {
let wdh = client.create_from_slice(f32::as_bytes(wd));
let wqh = client.create_from_slice(i32::as_bytes(wq));
let xh = client.create_from_slice(f32::as_bytes(x));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; rows]));
let block = 64u32;
let grid = (rows as u32).div_ceil(block);
let launch = |c: &ComputeClient<R>| unsafe {
matvec_q8::launch_unchecked::<f32, R>(
c,
Grid::Static(grid, 1, 1),
Block::new_1d(block),
ArrayArg::from_raw_parts(wdh.clone(), wd.len()),
ArrayArg::from_raw_parts(wqh.clone(), wq.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(oh.clone(), rows),
k,
);
};
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone()); let t = std::time::Instant::now();
for _ in 0..iters {
launch(client);
}
let _ = client.read_one_unchecked(oh); t.elapsed().as_secs_f64() * 1e3 / iters as f64
}
pub fn matvec_q8_ref(wd: &[f32], wq: &[i32], x: &[f32], rows: usize, k: usize) -> Vec<f32> {
let nb = k / 32;
(0..rows)
.map(|row| {
let mut acc = 0.0f32;
for i in 0..k {
acc += wd[row * nb + i / 32] * wq[row * k + i] as f32 * x[i];
}
acc
})
.collect()
}
pub const QK_K: usize = 256;
#[device]
fn byte_at(a: &Array<u32>, base: usize, i: usize) -> u32 {
(a[base + i / 4] >> ((8 * (i % 4)) as u32)) & 255
}
#[device]
fn q4k_sc(wsc: &Array<u32>, scbase: usize, j: usize) -> u32 {
let mut r = byte_at(wsc, scbase, j) & 63;
if j >= 4 {
r = (byte_at(wsc, scbase, j + 4) & 15) | ((byte_at(wsc, scbase, j - 4) >> 6) << 4);
}
r
}
#[device]
fn q4k_m(wsc: &Array<u32>, scbase: usize, j: usize) -> u32 {
let mut r = byte_at(wsc, scbase, j + 4) & 63;
if j >= 4 {
r = (byte_at(wsc, scbase, j + 4) >> 4) | ((byte_at(wsc, scbase, j) >> 6) << 4);
}
r
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn matvec_q4k<F: Float>(
wqs: &Array<u32>,
wsc: &Array<u32>,
wd: &Array<F>,
wdm: &Array<F>,
x: &Array<F>,
out: &mut Array<F>,
#[comptime] k: usize,
) {
let row = ABSOLUTE_POS;
if row < out.len() {
let nb = k / 256;
let mut acc = F::new(0.0);
for b in 0..nb {
let blk = row * nb + b;
let qbase = blk * 32;
let scbase = blk * 3;
let d = wd[blk];
let dmin = wdm[blk];
let xbase = b * 256;
for g in 0..4 {
let is = g * 2;
let d1 = d * F::cast_from(q4k_sc(wsc, scbase, is));
let mm1 = dmin * F::cast_from(q4k_m(wsc, scbase, is));
let d2 = d * F::cast_from(q4k_sc(wsc, scbase, is + 1));
let mm2 = dmin * F::cast_from(q4k_m(wsc, scbase, is + 1));
for qi in 0..32 {
let qb = byte_at(wqs, qbase, g * 32 + qi);
let wlo = d1 * F::cast_from(qb & 15) - mm1;
acc += wlo * x[xbase + g * 64 + qi];
let whi = d2 * F::cast_from(qb >> 4) - mm2;
acc += whi * x[xbase + g * 64 + 32 + qi];
}
}
}
out[row] = acc;
}
}
pub fn matvec_q4k_run<R: Runtime>(
client: &ComputeClient<R>,
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
x: &[f32],
rows: usize,
k: usize,
) -> Vec<f32> {
let qh = client.create_from_slice(u32::as_bytes(wqs));
let sh = client.create_from_slice(u32::as_bytes(wsc));
let dh = client.create_from_slice(f32::as_bytes(wd));
let mh = client.create_from_slice(f32::as_bytes(wdm));
let xh = client.create_from_slice(f32::as_bytes(x));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; rows]));
let block = 64u32;
let grid = (rows as u32).div_ceil(block);
unsafe {
matvec_q4k::launch_unchecked::<f32, R>(
client,
Grid::Static(grid, 1, 1),
Block::new_1d(block),
ArrayArg::from_raw_parts(qh.clone(), wqs.len()),
ArrayArg::from_raw_parts(sh.clone(), wsc.len()),
ArrayArg::from_raw_parts(dh.clone(), wd.len()),
ArrayArg::from_raw_parts(mh.clone(), wdm.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(oh.clone(), rows),
k,
);
}
f32::from_bytes(&client.read_one_unchecked(oh)).to_vec()
}
pub fn matvec_q4k_bench<R: Runtime>(
client: &ComputeClient<R>,
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
x: &[f32],
rows: usize,
k: usize,
iters: usize,
) -> f64 {
let qh = client.create_from_slice(u32::as_bytes(wqs));
let sh = client.create_from_slice(u32::as_bytes(wsc));
let dh = client.create_from_slice(f32::as_bytes(wd));
let mh = client.create_from_slice(f32::as_bytes(wdm));
let xh = client.create_from_slice(f32::as_bytes(x));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; rows]));
let block = 64u32;
let grid = (rows as u32).div_ceil(block);
let launch = |c: &ComputeClient<R>| unsafe {
matvec_q4k::launch_unchecked::<f32, R>(
c,
Grid::Static(grid, 1, 1),
Block::new_1d(block),
ArrayArg::from_raw_parts(qh.clone(), wqs.len()),
ArrayArg::from_raw_parts(sh.clone(), wsc.len()),
ArrayArg::from_raw_parts(dh.clone(), wd.len()),
ArrayArg::from_raw_parts(mh.clone(), wdm.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(oh.clone(), rows),
k,
);
};
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t = std::time::Instant::now();
for _ in 0..iters {
launch(client);
}
let _ = client.read_one_unchecked(oh);
t.elapsed().as_secs_f64() * 1e3 / iters as f64
}
#[inline]
fn cpu_byte(a: &[u32], base: usize, i: usize) -> u32 {
(a[base + i / 4] >> (8 * (i % 4))) & 255
}
#[inline]
fn cpu_sc(wsc: &[u32], scbase: usize, j: usize) -> u32 {
if j < 4 {
cpu_byte(wsc, scbase, j) & 63
} else {
(cpu_byte(wsc, scbase, j + 4) & 15) | ((cpu_byte(wsc, scbase, j - 4) >> 6) << 4)
}
}
#[inline]
fn cpu_m(wsc: &[u32], scbase: usize, j: usize) -> u32 {
if j < 4 {
cpu_byte(wsc, scbase, j + 4) & 63
} else {
(cpu_byte(wsc, scbase, j + 4) >> 4) | ((cpu_byte(wsc, scbase, j) >> 6) << 4)
}
}
pub fn matvec_q4k_ref(
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
x: &[f32],
rows: usize,
k: usize,
) -> Vec<f32> {
let nb = k / 256;
(0..rows)
.map(|row| {
let mut acc = 0.0f32;
for b in 0..nb {
let blk = row * nb + b;
let (qbase, scbase) = (blk * 32, blk * 3);
let (d, dmin) = (wd[blk], wdm[blk]);
let xbase = b * 256;
for g in 0..4 {
let is = g * 2;
let d1 = d * cpu_sc(wsc, scbase, is) as f32;
let mm1 = dmin * cpu_m(wsc, scbase, is) as f32;
let d2 = d * cpu_sc(wsc, scbase, is + 1) as f32;
let mm2 = dmin * cpu_m(wsc, scbase, is + 1) as f32;
for qi in 0..32 {
let qb = cpu_byte(wqs, qbase, g * 32 + qi);
acc += (d1 * (qb & 15) as f32 - mm1) * x[xbase + g * 64 + qi];
acc += (d2 * (qb >> 4) as f32 - mm2) * x[xbase + g * 64 + 32 + qi];
}
}
}
acc
})
.collect()
}
pub fn gen_q4k(rows: usize, k: usize) -> (Vec<u32>, Vec<u32>, Vec<f32>, Vec<f32>, Vec<f32>) {
let nb = k / 256;
let nblk = rows * nb;
let mut s = 0x9E3779B97F4A7C15u64;
let mut next = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
let wqs: Vec<u32> = (0..nblk * 32).map(|_| next() as u32).collect(); let wsc: Vec<u32> = (0..nblk * 3).map(|_| next() as u32).collect(); let wd: Vec<f32> = (0..nblk)
.map(|_| half::f16::from_f32((next() % 1000) as f32 / 20000.0 + 0.002).to_f32())
.collect();
let wdm: Vec<f32> = (0..nblk)
.map(|_| half::f16::from_f32((next() % 1000) as f32 / 40000.0).to_f32())
.collect();
let x: Vec<f32> = (0..k)
.map(|_| (next() % 2000) as f32 / 1000.0 - 1.0)
.collect();
(wqs, wsc, wd, wdm, x)
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn moe_matvec_q4k<F: Float>(
wqs: &Array<u32>,
wsc: &Array<u32>,
wd: &Array<F>,
wdm: &Array<F>,
x: &Array<F>,
ids: &Array<u32>,
out: &mut Array<F>,
#[comptime] n: usize,
#[comptime] k: usize,
) {
let gid = ABSOLUTE_POS;
if gid < out.len() {
let slot = gid / n;
let r = gid % n;
let wrow = ids[slot] as usize * n + r; let nb = k / 256;
let xbase = slot * k;
let mut acc = F::new(0.0);
for b in 0..nb {
let blk = wrow * nb + b;
let qbase = blk * 32;
let scbase = blk * 3;
let d = wd[blk];
let dmin = wdm[blk];
let xblk = xbase + b * 256;
for g in 0..4 {
let is = g * 2;
let d1 = d * F::cast_from(q4k_sc(wsc, scbase, is));
let mm1 = dmin * F::cast_from(q4k_m(wsc, scbase, is));
let d2 = d * F::cast_from(q4k_sc(wsc, scbase, is + 1));
let mm2 = dmin * F::cast_from(q4k_m(wsc, scbase, is + 1));
for qi in 0..32 {
let qb = byte_at(wqs, qbase, g * 32 + qi);
let wlo = d1 * F::cast_from(qb & 15) - mm1;
acc += wlo * x[xblk + g * 64 + qi];
let whi = d2 * F::cast_from(qb >> 4) - mm2;
acc += whi * x[xblk + g * 64 + 32 + qi];
}
}
}
out[gid] = acc;
}
}
#[allow(clippy::too_many_arguments)]
pub fn moe_matvec_q4k_run<R: Runtime>(
client: &ComputeClient<R>,
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
x: &[f32],
ids: &[u32],
slots: usize,
n: usize,
k: usize,
) -> Vec<f32> {
let qh = client.create_from_slice(u32::as_bytes(wqs));
let sh = client.create_from_slice(u32::as_bytes(wsc));
let dh = client.create_from_slice(f32::as_bytes(wd));
let mh = client.create_from_slice(f32::as_bytes(wdm));
let xh = client.create_from_slice(f32::as_bytes(x));
let ih = client.create_from_slice(u32::as_bytes(ids));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; slots * n]));
let block = 64u32;
let grid = ((slots * n) as u32).div_ceil(block);
unsafe {
moe_matvec_q4k::launch_unchecked::<f32, R>(
client,
Grid::Static(grid, 1, 1),
Block::new_1d(block),
ArrayArg::from_raw_parts(qh.clone(), wqs.len()),
ArrayArg::from_raw_parts(sh.clone(), wsc.len()),
ArrayArg::from_raw_parts(dh.clone(), wd.len()),
ArrayArg::from_raw_parts(mh.clone(), wdm.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(ih.clone(), ids.len()),
ArrayArg::from_raw_parts(oh.clone(), slots * n),
n,
k,
);
}
f32::from_bytes(&client.read_one_unchecked(oh)).to_vec()
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn moe_matvec_q4k_blk<F: Float>(
wqs: &Array<u32>,
wsc: &Array<u32>,
wd: &Array<F>,
wdm: &Array<F>,
x: &Array<F>,
ids: &Array<u32>,
out: &mut Array<F>,
#[comptime] n: usize, #[comptime] k: usize, #[comptime] nt: usize, ) {
let outrow = CUBE_POS as usize;
let t = UNIT_POS as usize;
let slot = outrow / n;
let r = outrow % n;
let wrow = ids[slot] as usize * n + r;
let nb = k / 256;
let nsub = k / 32; let per = (nsub + nt - 1) / nt; let xrow = slot * k;
let mut partial = F::new(0.0);
for j in 0..per {
let sb = j * nt + t;
if sb < nsub {
let sup = sb / 8;
let jloc = sb % 8;
let blk = wrow * nb + sup;
let qbase = blk * 32;
let scbase = blk * 3;
let dj = wd[blk] * F::cast_from(q4k_sc(wsc, scbase, jloc));
let mj = wdm[blk] * F::cast_from(q4k_m(wsc, scbase, jloc));
let shift = ((jloc % 2) * 4) as u32; let boff = (jloc / 2) * 32; let xoff = xrow + sup * 256 + jloc * 32;
let mut sbsum = F::new(0.0);
for qi in 0..32 {
let nib = (byte_at(wqs, qbase, boff + qi) >> shift) & 15;
sbsum += (dj * F::cast_from(nib) - mj) * x[xoff + qi];
}
partial += sbsum;
}
}
let mut smem = SharedMemory::<F>::new(nt);
smem[t] = partial;
sync_cube();
let mut stride = CUBE_DIM / 2;
while stride > 0 {
if UNIT_POS < stride {
let v = smem[(UNIT_POS + stride) as usize];
smem[t] += v;
}
sync_cube();
stride /= 2;
}
if t == 0 {
out[outrow] = smem[0];
}
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
#[allow(clippy::too_many_arguments)]
pub fn moe_matvec_q4k_dp4a_blk<F: Float>(
wqs: &Array<u32>,
wsc: &Array<u32>,
wd: &Array<F>,
wdm: &Array<F>,
xq: &Array<u32>, xs: &Array<F>, xsum: &Array<F>, ids: &Array<u32>,
out: &mut Array<F>,
#[comptime] n: usize,
#[comptime] k: usize,
#[comptime] nt: usize,
) {
let outrow = CUBE_POS as usize;
let t = UNIT_POS as usize;
let slot = outrow / n;
let r = outrow % n;
let wrow = ids[slot] as usize * n + r;
let nb = k / 256;
let nsub = k / 32;
let per = (nsub + nt - 1) / nt;
let mut partial = F::new(0.0);
for j in 0..per {
let sb = j * nt + t;
if sb < nsub {
let sup = sb / 8;
let jloc = sb % 8;
let blk = wrow * nb + sup;
let qbase = blk * 32; let scbase = blk * 3;
let dj = wd[blk] * F::cast_from(q4k_sc(wsc, scbase, jloc));
let mj = wdm[blk] * F::cast_from(q4k_m(wsc, scbase, jloc));
let shift = ((jloc % 2) * 4) as u32; let cw = qbase + (jloc / 2) * 8; let xg = (slot * nsub + sb) * 8; let mut idot = 0i32;
#[unroll]
for g in 0..8usize {
let nibs = (wqs[cw + g] >> shift) & 0x0F0F0F0F; let wv = Vector::<i32, Const<4>>::cast_from(Vector::<i8, Const<4>>::reinterpret::<
u32,
>(nibs));
let xv = Vector::<i32, Const<4>>::cast_from(Vector::<i8, Const<4>>::reinterpret::<
u32,
>(xq[xg + g]));
idot += wv.dot(xv); }
let xi = slot * nsub + sb;
partial += dj * xs[xi] * F::cast_from(idot) - mj * xsum[xi];
}
}
let mut smem = SharedMemory::<F>::new(nt);
smem[t] = partial;
sync_cube();
let mut stride = CUBE_DIM / 2;
while stride > 0 {
if UNIT_POS < stride {
let v = smem[(UNIT_POS + stride) as usize];
smem[t] += v;
}
sync_cube();
stride /= 2;
}
if t == 0 {
out[outrow] = smem[0];
}
}
pub fn quant_act_q8_cpu(x: &[f32], slots: usize, k: usize) -> (Vec<u32>, Vec<f32>, Vec<f32>) {
let kb = k / 32;
let mut xq = vec![0u32; slots * kb * 8];
let mut xs = vec![0f32; slots * kb];
let mut xsum = vec![0f32; slots * kb];
for s in 0..slots {
for b in 0..kb {
let base = s * k + b * 32;
let amax = (0..32).fold(0f32, |m, i| m.max(x[base + i].abs()));
let inv = if amax > 0.0 { 127.0 / amax } else { 0.0 };
let scale = if amax > 0.0 { amax / 127.0 } else { 1.0 };
let gid = s * kb + b;
let mut isum = 0i32;
for j in 0..8 {
let mut word = 0u32;
for l in 0..4 {
let q = (x[base + j * 4 + l] * inv).round().clamp(-127.0, 127.0) as i32;
isum += q;
word |= ((q as u32) & 0xFF) << (l * 8);
}
xq[gid * 8 + j] = word;
}
xs[gid] = scale;
xsum[gid] = scale * isum as f32;
}
}
(xq, xs, xsum)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_matvec_q4k_dp4a_blk_run<R: Runtime>(
client: &ComputeClient<R>,
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
x: &[f32],
ids: &[u32],
slots: usize,
n: usize,
k: usize,
nt: usize,
) -> Vec<f32> {
let (xq, xs, xsum) = quant_act_q8_cpu(x, slots, k);
let qh = client.create_from_slice(u32::as_bytes(wqs));
let sh = client.create_from_slice(u32::as_bytes(wsc));
let dh = client.create_from_slice(f32::as_bytes(wd));
let mh = client.create_from_slice(f32::as_bytes(wdm));
let xqh = client.create_from_slice(u32::as_bytes(&xq));
let xsh = client.create_from_slice(f32::as_bytes(&xs));
let xsumh = client.create_from_slice(f32::as_bytes(&xsum));
let ih = client.create_from_slice(u32::as_bytes(ids));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; slots * n]));
unsafe {
moe_matvec_q4k_dp4a_blk::launch_unchecked::<f32, R>(
client,
Grid::Static((slots * n) as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(qh.clone(), wqs.len()),
ArrayArg::from_raw_parts(sh.clone(), wsc.len()),
ArrayArg::from_raw_parts(dh.clone(), wd.len()),
ArrayArg::from_raw_parts(mh.clone(), wdm.len()),
ArrayArg::from_raw_parts(xqh.clone(), xq.len()),
ArrayArg::from_raw_parts(xsh.clone(), xs.len()),
ArrayArg::from_raw_parts(xsumh.clone(), xsum.len()),
ArrayArg::from_raw_parts(ih.clone(), ids.len()),
ArrayArg::from_raw_parts(oh.clone(), slots * n),
n,
k,
nt,
);
}
f32::from_bytes(&client.read_one_unchecked(oh)).to_vec()
}
pub fn pack_q4k(wqs: &[u32], wsc: &[u32], wd: &[f32], wdm: &[f32]) -> Vec<u32> {
let nblk = wd.len();
let mut p = vec![0u32; nblk * 36];
for b in 0..nblk {
let d = half::f16::from_f32(wd[b]).to_bits() as u32;
let dm = half::f16::from_f32(wdm[b]).to_bits() as u32;
p[b * 36] = d | (dm << 16);
p[b * 36 + 1..b * 36 + 4].copy_from_slice(&wsc[b * 3..b * 3 + 3]);
p[b * 36 + 4..b * 36 + 36].copy_from_slice(&wqs[b * 32..b * 32 + 32]);
}
p
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn matvec_q4k_dp4a_blk<F: Float>(
wq: &Array<u32>, xq: &Array<u32>, xs: &Array<F>, xsum: &Array<F>, out: &mut Array<F>,
meta: &Array<u32>, #[comptime] nt: usize,
) {
let row = CUBE_POS as usize;
let t = UNIT_POS as usize;
let k = meta[0] as usize;
let nb = k / 256;
let nsub = k / 32;
let per = (nsub + nt - 1) / nt;
let mut partial = F::new(0.0);
for j in 0..per {
let sb = j * nt + t;
if sb < nsub {
let sup = sb / 8;
let jloc = sb % 8;
let bb = (row * nb + sup) * 36; let d = F::cast_from(f16lo_to_f32(wq[bb]));
let dmin = F::cast_from(f16lo_to_f32(wq[bb] >> 16));
let scbase = bb + 1; let dj = d * F::cast_from(q4k_sc(wq, scbase, jloc));
let mj = dmin * F::cast_from(q4k_m(wq, scbase, jloc));
let shift = ((jloc % 2) * 4) as u32;
let cw = bb + 4 + (jloc / 2) * 8; let xg = sb * 8;
let mut idot = 0i32;
#[unroll]
for g in 0..8usize {
let nibs = (wq[cw + g] >> shift) & 0x0F0F0F0F;
let wv = Vector::<i32, Const<4>>::cast_from(Vector::<i8, Const<4>>::reinterpret::<
u32,
>(nibs));
let xv = Vector::<i32, Const<4>>::cast_from(Vector::<i8, Const<4>>::reinterpret::<
u32,
>(xq[xg + g]));
idot += wv.dot(xv);
}
partial += dj * xs[sb] * F::cast_from(idot) - mj * xsum[sb];
}
}
let mut smem = SharedMemory::<F>::new(nt);
smem[t] = partial;
sync_cube();
let mut stride = CUBE_DIM / 2;
while stride > 0 {
if UNIT_POS < stride {
let v = smem[(UNIT_POS + stride) as usize];
smem[t] += v;
}
sync_cube();
stride /= 2;
}
if t == 0 {
out[row] = smem[0];
}
}
pub fn matvec_q4k_dp4a_blk_run<R: Runtime>(
client: &ComputeClient<R>,
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
x: &[f32],
nout: usize,
k: usize,
nt: usize,
) -> Vec<f32> {
let packed = pack_q4k(wqs, wsc, wd, wdm);
let (xq, xs, xsum) = quant_act_q8_cpu(x, 1, k);
let wh = client.create_from_slice(u32::as_bytes(&packed));
let xqh = client.create_from_slice(u32::as_bytes(&xq));
let xsh = client.create_from_slice(f32::as_bytes(&xs));
let xsumh = client.create_from_slice(f32::as_bytes(&xsum));
let meta = client.create_from_slice(u32::as_bytes(&[k as u32]));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; nout]));
unsafe {
matvec_q4k_dp4a_blk::launch_unchecked::<f32, R>(
client,
Grid::Static(nout as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(wh.clone(), packed.len()),
ArrayArg::from_raw_parts(xqh.clone(), xq.len()),
ArrayArg::from_raw_parts(xsh.clone(), xs.len()),
ArrayArg::from_raw_parts(xsumh.clone(), xsum.len()),
ArrayArg::from_raw_parts(oh.clone(), nout),
ArrayArg::from_raw_parts(meta.clone(), 1),
nt,
);
}
f32::from_bytes(&client.read_one_unchecked(oh)).to_vec()
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn matvec_q4k_f32_blk<F: Float>(
wq: &Array<u32>, x: &Array<F>, out: &mut Array<F>,
meta: &Array<u32>, #[comptime] nt: usize, #[comptime] nr: usize, #[comptime] tgt: Target, ) {
let t = UNIT_POS as usize;
let k = meta[0] as usize;
let nblocks = k / 256;
let nout = meta[1] as usize; let row0 = CUBE_POS as usize * nr;
island! {
vulkan => {
let its = nt / 16;
let itid = t & 15;
let ix = t / 16;
let hg = itid / 8; let qq = (itid & 7) * 4; let qw = itid; let j0 = 2 * hg;
let j1 = j0 + 1;
let j2 = 4 + 2 * hg;
let j3 = j2 + 1;
let ya = hg * 64 + qq;
let yb = ya + 32;
let yc = 128 + hg * 64 + qq;
let yd = yc + 32;
let mut acc = Array::<F>::new(nr);
#[unroll]
for n in 0..nr {
acc[n] = F::new(0.0);
}
#[unroll]
for n in 0..nr {
let row = row0 + n;
if row < nout {
let rowbase = row * nblocks * 36;
let mut blk = ix;
while blk < nblocks {
let b = rowbase + blk * 36;
let d = F::cast_from(f16lo_to_f32(wq[b]));
let dmin = F::cast_from(f16lo_to_f32(wq[b] >> 16));
let wlo = wq[b + 4 + qw]; let whi = wq[b + 4 + qw + 16]; let dj0 = d * F::cast_from(q4k_sc(wq, b + 1, j0));
let mj0 = dmin * F::cast_from(q4k_m(wq, b + 1, j0));
let dj1 = d * F::cast_from(q4k_sc(wq, b + 1, j1));
let mj1 = dmin * F::cast_from(q4k_m(wq, b + 1, j1));
let dj2 = d * F::cast_from(q4k_sc(wq, b + 1, j2));
let mj2 = dmin * F::cast_from(q4k_m(wq, b + 1, j2));
let dj3 = d * F::cast_from(q4k_sc(wq, b + 1, j3));
let mj3 = dmin * F::cast_from(q4k_m(wq, b + 1, j3));
let yblk = blk * 256;
#[unroll]
for l in 0..4usize {
let s8 = (l * 8) as u32;
acc[n] += (dj0 * F::cast_from((wlo >> s8) & 0x0F) - mj0) * x[yblk + ya + l];
acc[n] += (dj1 * F::cast_from((wlo >> (s8 + 4)) & 0x0F) - mj1) * x[yblk + yb + l];
acc[n] += (dj2 * F::cast_from((whi >> s8) & 0x0F) - mj2) * x[yblk + yc + l];
acc[n] += (dj3 * F::cast_from((whi >> (s8 + 4)) & 0x0F) - mj3) * x[yblk + yd + l];
}
blk += its;
}
}
}
#[unroll]
for n in 0..nr {
let total = plane_sum(acc[n]);
if t == 0 {
let row = row0 + n;
if row < nout {
out[row] = total;
}
}
}
}
default => {
let nsub = k / 32;
let per = (nsub + nt - 1) / nt;
let mut acc = Array::<F>::new(nr);
#[unroll]
for n in 0..nr {
acc[n] = F::new(0.0);
}
for j in 0..per {
let sb = j * nt + t;
if sb < nsub {
let sup = sb / 8;
let jloc = sb % 8;
let shift = ((jloc % 2) * 4) as u32;
let cwoff = 4 + (jloc / 2) * 8;
let xbase = sup * 256 + (jloc / 2) * 64 + (jloc % 2) * 32;
#[unroll]
for n in 0..nr {
let row = row0 + n;
if row < nout {
let bb = (row * nblocks + sup) * 36;
let d = F::cast_from(f16lo_to_f32(wq[bb]));
let dmin = F::cast_from(f16lo_to_f32(wq[bb] >> 16));
let dj = d * F::cast_from(q4k_sc(wq, bb + 1, jloc));
let mj = dmin * F::cast_from(q4k_m(wq, bb + 1, jloc));
let cw = bb + cwoff;
let mut sdot = F::new(0.0);
let mut sx = F::new(0.0);
#[unroll]
for g in 0..8usize {
let word = (wq[cw + g] >> shift) & 0x0F0F0F0F;
let xb = xbase + g * 4;
let x0 = x[xb];
let x1 = x[xb + 1];
let x2 = x[xb + 2];
let x3 = x[xb + 3];
sdot += F::cast_from(word & 0xFF) * x0
+ F::cast_from((word >> 8) & 0xFF) * x1
+ F::cast_from((word >> 16) & 0xFF) * x2
+ F::cast_from((word >> 24) & 0xFF) * x3;
sx += x0 + x1 + x2 + x3;
}
acc[n] += dj * sdot - mj * sx;
}
}
}
}
let mut smem = SharedMemory::<F>::new(nt);
#[unroll]
for n in 0..nr {
smem[t] = acc[n];
sync_cube();
let mut stride = CUBE_DIM / 2;
while stride > 0 {
if UNIT_POS < stride {
let v = smem[(UNIT_POS + stride) as usize];
smem[t] += v;
}
sync_cube();
stride /= 2;
}
if t == 0 {
let row = row0 + n;
if row < nout {
out[row] = smem[0];
}
}
sync_cube();
}
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn matvec_q4k_f32_blk_run<R: Runtime>(
client: &ComputeClient<R>,
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
x: &[f32],
nout: usize,
k: usize,
nt: usize,
nr: usize,
) -> Vec<f32> {
let packed = pack_q4k(wqs, wsc, wd, wdm);
let wh = client.create_from_slice(u32::as_bytes(&packed));
let xh = client.create_from_slice(f32::as_bytes(x));
let meta = client.create_from_slice(u32::as_bytes(&[k as u32, nout as u32]));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; nout]));
let grid = (nout as u32).div_ceil(nr as u32); unsafe {
matvec_q4k_f32_blk::launch_unchecked::<f32, R>(
client,
Grid::Static(grid, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(wh.clone(), packed.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(oh.clone(), nout),
ArrayArg::from_raw_parts(meta.clone(), 2),
nt,
nr,
Target::of(client),
);
}
f32::from_bytes(&client.read_one_unchecked(oh)).to_vec()
}
pub fn matvec_q4k_f32_space() -> Space {
Space::new()
.param("WG", [16, 32, 64, 128, 256])
.param("NR", [1, 2, 4])
.constraint(|c, s| {
let w = c.get(s, "WG");
w >= 2 && (w & (w - 1)) == 0 && w <= 1024
})
}
fn matvec_q4k_f32_cfg(c: &Config, s: &Space) -> (usize, usize) {
(c.get(s, "WG") as usize, c.get(s, "NR") as usize)
}
pub struct MatvecQ4kF32Eval<'a, R: Runtime> {
client: &'a ComputeClient<R>,
space: &'a Space,
banks: Vec<Handle>, xh: Handle,
meta: Handle,
outh: Handle,
packed_len: usize,
x_len: usize,
rows: usize,
k: usize,
oracle: Vec<f32>,
maxref: f32,
repeats: usize,
worst_rel: std::cell::Cell<f32>,
}
impl<'a, R: Runtime> MatvecQ4kF32Eval<'a, R> {
#[allow(clippy::too_many_arguments)]
pub fn new(
client: &'a ComputeClient<R>,
space: &'a Space,
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
x: &[f32],
rows: usize,
k: usize,
nbanks: usize,
repeats: usize,
) -> Self {
let packed = pack_q4k(wqs, wsc, wd, wdm);
let banks: Vec<Handle> = (0..nbanks.max(1))
.map(|_| client.create_from_slice(u32::as_bytes(&packed)))
.collect();
let xh = client.create_from_slice(f32::as_bytes(x));
let meta = client.create_from_slice(u32::as_bytes(&[k as u32, rows as u32]));
let outh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; rows]));
let oracle = matvec_q4k_ref(wqs, wsc, wd, wdm, x, rows, k);
let maxref = oracle.iter().fold(0f32, |a, &v| a.max(v.abs())).max(1e-30);
Self {
client,
space,
banks,
xh,
meta,
outh,
packed_len: packed.len(),
x_len: x.len(),
rows,
k,
oracle,
maxref,
repeats,
worst_rel: std::cell::Cell::new(0.0),
}
}
pub fn space(&self) -> &Space {
self.space
}
pub fn worst_rel(&self) -> f32 {
self.worst_rel.get()
}
fn dispatch(&self, bank: &Handle, nt: usize, nr: usize) {
let grid = (self.rows as u32).div_ceil(nr as u32);
unsafe {
matvec_q4k_f32_blk::launch_unchecked::<f32, R>(
self.client,
Grid::Static(grid, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(bank.clone(), self.packed_len),
ArrayArg::from_raw_parts(self.xh.clone(), self.x_len),
ArrayArg::from_raw_parts(self.outh.clone(), self.rows),
ArrayArg::from_raw_parts(self.meta.clone(), 2),
nt,
nr,
Target::of(self.client),
);
}
}
fn read_out(&self) -> Vec<f32> {
f32::from_bytes(&self.client.read_one_unchecked(self.outh.clone())).to_vec()
}
}
impl<'a, R: Runtime> Evaluator for MatvecQ4kF32Eval<'a, R> {
fn static_check(&self, cfg: &Config) -> Verdict {
let (nt, _nr) = matvec_q4k_f32_cfg(cfg, self.space);
if nt < 2 || nt > 1024 || (nt & (nt - 1)) != 0 {
return Verdict::Reject(format!(
"workgroup width {nt} must be a power of two in [2,1024]"
));
}
Verdict::Pass
}
fn measure(&self, cfg: &Config, iters: usize) -> f64 {
let (nt, nr) = matvec_q4k_f32_cfg(cfg, self.space);
self.dispatch(&self.banks[0], nt, nr);
let got = self.read_out();
let rel = got
.iter()
.zip(&self.oracle)
.map(|(a, b)| (a - b).abs())
.fold(0f32, f32::max)
/ self.maxref;
self.worst_rel.set(self.worst_rel.get().max(rel));
if rel > 1e-3 {
return f64::INFINITY;
}
let nb = self.banks.len();
for i in 0..(2 * nb) {
self.dispatch(&self.banks[i % nb], nt, nr);
}
let _ = self.read_out();
(0..self.repeats)
.map(|_| {
let t = std::time::Instant::now();
for i in 0..iters {
self.dispatch(&self.banks[i % nb], nt, nr);
}
let _ = self.read_out();
t.elapsed().as_secs_f64() * 1e3 / iters as f64
})
.fold(f64::INFINITY, f64::min)
}
}
#[allow(clippy::too_many_arguments)]
pub fn matvec_q4k_f32_hunt<R: Runtime>(
tuner: &Tuner,
device: &str,
rows: usize,
k: usize,
eval: &MatvecQ4kF32Eval<R>,
evo: &Evolution,
seed: u64,
) -> Evolved {
tuner.evolve(
device,
"matvec_q4k_f32",
&format!("rows={rows},k={k}"),
eval.space(),
eval,
evo,
seed,
)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_matvec_q4k_blk_run<R: Runtime>(
client: &ComputeClient<R>,
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
x: &[f32],
ids: &[u32],
slots: usize,
n: usize,
k: usize,
nt: usize,
) -> Vec<f32> {
let (qh, sh, dh, mh, xh, ih, oh) = (
client.create_from_slice(u32::as_bytes(wqs)),
client.create_from_slice(u32::as_bytes(wsc)),
client.create_from_slice(f32::as_bytes(wd)),
client.create_from_slice(f32::as_bytes(wdm)),
client.create_from_slice(f32::as_bytes(x)),
client.create_from_slice(u32::as_bytes(ids)),
client.create_from_slice(f32::as_bytes(&vec![0.0f32; slots * n])),
);
unsafe {
moe_matvec_q4k_blk::launch_unchecked::<f32, R>(
client,
Grid::Static((slots * n) as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(qh.clone(), wqs.len()),
ArrayArg::from_raw_parts(sh.clone(), wsc.len()),
ArrayArg::from_raw_parts(dh.clone(), wd.len()),
ArrayArg::from_raw_parts(mh.clone(), wdm.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(ih.clone(), ids.len()),
ArrayArg::from_raw_parts(oh.clone(), slots * n),
n,
k,
nt,
);
}
f32::from_bytes(&client.read_one_unchecked(oh)).to_vec()
}
#[allow(clippy::too_many_arguments)]
pub fn moe_matvec_q4k_blk_bench<R: Runtime>(
client: &ComputeClient<R>,
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
x: &[f32],
ids: &[u32],
slots: usize,
n: usize,
k: usize,
nt: usize,
iters: usize,
) -> f64 {
let (qh, sh, dh, mh, xh, ih, oh) = (
client.create_from_slice(u32::as_bytes(wqs)),
client.create_from_slice(u32::as_bytes(wsc)),
client.create_from_slice(f32::as_bytes(wd)),
client.create_from_slice(f32::as_bytes(wdm)),
client.create_from_slice(f32::as_bytes(x)),
client.create_from_slice(u32::as_bytes(ids)),
client.create_from_slice(f32::as_bytes(&vec![0.0f32; slots * n])),
);
let launch = |c: &ComputeClient<R>| unsafe {
moe_matvec_q4k_blk::launch_unchecked::<f32, R>(
c,
Grid::Static((slots * n) as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(qh.clone(), wqs.len()),
ArrayArg::from_raw_parts(sh.clone(), wsc.len()),
ArrayArg::from_raw_parts(dh.clone(), wd.len()),
ArrayArg::from_raw_parts(mh.clone(), wdm.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(ih.clone(), ids.len()),
ArrayArg::from_raw_parts(oh.clone(), slots * n),
n,
k,
nt,
);
};
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t = std::time::Instant::now();
for _ in 0..iters {
launch(client);
}
let _ = client.read_one_unchecked(oh);
t.elapsed().as_secs_f64() * 1e3 / iters as f64
}
#[allow(clippy::too_many_arguments)]
pub fn moe_matvec_q4k_bench<R: Runtime>(
client: &ComputeClient<R>,
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
x: &[f32],
ids: &[u32],
slots: usize,
n: usize,
k: usize,
iters: usize,
) -> f64 {
let qh = client.create_from_slice(u32::as_bytes(wqs));
let sh = client.create_from_slice(u32::as_bytes(wsc));
let dh = client.create_from_slice(f32::as_bytes(wd));
let mh = client.create_from_slice(f32::as_bytes(wdm));
let xh = client.create_from_slice(f32::as_bytes(x));
let ih = client.create_from_slice(u32::as_bytes(ids));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; slots * n]));
let block = 64u32;
let grid = ((slots * n) as u32).div_ceil(block);
let launch = |c: &ComputeClient<R>| unsafe {
moe_matvec_q4k::launch_unchecked::<f32, R>(
c,
Grid::Static(grid, 1, 1),
Block::new_1d(block),
ArrayArg::from_raw_parts(qh.clone(), wqs.len()),
ArrayArg::from_raw_parts(sh.clone(), wsc.len()),
ArrayArg::from_raw_parts(dh.clone(), wd.len()),
ArrayArg::from_raw_parts(mh.clone(), wdm.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(ih.clone(), ids.len()),
ArrayArg::from_raw_parts(oh.clone(), slots * n),
n,
k,
);
};
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t = std::time::Instant::now();
for _ in 0..iters {
launch(client);
}
let _ = client.read_one_unchecked(oh);
t.elapsed().as_secs_f64() * 1e3 / iters as f64
}
#[allow(clippy::too_many_arguments)]
pub fn moe_matvec_q4k_ref(
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
x: &[f32],
ids: &[u32],
slots: usize,
n: usize,
k: usize,
) -> Vec<f32> {
let nb = k / 256;
let mut out = vec![0.0f32; slots * n];
for slot in 0..slots {
let expert = ids[slot] as usize;
for r in 0..n {
let wrow = expert * n + r;
let mut acc = 0.0f32;
for b in 0..nb {
let blk = wrow * nb + b;
let (qbase, scbase) = (blk * 32, blk * 3);
let (d, dmin) = (wd[blk], wdm[blk]);
let xblk = slot * k + b * 256;
for g in 0..4 {
let is = g * 2;
let d1 = d * cpu_sc(wsc, scbase, is) as f32;
let mm1 = dmin * cpu_m(wsc, scbase, is) as f32;
let d2 = d * cpu_sc(wsc, scbase, is + 1) as f32;
let mm2 = dmin * cpu_m(wsc, scbase, is + 1) as f32;
for qi in 0..32 {
let qb = cpu_byte(wqs, qbase, g * 32 + qi);
acc += (d1 * (qb & 15) as f32 - mm1) * x[xblk + g * 64 + qi];
acc += (d2 * (qb >> 4) as f32 - mm2) * x[xblk + g * 64 + 32 + qi];
}
}
}
out[slot * n + r] = acc;
}
}
out
}
pub fn gen_moe_q4k(
e: usize,
n: usize,
slots: usize,
k: usize,
) -> (Vec<u32>, Vec<u32>, Vec<f32>, Vec<f32>, Vec<f32>, Vec<u32>) {
let (wqs, wsc, wd, wdm, _x1) = gen_q4k(e * n, k);
let mut s = 0xC2B2AE3D27D4EB4Fu64;
let mut next = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
let x: Vec<f32> = (0..slots * k)
.map(|_| (next() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let ids: Vec<u32> = (0..slots).map(|_| (next() % e as u64) as u32).collect();
(wqs, wsc, wd, wdm, x, ids)
}
#[device]
fn sbyte_at(a: &Array<u32>, base: usize, i: usize) -> i32 {
let b = (a[base + i / 4] >> ((8 * (i % 4)) as u32)) & 255;
(b ^ 128) as i32 - 128
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn moe_matvec_q6k<F: Float>(
wql: &Array<u32>,
wqh: &Array<u32>,
wsc: &Array<u32>,
wd: &Array<F>,
x: &Array<F>,
ids: &Array<u32>,
out: &mut Array<F>,
#[comptime] n: usize,
#[comptime] k: usize,
) {
let gid = ABSOLUTE_POS;
if gid < out.len() {
let slot = gid / n;
let r = gid % n;
let wrow = ids[slot] as usize * n + r;
let nb = k / 256;
let xbase = slot * k;
let mut acc = F::new(0.0);
for b in 0..nb {
let blk = wrow * nb + b;
let qlb = blk * 32; let qhb = blk * 16; let scb = blk * 4; let d = wd[blk];
let xblk = xbase + b * 256;
for idx in 0..2 {
let scoff = idx * 8;
let qloff = idx * 64;
let qhoff = idx * 32;
let xb = xblk + idx * 128;
for l in 0..32 {
let is = l / 16;
let qll = byte_at(wql, qlb, qloff + l);
let qlh = byte_at(wql, qlb, qloff + l + 32);
let qhv = byte_at(wqh, qhb, qhoff + l);
let q1 = ((qll & 15) | ((qhv & 3) << 4)) as i32 - 32;
let q2 = ((qlh & 15) | (((qhv >> 2) & 3) << 4)) as i32 - 32;
let q3 = ((qll >> 4) | (((qhv >> 4) & 3) << 4)) as i32 - 32;
let q4 = ((qlh >> 4) | (((qhv >> 6) & 3) << 4)) as i32 - 32;
let s1 = d * F::cast_from(sbyte_at(wsc, scb, scoff + is));
let s2 = d * F::cast_from(sbyte_at(wsc, scb, scoff + is + 2));
let s3 = d * F::cast_from(sbyte_at(wsc, scb, scoff + is + 4));
let s4 = d * F::cast_from(sbyte_at(wsc, scb, scoff + is + 6));
acc += s1 * F::cast_from(q1) * x[xb + l];
acc += s2 * F::cast_from(q2) * x[xb + l + 32];
acc += s3 * F::cast_from(q3) * x[xb + l + 64];
acc += s4 * F::cast_from(q4) * x[xb + l + 96];
}
}
}
out[gid] = acc;
}
}
#[allow(clippy::too_many_arguments)]
pub fn moe_matvec_q6k_run<R: Runtime>(
client: &ComputeClient<R>,
wql: &[u32],
wqh: &[u32],
wsc: &[u32],
wd: &[f32],
x: &[f32],
ids: &[u32],
slots: usize,
n: usize,
k: usize,
) -> Vec<f32> {
let qlh = client.create_from_slice(u32::as_bytes(wql));
let qhh = client.create_from_slice(u32::as_bytes(wqh));
let sh = client.create_from_slice(u32::as_bytes(wsc));
let dh = client.create_from_slice(f32::as_bytes(wd));
let xh = client.create_from_slice(f32::as_bytes(x));
let ih = client.create_from_slice(u32::as_bytes(ids));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; slots * n]));
let block = 64u32;
let grid = ((slots * n) as u32).div_ceil(block);
unsafe {
moe_matvec_q6k::launch_unchecked::<f32, R>(
client,
Grid::Static(grid, 1, 1),
Block::new_1d(block),
ArrayArg::from_raw_parts(qlh.clone(), wql.len()),
ArrayArg::from_raw_parts(qhh.clone(), wqh.len()),
ArrayArg::from_raw_parts(sh.clone(), wsc.len()),
ArrayArg::from_raw_parts(dh.clone(), wd.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(ih.clone(), ids.len()),
ArrayArg::from_raw_parts(oh.clone(), slots * n),
n,
k,
);
}
f32::from_bytes(&client.read_one_unchecked(oh)).to_vec()
}
#[allow(clippy::too_many_arguments)]
pub fn moe_matvec_q6k_bench<R: Runtime>(
client: &ComputeClient<R>,
wql: &[u32],
wqh: &[u32],
wsc: &[u32],
wd: &[f32],
x: &[f32],
ids: &[u32],
slots: usize,
n: usize,
k: usize,
iters: usize,
) -> f64 {
let qlh = client.create_from_slice(u32::as_bytes(wql));
let qhh = client.create_from_slice(u32::as_bytes(wqh));
let sh = client.create_from_slice(u32::as_bytes(wsc));
let dh = client.create_from_slice(f32::as_bytes(wd));
let xh = client.create_from_slice(f32::as_bytes(x));
let ih = client.create_from_slice(u32::as_bytes(ids));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; slots * n]));
let block = 64u32;
let grid = ((slots * n) as u32).div_ceil(block);
let launch = |c: &ComputeClient<R>| unsafe {
moe_matvec_q6k::launch_unchecked::<f32, R>(
c,
Grid::Static(grid, 1, 1),
Block::new_1d(block),
ArrayArg::from_raw_parts(qlh.clone(), wql.len()),
ArrayArg::from_raw_parts(qhh.clone(), wqh.len()),
ArrayArg::from_raw_parts(sh.clone(), wsc.len()),
ArrayArg::from_raw_parts(dh.clone(), wd.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(ih.clone(), ids.len()),
ArrayArg::from_raw_parts(oh.clone(), slots * n),
n,
k,
);
};
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t = std::time::Instant::now();
for _ in 0..iters {
launch(client);
}
let _ = client.read_one_unchecked(oh);
t.elapsed().as_secs_f64() * 1e3 / iters as f64
}
#[inline]
fn cpu_sbyte(a: &[u32], base: usize, i: usize) -> i32 {
let b = (a[base + i / 4] >> (8 * (i % 4))) & 255;
(b ^ 128) as i32 - 128
}
#[allow(clippy::too_many_arguments)]
pub fn moe_matvec_q6k_ref(
wql: &[u32],
wqh: &[u32],
wsc: &[u32],
wd: &[f32],
x: &[f32],
ids: &[u32],
slots: usize,
n: usize,
k: usize,
) -> Vec<f32> {
let nb = k / 256;
let mut out = vec![0.0f32; slots * n];
for slot in 0..slots {
let expert = ids[slot] as usize;
for r in 0..n {
let wrow = expert * n + r;
let mut acc = 0.0f32;
for b in 0..nb {
let blk = wrow * nb + b;
let (qlb, qhb, scb) = (blk * 32, blk * 16, blk * 4);
let d = wd[blk];
let xblk = slot * k + b * 256;
for idx in 0..2 {
let (scoff, qloff, qhoff) = (idx * 8, idx * 64, idx * 32);
let xb = xblk + idx * 128;
for l in 0..32 {
let is = l / 16;
let qll = cpu_byte(wql, qlb, qloff + l);
let qlh = cpu_byte(wql, qlb, qloff + l + 32);
let qhv = cpu_byte(wqh, qhb, qhoff + l);
let q1 = ((qll & 15) | ((qhv & 3) << 4)) as i32 - 32;
let q2 = ((qlh & 15) | (((qhv >> 2) & 3) << 4)) as i32 - 32;
let q3 = ((qll >> 4) | (((qhv >> 4) & 3) << 4)) as i32 - 32;
let q4 = ((qlh >> 4) | (((qhv >> 6) & 3) << 4)) as i32 - 32;
let s1 = d * cpu_sbyte(wsc, scb, scoff + is) as f32;
let s2 = d * cpu_sbyte(wsc, scb, scoff + is + 2) as f32;
let s3 = d * cpu_sbyte(wsc, scb, scoff + is + 4) as f32;
let s4 = d * cpu_sbyte(wsc, scb, scoff + is + 6) as f32;
acc += s1 * q1 as f32 * x[xb + l];
acc += s2 * q2 as f32 * x[xb + l + 32];
acc += s3 * q3 as f32 * x[xb + l + 64];
acc += s4 * q4 as f32 * x[xb + l + 96];
}
}
}
out[slot * n + r] = acc;
}
}
out
}
pub fn gen_moe_q6k(
e: usize,
n: usize,
slots: usize,
k: usize,
) -> (Vec<u32>, Vec<u32>, Vec<u32>, Vec<f32>, Vec<f32>, Vec<u32>) {
let nb = k / 256;
let nblk = e * n * nb;
let mut s = 0x27D4EB2F165667C5u64;
let mut next = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
let wql: Vec<u32> = (0..nblk * 32).map(|_| next() as u32).collect(); let wqh: Vec<u32> = (0..nblk * 16).map(|_| next() as u32).collect(); let wsc: Vec<u32> = (0..nblk * 4).map(|_| next() as u32).collect(); let wd: Vec<f32> = (0..nblk)
.map(|_| half::f16::from_f32((next() % 1000) as f32 / 20000.0 + 0.002).to_f32())
.collect();
let x: Vec<f32> = (0..slots * k)
.map(|_| (next() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let ids: Vec<u32> = (0..slots).map(|_| (next() % e as u64) as u32).collect();
(wql, wqh, wsc, wd, x, ids)
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn moe_matvec_q6k_blk<F: Float>(
wql: &Array<u32>,
wqh: &Array<u32>,
wsc: &Array<u32>,
wd: &Array<F>,
x: &Array<F>,
ids: &Array<u32>,
out: &mut Array<F>,
#[comptime] n: usize, #[comptime] k: usize,
#[comptime] nt: usize,
) {
let outrow = CUBE_POS as usize;
let t = UNIT_POS as usize;
let slot = outrow / n;
let r = outrow % n;
let wrow = ids[slot] as usize * n + r;
let nb = k / 256;
let nsub = k / 32;
let per = (nsub + nt - 1) / nt; let xrow = slot * k;
let mut partial = F::new(0.0);
for j in 0..per {
let sb = j * nt + t;
if sb < nsub {
let sup = sb / 8;
let sbloc = sb % 8;
let idx = sbloc / 4; let sub = sbloc % 4; let blk = wrow * nb + sup;
let d = wd[blk];
let qlbase = blk * 32;
let qhbase = blk * 16;
let scbase = blk * 4;
let qloff = idx * 64 + (sub % 2) * 32; let qhoff = idx * 32; let qlshift = ((sub / 2) * 4) as u32; let qhshift = (sub * 2) as u32; let scb0 = idx * 8 + sub * 2; let xoff = xrow + sup * 256 + idx * 128 + sub * 32;
let mut sbsum = F::new(0.0);
for l in 0..32 {
let is = l / 16;
let sc = sbyte_at(wsc, scbase, scb0 + is);
let qlb = byte_at(wql, qlbase, qloff + l);
let qhb = byte_at(wqh, qhbase, qhoff + l);
let q = ((qlb >> qlshift) & 15) | (((qhb >> qhshift) & 3) << 4);
let qv = q as i32 - 32;
sbsum += d * F::cast_from(sc) * F::cast_from(qv) * x[xoff + l];
}
partial += sbsum;
}
}
let mut smem = SharedMemory::<F>::new(nt);
smem[t] = partial;
sync_cube();
let mut stride = CUBE_DIM / 2;
while stride > 0 {
if UNIT_POS < stride {
let v = smem[(UNIT_POS + stride) as usize];
smem[t] += v;
}
sync_cube();
stride /= 2;
}
if t == 0 {
out[outrow] = smem[0];
}
}
#[allow(clippy::too_many_arguments)]
pub fn moe_matvec_q6k_blk_run<R: Runtime>(
client: &ComputeClient<R>,
wql: &[u32],
wqh: &[u32],
wsc: &[u32],
wd: &[f32],
x: &[f32],
ids: &[u32],
slots: usize,
n: usize,
k: usize,
nt: usize,
) -> Vec<f32> {
let (qlh, qhh, sh, dh, xh, ih, oh) = (
client.create_from_slice(u32::as_bytes(wql)),
client.create_from_slice(u32::as_bytes(wqh)),
client.create_from_slice(u32::as_bytes(wsc)),
client.create_from_slice(f32::as_bytes(wd)),
client.create_from_slice(f32::as_bytes(x)),
client.create_from_slice(u32::as_bytes(ids)),
client.create_from_slice(f32::as_bytes(&vec![0.0f32; slots * n])),
);
unsafe {
moe_matvec_q6k_blk::launch_unchecked::<f32, R>(
client,
Grid::Static((slots * n) as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(qlh.clone(), wql.len()),
ArrayArg::from_raw_parts(qhh.clone(), wqh.len()),
ArrayArg::from_raw_parts(sh.clone(), wsc.len()),
ArrayArg::from_raw_parts(dh.clone(), wd.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(ih.clone(), ids.len()),
ArrayArg::from_raw_parts(oh.clone(), slots * n),
n,
k,
nt,
);
}
f32::from_bytes(&client.read_one_unchecked(oh)).to_vec()
}
#[allow(clippy::too_many_arguments)]
pub fn moe_matvec_q6k_blk_bench<R: Runtime>(
client: &ComputeClient<R>,
wql: &[u32],
wqh: &[u32],
wsc: &[u32],
wd: &[f32],
x: &[f32],
ids: &[u32],
slots: usize,
n: usize,
k: usize,
nt: usize,
iters: usize,
) -> f64 {
let (qlh, qhh, sh, dh, xh, ih, oh) = (
client.create_from_slice(u32::as_bytes(wql)),
client.create_from_slice(u32::as_bytes(wqh)),
client.create_from_slice(u32::as_bytes(wsc)),
client.create_from_slice(f32::as_bytes(wd)),
client.create_from_slice(f32::as_bytes(x)),
client.create_from_slice(u32::as_bytes(ids)),
client.create_from_slice(f32::as_bytes(&vec![0.0f32; slots * n])),
);
let launch = |c: &ComputeClient<R>| unsafe {
moe_matvec_q6k_blk::launch_unchecked::<f32, R>(
c,
Grid::Static((slots * n) as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(qlh.clone(), wql.len()),
ArrayArg::from_raw_parts(qhh.clone(), wqh.len()),
ArrayArg::from_raw_parts(sh.clone(), wsc.len()),
ArrayArg::from_raw_parts(dh.clone(), wd.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(ih.clone(), ids.len()),
ArrayArg::from_raw_parts(oh.clone(), slots * n),
n,
k,
nt,
);
};
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t = std::time::Instant::now();
for _ in 0..iters {
launch(client);
}
let _ = client.read_one_unchecked(oh);
t.elapsed().as_secs_f64() * 1e3 / iters as f64
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
#[allow(clippy::too_many_arguments)]
pub fn moe_matvec_q6k_dp4a_blk<F: Float>(
wql: &Array<u32>,
wqh: &Array<u32>,
wsc: &Array<u32>,
wd: &Array<F>,
xq: &Array<u32>, xs: &Array<F>, ids: &Array<u32>,
out: &mut Array<F>,
#[comptime] n: usize,
#[comptime] k: usize,
#[comptime] nt: usize,
) {
let outrow = CUBE_POS as usize;
let t = UNIT_POS as usize;
let slot = outrow / n;
let r = outrow % n;
let wrow = ids[slot] as usize * n + r;
let nb = k / 256;
let nsub = k / 32;
let per = (nsub + nt - 1) / nt;
let xrow = slot * nsub;
let ones = Vector::<i32, Const<4>>::cast_from(Vector::<i8, Const<4>>::reinterpret::<u32>(
0x0101_0101u32,
));
let mut partial = F::new(0.0);
for j in 0..per {
let sb = j * nt + t;
if sb < nsub {
let sup = sb / 8;
let sbloc = sb % 8;
let idx = sbloc / 4;
let sub = sbloc % 4;
let blk = wrow * nb + sup;
let d = wd[blk];
let qlw0 = blk * 32 + (idx * 64 + (sub % 2) * 32) / 4; let qhw0 = blk * 16 + (idx * 32) / 4; let scbase = blk * 4;
let qlshift = ((sub / 2) * 4) as u32;
let qhshift = (sub * 2) as u32;
let scb0 = idx * 8 + sub * 2;
let sc0 = F::cast_from(sbyte_at(wsc, scbase, scb0)); let sc1 = F::cast_from(sbyte_at(wsc, scbase, scb0 + 1)); let xg = (xrow + sb) * 8;
let mut idot0 = 0i32;
let mut isum0 = 0i32;
let mut idot1 = 0i32;
let mut isum1 = 0i32;
for g in 0..4 {
let qw = ((wql[qlw0 + g] >> qlshift) & 0x0F0F_0F0F)
| (((wqh[qhw0 + g] >> qhshift) & 0x0303_0303) << 4);
let wv = Vector::<i32, Const<4>>::cast_from(Vector::<i8, Const<4>>::reinterpret::<
u32,
>(qw));
let xv = Vector::<i32, Const<4>>::cast_from(Vector::<i8, Const<4>>::reinterpret::<
u32,
>(xq[xg + g]));
idot0 += wv.dot(xv);
isum0 += xv.dot(ones);
}
for g in 4..8 {
let qw = ((wql[qlw0 + g] >> qlshift) & 0x0F0F_0F0F)
| (((wqh[qhw0 + g] >> qhshift) & 0x0303_0303) << 4);
let wv = Vector::<i32, Const<4>>::cast_from(Vector::<i8, Const<4>>::reinterpret::<
u32,
>(qw));
let xv = Vector::<i32, Const<4>>::cast_from(Vector::<i8, Const<4>>::reinterpret::<
u32,
>(xq[xg + g]));
idot1 += wv.dot(xv);
isum1 += xv.dot(ones);
}
partial += d
* xs[xrow + sb]
* (sc0 * F::cast_from(idot0 - 32 * isum0) + sc1 * F::cast_from(idot1 - 32 * isum1));
}
}
let mut smem = SharedMemory::<F>::new(nt);
smem[t] = partial;
sync_cube();
let mut stride = CUBE_DIM / 2;
while stride > 0 {
if UNIT_POS < stride {
let v = smem[(UNIT_POS + stride) as usize];
smem[t] += v;
}
sync_cube();
stride /= 2;
}
if t == 0 {
out[outrow] = smem[0];
}
}
#[allow(clippy::too_many_arguments)]
pub fn moe_matvec_q6k_dp4a_blk_run<R: Runtime>(
client: &ComputeClient<R>,
wql: &[u32],
wqh: &[u32],
wsc: &[u32],
wd: &[f32],
x: &[f32],
ids: &[u32],
slots: usize,
n: usize,
k: usize,
nt: usize,
) -> Vec<f32> {
let (xq, xs, _) = quant_act_q8_cpu(x, slots, k);
let qlh = client.create_from_slice(u32::as_bytes(wql));
let qhh = client.create_from_slice(u32::as_bytes(wqh));
let sh = client.create_from_slice(u32::as_bytes(wsc));
let dh = client.create_from_slice(f32::as_bytes(wd));
let xqh = client.create_from_slice(u32::as_bytes(&xq));
let xsh = client.create_from_slice(f32::as_bytes(&xs));
let ih = client.create_from_slice(u32::as_bytes(ids));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; slots * n]));
unsafe {
moe_matvec_q6k_dp4a_blk::launch_unchecked::<f32, R>(
client,
Grid::Static((slots * n) as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(qlh.clone(), wql.len()),
ArrayArg::from_raw_parts(qhh.clone(), wqh.len()),
ArrayArg::from_raw_parts(sh.clone(), wsc.len()),
ArrayArg::from_raw_parts(dh.clone(), wd.len()),
ArrayArg::from_raw_parts(xqh.clone(), xq.len()),
ArrayArg::from_raw_parts(xsh.clone(), xs.len()),
ArrayArg::from_raw_parts(ih.clone(), ids.len()),
ArrayArg::from_raw_parts(oh.clone(), slots * n),
n,
k,
nt,
);
}
f32::from_bytes(&client.read_one_unchecked(oh)).to_vec()
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn matvec_q8_dp4a<F: Float>(
wq: &Array<Vector<i32, Const<4>>>, xq: &Array<Vector<i32, Const<4>>>, wd: &Array<F>, out: &mut Array<F>,
#[comptime] k: usize,
) {
let row = ABSOLUTE_POS;
if row < out.len() {
let ng = k / 4;
let nb = k / 32;
let wbase = row * ng;
let dbase = row * nb;
let mut acc = F::new(0.0);
for g in 0..ng {
let dp = wq[wbase + g].dot(xq[g]); acc += wd[dbase + g / 8] * F::cast_from(dp);
}
out[row] = acc;
}
}
pub fn matvec_q8_dp4a_run<R: Runtime>(
client: &ComputeClient<R>,
wq: &[i32],
xq: &[i32],
wd: &[f32],
rows: usize,
k: usize,
bench_iters: usize,
) -> (Vec<f32>, f64) {
let wqh = client.create_from_slice(i32::as_bytes(wq));
let xqh = client.create_from_slice(i32::as_bytes(xq));
let wdh = client.create_from_slice(f32::as_bytes(wd));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; rows]));
let block = 64u32;
let grid = (rows as u32).div_ceil(block);
let ng = k / 4;
let launch = |c: &ComputeClient<R>| unsafe {
matvec_q8_dp4a::launch_unchecked::<f32, R>(
c,
Grid::Static(grid, 1, 1),
Block::new_1d(block),
ArrayArg::from_raw_parts(wqh.clone(), rows * ng),
ArrayArg::from_raw_parts(xqh.clone(), ng),
ArrayArg::from_raw_parts(wdh.clone(), wd.len()),
ArrayArg::from_raw_parts(oh.clone(), rows),
k,
);
};
launch(client);
let bytes = client.read_one_unchecked(oh.clone());
let out = f32::from_bytes(&bytes).to_vec();
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t = std::time::Instant::now();
for _ in 0..bench_iters {
launch(client);
}
let _ = client.read_one_unchecked(oh);
let ms = t.elapsed().as_secs_f64() * 1e3 / bench_iters as f64;
(out, ms)
}
pub fn matvec_q8_dp4a_ref(wq: &[i32], xq: &[i32], wd: &[f32], rows: usize, k: usize) -> Vec<f32> {
let nb = k / 32;
(0..rows)
.map(|row| {
let mut acc = 0.0f32;
for i in 0..k {
acc += wd[row * nb + i / 32] * (wq[row * k + i] * xq[i]) as f32;
}
acc
})
.collect()
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn matvec_q8_dp4a_i8<F: Float>(
wq: &Array<Vector<i8, Const<4>>>, xq: &Array<Vector<i8, Const<4>>>, wd: &Array<F>, out: &mut Array<F>,
#[comptime] k: usize,
) {
let row = ABSOLUTE_POS;
if row < out.len() {
let ng = k / 4;
let nb = k / 32;
let wbase = row * ng;
let dbase = row * nb;
let mut acc = F::new(0.0);
for g in 0..ng {
let wi = Vector::<i32, Const<4>>::cast_from(wq[wbase + g]);
let xi = Vector::<i32, Const<4>>::cast_from(xq[g]);
let dp = wi.dot(xi); acc += wd[dbase + g / 8] * F::cast_from(dp);
}
out[row] = acc;
}
}
pub fn matvec_q8_dp4a_i8_run<R: Runtime>(
client: &ComputeClient<R>,
wq: &[i8],
xq: &[i8],
wd: &[f32],
rows: usize,
k: usize,
bench_iters: usize,
) -> (Vec<f32>, f64) {
let wqh = client.create_from_slice(i8::as_bytes(wq));
let xqh = client.create_from_slice(i8::as_bytes(xq));
let wdh = client.create_from_slice(f32::as_bytes(wd));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; rows]));
let block = 64u32;
let grid = (rows as u32).div_ceil(block);
let ng = k / 4;
let launch = |c: &ComputeClient<R>| unsafe {
matvec_q8_dp4a_i8::launch_unchecked::<f32, R>(
c,
Grid::Static(grid, 1, 1),
Block::new_1d(block),
ArrayArg::from_raw_parts(wqh.clone(), rows * ng),
ArrayArg::from_raw_parts(xqh.clone(), ng),
ArrayArg::from_raw_parts(wdh.clone(), wd.len()),
ArrayArg::from_raw_parts(oh.clone(), rows),
k,
);
};
launch(client);
let bytes = client.read_one_unchecked(oh.clone());
let out = f32::from_bytes(&bytes).to_vec();
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t = std::time::Instant::now();
for _ in 0..bench_iters {
launch(client);
}
let _ = client.read_one_unchecked(oh);
let ms = t.elapsed().as_secs_f64() * 1e3 / bench_iters as f64;
(out, ms)
}
#[device]
fn f16lo_to_f32(h: u32) -> f32 {
let magic = 113u32 << 23; let shifted_exp = 0x7C00u32 << 13; let mut o = (h & 0x7FFF) << 13; let exp = shifted_exp & o;
o += (127u32 - 15) << 23; if exp == shifted_exp {
o += (128u32 - 16) << 23; }
if exp == 0u32 {
o += 1u32 << 23; let f = f32::reinterpret(o) - f32::reinterpret(magic);
o = u32::reinterpret(f);
}
o = o | ((h & 0x8000) << 16); f32::reinterpret(o)
}
#[device]
fn i8lane<F: Float>(word: u32, p: u32) -> F {
let b = (word >> p) & 255;
let mut v = F::cast_from(b);
if b >= 128 {
v -= F::new(256.0);
}
v
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn matvec_q8_0_packed_blk<F: Float>(
w: &Array<u32>, x: &Array<F>, out: &mut Array<F>,
#[comptime] k: usize,
#[comptime] nt: usize, ) {
let row = CUBE_POS as usize;
let t = UNIT_POS as usize;
let nblocks = k / 32;
let wbase = row * nblocks * 9;
let per = nblocks / nt; let mut partial = F::new(0.0);
for j in 0..per {
let b = j * nt + t;
let off = wbase + b * 9;
let scale = F::cast_from(f16lo_to_f32(w[off]));
let xb = b * 32;
let mut bsum = F::new(0.0);
for jj in 0..8 {
let word = w[off + 1 + jj];
let xo = xb + jj * 4;
bsum += i8lane::<F>(word, 0) * x[xo];
bsum += i8lane::<F>(word, 8) * x[xo + 1];
bsum += i8lane::<F>(word, 16) * x[xo + 2];
bsum += i8lane::<F>(word, 24) * x[xo + 3];
}
partial += scale * bsum;
}
let mut smem = SharedMemory::<F>::new(nt);
smem[t] = partial;
sync_cube();
let mut stride = CUBE_DIM / 2;
while stride > 0 {
if UNIT_POS < stride {
let v = smem[(UNIT_POS + stride) as usize];
smem[t] += v;
}
sync_cube();
stride /= 2;
}
if t == 0 {
out[row] = smem[0];
}
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn matvec_q8_0_packed_sg<F: Float>(
w: &Array<u32>,
x: &Array<F>,
out: &mut Array<F>,
#[comptime] k: usize,
#[comptime] nt: usize, ) {
let row = CUBE_POS as usize;
let t = UNIT_POS as usize;
let nblocks = k / 32;
let wbase = row * nblocks * 9;
let per = nblocks / nt;
let mut partial = F::new(0.0);
for j in 0..per {
let b = j * nt + t;
let off = wbase + b * 9;
let scale = F::cast_from(f16lo_to_f32(w[off]));
let xb = b * 32;
let mut bsum = F::new(0.0);
for jj in 0..8 {
let word = w[off + 1 + jj];
let xo = xb + jj * 4;
bsum += i8lane::<F>(word, 0) * x[xo];
bsum += i8lane::<F>(word, 8) * x[xo + 1];
bsum += i8lane::<F>(word, 16) * x[xo + 2];
bsum += i8lane::<F>(word, 24) * x[xo + 3];
}
partial += scale * bsum;
}
let total = plane_sum(partial);
if t == 0 {
out[row] = total;
}
}
pub fn matvec_q8_0_packed_ref(w: &[u32], x: &[f32], rows: usize, k: usize) -> Vec<f32> {
let nblocks = k / 32;
let mut out = vec![0f32; rows];
for row in 0..rows {
let wbase = row * nblocks * 9;
let mut acc = 0f32;
for b in 0..nblocks {
let off = wbase + b * 9;
let scale = half::f16::from_bits((w[off] & 0xFFFF) as u16).to_f32();
let xb = b * 32;
let mut bsum = 0f32;
for jj in 0..8 {
let word = w[off + 1 + jj];
let xo = xb + jj * 4;
for lane in 0..4 {
let q = ((word >> (8 * lane as u32)) & 0xFF) as u8 as i8 as f32;
bsum += q * x[xo + lane];
}
}
acc += scale * bsum;
}
out[row] = acc;
}
out
}
pub fn gen_q8_0_packed(rows: usize, k: usize) -> (Vec<u32>, Vec<f32>) {
let nblocks = k / 32;
let mut s = 0x1234_5678_9ABC_DEF1u64;
let mut next = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
let mut w = Vec::with_capacity(rows * nblocks * 9);
for _ in 0..rows * nblocks {
let scale = (next() % 1000) as f32 / 8000.0 + 0.01;
w.push(half::f16::from_f32(scale).to_bits() as u32);
for _ in 0..8 {
let mut word = 0u32;
for lane in 0..4 {
let q = (next() % 255) as u8; word |= (q as u32) << (8 * lane);
}
w.push(word);
}
}
let x: Vec<f32> = (0..k)
.map(|_| (next() % 2000) as f32 / 1000.0 - 1.0)
.collect();
(w, x)
}
pub fn matvec_q8_0_packed_run<R: Runtime>(
client: &ComputeClient<R>,
w: &[u32],
x: &[f32],
rows: usize,
k: usize,
nt: usize,
bench_iters: usize,
) -> (Vec<f32>, f64) {
let wh = client.create_from_slice(u32::as_bytes(w));
let xh = client.create_from_slice(f32::as_bytes(x));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; rows]));
let launch = |c: &ComputeClient<R>| unsafe {
matvec_q8_0_packed_blk::launch_unchecked::<f32, R>(
c,
Grid::Static(rows as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(wh.clone(), w.len()),
ArrayArg::from_raw_parts(xh.clone(), k),
ArrayArg::from_raw_parts(oh.clone(), rows),
k,
nt,
);
};
launch(client);
let bytes = client.read_one_unchecked(oh.clone());
let out = f32::from_bytes(&bytes).to_vec();
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t = std::time::Instant::now();
for _ in 0..bench_iters {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let ms = t.elapsed().as_secs_f64() * 1000.0 / bench_iters as f64;
(out, ms)
}
pub fn matvec_q8_0_packed_sg_run<R: Runtime>(
client: &ComputeClient<R>,
w: &[u32],
x: &[f32],
rows: usize,
k: usize,
nt: usize,
bench_iters: usize,
) -> (Vec<f32>, f64) {
let wh = client.create_from_slice(u32::as_bytes(w));
let xh = client.create_from_slice(f32::as_bytes(x));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; rows]));
let launch = |c: &ComputeClient<R>| unsafe {
matvec_q8_0_packed_sg::launch_unchecked::<f32, R>(
c,
Grid::Static(rows as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(wh.clone(), w.len()),
ArrayArg::from_raw_parts(xh.clone(), k),
ArrayArg::from_raw_parts(oh.clone(), rows),
k,
nt,
);
};
launch(client);
let bytes = client.read_one_unchecked(oh.clone());
let out = f32::from_bytes(&bytes).to_vec();
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t = std::time::Instant::now();
for _ in 0..bench_iters {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let ms = t.elapsed().as_secs_f64() * 1000.0 / bench_iters as f64;
(out, ms)
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn matvec_q8_dp4a_blk<F: Float>(
wq: &Array<Vector<i8, Const<4>>>, xq: &Array<Vector<i8, Const<4>>>, wd: &Array<F>, out: &mut Array<F>,
#[comptime] k: usize,
#[comptime] nt: usize, ) {
let row = CUBE_POS as usize;
let t = UNIT_POS as usize;
let ng = k / 4;
let wbase = row * ng;
let dbase = row * (k / 32);
let mut partial = F::new(0.0);
let per = ng / nt; for j in 0..per {
let g = j * nt + t;
let wi = Vector::<i32, Const<4>>::cast_from(wq[wbase + g]);
let xi = Vector::<i32, Const<4>>::cast_from(xq[g]);
partial += wd[dbase + g / 8] * F::cast_from(wi.dot(xi)); }
let mut smem = SharedMemory::<F>::new(nt);
smem[t] = partial;
sync_cube();
let mut stride = CUBE_DIM / 2; while stride > 0 {
if UNIT_POS < stride {
let v = smem[(UNIT_POS + stride) as usize];
smem[t] += v;
}
sync_cube();
stride /= 2;
}
if t == 0 {
out[row] = smem[0];
}
}
pub fn matvec_q8_dp4a_blk_run<R: Runtime>(
client: &ComputeClient<R>,
wq: &[i8],
xq: &[i8],
wd: &[f32],
rows: usize,
k: usize,
nt: usize,
bench_iters: usize,
) -> (Vec<f32>, f64) {
let wqh = client.create_from_slice(i8::as_bytes(wq));
let xqh = client.create_from_slice(i8::as_bytes(xq));
let wdh = client.create_from_slice(f32::as_bytes(wd));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; rows]));
let ng = k / 4;
let launch = |c: &ComputeClient<R>| unsafe {
matvec_q8_dp4a_blk::launch_unchecked::<f32, R>(
c,
Grid::Static(rows as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(wqh.clone(), rows * ng),
ArrayArg::from_raw_parts(xqh.clone(), ng),
ArrayArg::from_raw_parts(wdh.clone(), wd.len()),
ArrayArg::from_raw_parts(oh.clone(), rows),
k,
nt,
);
};
launch(client);
let bytes = client.read_one_unchecked(oh.clone());
let out = f32::from_bytes(&bytes).to_vec();
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t = std::time::Instant::now();
for _ in 0..bench_iters {
launch(client);
}
let _ = client.read_one_unchecked(oh);
let ms = t.elapsed().as_secs_f64() * 1e3 / bench_iters as f64;
(out, ms)
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn matvec_q8_dp4a_tuned<F: Float>(
wq: &Array<Vector<i8, Const<4>>>,
xq: &Array<Vector<i8, Const<4>>>,
wd: &Array<F>,
out: &mut Array<F>,
#[comptime] k: usize,
#[comptime] vw: usize, #[comptime] nr: usize, ) {
let nout = out.len();
let ng = k / 4;
let nb = k / 32;
let row0 = ABSOLUTE_POS * nr; let mut acc = Array::<F>::new(nr);
#[unroll]
for n in 0..nr {
acc[n] = F::new(0.0);
}
let steps = ng / vw;
for s in 0..steps {
#[unroll]
for j in 0..vw {
let g = s * vw + j;
let xi = Vector::<i32, Const<4>>::cast_from(xq[g]); #[unroll]
for n in 0..nr {
let row = row0 + n;
if row < nout {
let wi = Vector::<i32, Const<4>>::cast_from(wq[row * ng + g]);
acc[n] += wd[row * nb + g / 8] * F::cast_from(wi.dot(xi));
}
}
}
}
#[unroll]
for n in 0..nr {
let row = row0 + n;
if row < nout {
out[row] = acc[n];
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn matvec_q8_dp4a_tuned_bench<R: Runtime>(
client: &ComputeClient<R>,
wq: &[i8],
xq: &[i8],
wd: &[f32],
rows: usize,
k: usize,
vw: usize,
block: u32,
nr: usize,
iters: usize,
) -> (Vec<f32>, f64) {
let wqh = client.create_from_slice(i8::as_bytes(wq));
let xqh = client.create_from_slice(i8::as_bytes(xq));
let wdh = client.create_from_slice(f32::as_bytes(wd));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; rows]));
let grid = (rows as u32).div_ceil(nr as u32).div_ceil(block);
let ng = k / 4;
let launch = |c: &ComputeClient<R>| unsafe {
matvec_q8_dp4a_tuned::launch_unchecked::<f32, R>(
c,
Grid::Static(grid, 1, 1),
Block::new_1d(block),
ArrayArg::from_raw_parts(wqh.clone(), rows * ng),
ArrayArg::from_raw_parts(xqh.clone(), ng),
ArrayArg::from_raw_parts(wdh.clone(), wd.len()),
ArrayArg::from_raw_parts(oh.clone(), rows),
k,
vw,
nr,
);
};
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t = std::time::Instant::now();
for _ in 0..iters {
launch(client);
}
let out = client.read_one_unchecked(oh);
let ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
(f32::from_bytes(&out).to_vec(), ms)
}
pub fn matvec_q8_dp4a_tuned_set<'a, R: Runtime>(
client: &'a ComputeClient<R>,
wq: &'a [i8],
xq: &'a [i8],
wd: &'a [f32],
rows: usize,
k: usize,
) -> Tuned<'a, Vec<f32>> {
Tuned::new("matvec_q8_dp4a", format!("rows={rows},k={k}"))
.variant("b64_v1", move |it| {
matvec_q8_dp4a_tuned_bench(client, wq, xq, wd, rows, k, 1, 64, 1, it)
})
.variant("b64_v4", move |it| {
matvec_q8_dp4a_tuned_bench(client, wq, xq, wd, rows, k, 4, 64, 1, it)
})
.variant("b128_v2", move |it| {
matvec_q8_dp4a_tuned_bench(client, wq, xq, wd, rows, k, 2, 128, 1, it)
})
.variant("b256_v8", move |it| {
matvec_q8_dp4a_tuned_bench(client, wq, xq, wd, rows, k, 8, 256, 1, it)
})
}
pub fn matvec_q8_dp4a_autotuned<R: Runtime>(
client: &ComputeClient<R>,
wq: &[i8],
xq: &[i8],
wd: &[f32],
rows: usize,
k: usize,
) -> Vec<f32> {
matvec_q8_dp4a_tuned_set(client, wq, xq, wd, rows, k).run(client)
}
pub fn matvec_dp4a_space(k: usize) -> Space {
let ng = (k / 4) as i64;
Space::new()
.param("WG", [64, 128, 256])
.param("VW", [1, 2, 4, 8])
.param("NR", [1, 2, 4])
.constraint(move |c, s| ng % c.get(s, "VW") == 0)
.constraint(|c, s| {
let wg = c.get(s, "WG");
wg > 0 && wg <= 1024
})
}
fn dp4a_cfg(c: &Config, s: &Space) -> (u32, usize, usize) {
(
c.get(s, "WG") as u32,
c.get(s, "VW") as usize,
c.get(s, "NR") as usize,
)
}
pub struct MatvecDp4aEval<'a, R: Runtime> {
client: &'a ComputeClient<R>,
space: &'a Space,
banks: Vec<Handle>, xqh: Handle,
wdh: Handle,
outh: Handle,
wd_len: usize,
rows: usize,
k: usize,
oracle: Vec<f32>,
maxref: f32,
repeats: usize,
worst_rel: std::cell::Cell<f32>,
}
impl<'a, R: Runtime> MatvecDp4aEval<'a, R> {
#[allow(clippy::too_many_arguments)]
pub fn new(
client: &'a ComputeClient<R>,
space: &'a Space,
weight_banks: &[Vec<i8>],
xq: &[i8],
wd: &[f32],
rows: usize,
k: usize,
repeats: usize,
) -> Self {
assert!(
!weight_banks.is_empty(),
"MatvecDp4aEval needs at least one weight bank"
);
let banks: Vec<Handle> = weight_banks
.iter()
.map(|w| client.create_from_slice(i8::as_bytes(w)))
.collect();
let xqh = client.create_from_slice(i8::as_bytes(xq));
let wdh = client.create_from_slice(f32::as_bytes(wd));
let outh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; rows]));
let wq32: Vec<i32> = weight_banks[0].iter().map(|&v| v as i32).collect();
let xq32: Vec<i32> = xq.iter().map(|&v| v as i32).collect();
let oracle = matvec_q8_dp4a_ref(&wq32, &xq32, wd, rows, k);
let maxref = oracle.iter().fold(0f32, |a, &v| a.max(v.abs())).max(1e-30);
Self {
client,
space,
banks,
xqh,
wdh,
outh,
wd_len: wd.len(),
rows,
k,
oracle,
maxref,
repeats,
worst_rel: std::cell::Cell::new(0.0),
}
}
pub fn space(&self) -> &Space {
self.space
}
pub fn worst_rel(&self) -> f32 {
self.worst_rel.get()
}
fn dispatch(&self, bank: &Handle, block: u32, vw: usize, nr: usize) {
let ng = self.k / 4;
let grid = (self.rows as u32).div_ceil(nr as u32).div_ceil(block);
unsafe {
matvec_q8_dp4a_tuned::launch_unchecked::<f32, R>(
self.client,
Grid::Static(grid, 1, 1),
Block::new_1d(block),
ArrayArg::from_raw_parts(bank.clone(), self.rows * ng),
ArrayArg::from_raw_parts(self.xqh.clone(), ng),
ArrayArg::from_raw_parts(self.wdh.clone(), self.wd_len),
ArrayArg::from_raw_parts(self.outh.clone(), self.rows),
self.k,
vw,
nr,
);
}
}
fn read_out(&self) -> Vec<f32> {
f32::from_bytes(&self.client.read_one_unchecked(self.outh.clone())).to_vec()
}
}
impl<'a, R: Runtime> Evaluator for MatvecDp4aEval<'a, R> {
fn static_check(&self, cfg: &Config) -> Verdict {
let (block, vw, _nr) = dp4a_cfg(cfg, self.space);
let ng = self.k / 4;
if vw == 0 || ng % vw != 0 {
return Verdict::Reject(format!("ng={ng} not a multiple of vw={vw}"));
}
if block == 0 || block > 1024 {
return Verdict::Reject(format!("illegal workgroup width {block}"));
}
Verdict::Pass
}
fn measure(&self, cfg: &Config, iters: usize) -> f64 {
let (block, vw, nr) = dp4a_cfg(cfg, self.space);
self.dispatch(&self.banks[0], block, vw, nr);
let got = self.read_out();
let rel = got
.iter()
.zip(&self.oracle)
.map(|(a, b)| (a - b).abs())
.fold(0f32, f32::max)
/ self.maxref;
self.worst_rel.set(self.worst_rel.get().max(rel));
if rel > 2e-3 {
return f64::INFINITY;
}
let nb = self.banks.len();
for i in 0..(2 * nb) {
self.dispatch(&self.banks[i % nb], block, vw, nr);
}
let _ = self.read_out();
(0..self.repeats)
.map(|_| {
let t = std::time::Instant::now();
for i in 0..iters {
self.dispatch(&self.banks[i % nb], block, vw, nr);
}
let _ = self.read_out();
t.elapsed().as_secs_f64() * 1e3 / iters as f64
})
.fold(f64::INFINITY, f64::min)
}
}
#[allow(clippy::too_many_arguments)]
pub fn matvec_dp4a_hunt<R: Runtime>(
tuner: &Tuner,
device: &str,
rows: usize,
k: usize,
eval: &MatvecDp4aEval<R>,
evo: &Evolution,
seed: u64,
) -> Evolved {
tuner.evolve(
device,
"matvec_q8_dp4a",
&format!("rows={rows},k={k}"),
eval.space(),
eval,
evo,
seed,
)
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn moe_route<F: Float>(
logits: &Array<F>, ids_out: &mut Array<u32>, w_out: &mut Array<F>, #[comptime] n_experts: usize,
#[comptime] topk: usize,
#[comptime] nt: usize,
) {
let tok = CUBE_POS as usize;
let t = UNIT_POS as usize;
let base = tok * n_experts;
let ninf = F::new(-3.4e38); let mut slog = SharedMemory::<F>::new(n_experts);
let mut i = t;
while i < n_experts {
slog[i] = logits[base + i];
i += nt;
}
sync_cube();
let mut sred = SharedMemory::<F>::new(nt);
let mut lmax = ninf;
let mut a = t;
while a < n_experts {
if slog[a] > lmax {
lmax = slog[a];
}
a += nt;
}
sred[t] = lmax;
sync_cube();
let mut stride = CUBE_DIM / 2;
while stride > 0 {
if UNIT_POS < stride {
let v = sred[(UNIT_POS + stride) as usize];
let cur = sred[t];
if v > cur {
sred[t] = v;
}
}
sync_cube();
stride /= 2;
}
let m = sred[0];
sync_cube();
let mut lsum = F::new(0.0);
let mut b = t;
while b < n_experts {
lsum += (slog[b] - m).exp();
b += nt;
}
sred[t] = lsum;
sync_cube();
let mut st2 = CUBE_DIM / 2;
while st2 > 0 {
if UNIT_POS < st2 {
let v = sred[(UNIT_POS + st2) as usize];
sred[t] += v;
}
sync_cube();
st2 /= 2;
}
let denom = sred[0];
sync_cube();
let mut sidx = SharedMemory::<u32>::new(nt);
let mut wsum = F::new(0.0);
for _r in 0..topk {
let mut lv = ninf;
let mut li = 0u32;
let mut c = t;
while c < n_experts {
if slog[c] > lv {
lv = slog[c];
li = c as u32;
}
c += nt;
}
sred[t] = lv;
sidx[t] = li;
sync_cube();
let mut sr = CUBE_DIM / 2;
while sr > 0 {
if UNIT_POS < sr {
let ov = sred[(UNIT_POS + sr) as usize];
let oi = sidx[(UNIT_POS + sr) as usize];
let curv = sred[t];
if ov > curv {
sred[t] = ov;
sidx[t] = oi;
}
}
sync_cube();
sr /= 2;
}
let best = sidx[0];
let wr = (sred[0] - m).exp() / denom; wsum += wr;
if t == 0 {
ids_out[tok * topk + _r] = best;
w_out[tok * topk + _r] = wr;
slog[best as usize] = ninf; }
sync_cube();
}
if t == 0 {
for _r in 0..topk {
let cur = w_out[tok * topk + _r];
w_out[tok * topk + _r] = cur / wsum;
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn moe_route_run<R: Runtime>(
client: &ComputeClient<R>,
logits: &[f32],
ntok: usize,
n_experts: usize,
topk: usize,
nt: usize,
) -> (Vec<u32>, Vec<f32>) {
let lh = client.create_from_slice(f32::as_bytes(logits));
let ih = client.create_from_slice(u32::as_bytes(&vec![0u32; ntok * topk]));
let wh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; ntok * topk]));
unsafe {
moe_route::launch_unchecked::<f32, R>(
client,
Grid::Static(ntok as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(lh.clone(), logits.len()),
ArrayArg::from_raw_parts(ih.clone(), ntok * topk),
ArrayArg::from_raw_parts(wh.clone(), ntok * topk),
n_experts,
topk,
nt,
);
}
let ids = u32::from_bytes(&client.read_one_unchecked(ih)).to_vec();
let w = f32::from_bytes(&client.read_one_unchecked(wh)).to_vec();
(ids, w)
}
#[kernel(targets(cuda, metal, vulkan, webgpu, cpu), unchecked)]
pub fn gemv<F: Float>(
w: &Array<F>,
x: &Array<F>,
out: &mut Array<F>,
meta: &Array<u32>, #[comptime] nt: usize,
) {
let row = CUBE_POS as usize;
let t = UNIT_POS as usize;
let k = meta[0] as usize;
let wbase = row * k;
let mut acc = F::new(0.0);
let mut i = t;
while i < k {
acc += w[wbase + i] * x[i];
i += nt;
}
let mut smem = SharedMemory::<F>::new(nt);
smem[t] = acc;
sync_cube();
let mut stride = CUBE_DIM / 2;
while stride > 0 {
if UNIT_POS < stride {
let v = smem[(UNIT_POS + stride) as usize];
smem[t] += v;
}
sync_cube();
stride /= 2;
}
if t == 0 {
out[row] = smem[0];
}
}
pub fn gemv_run<R: Runtime>(
client: &ComputeClient<R>,
w: &[f32],
x: &[f32],
n: usize,
k: usize,
nt: usize,
) -> Vec<f32> {
let wh = client.create_from_slice(f32::as_bytes(w));
let xh = client.create_from_slice(f32::as_bytes(x));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; n]));
let mh = client.create_from_slice(u32::as_bytes(&[k as u32]));
unsafe {
gemv::launch_unchecked::<f32, R>(
client,
Grid::Static(n as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(wh.clone(), w.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(oh.clone(), n),
ArrayArg::from_raw_parts(mh.clone(), 1),
nt,
);
}
f32::from_bytes(&client.read_one_unchecked(oh)).to_vec()
}
pub fn moe_route_ref(
logits: &[f32],
ntok: usize,
n_experts: usize,
topk: usize,
) -> (Vec<u32>, Vec<f32>) {
let mut ids = vec![0u32; ntok * topk];
let mut w = vec![0.0f32; ntok * topk];
for tok in 0..ntok {
let row = &logits[tok * n_experts..(tok + 1) * n_experts];
let m = row.iter().cloned().fold(f32::MIN, f32::max);
let exps: Vec<f32> = row.iter().map(|&x| (x - m).exp()).collect();
let denom: f32 = exps.iter().sum();
let mut p: Vec<(usize, f32)> = exps.iter().map(|&e| e / denom).enumerate().collect();
p.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
let wsum: f32 = p[..topk].iter().map(|&(_, v)| v).sum();
for r in 0..topk {
ids[tok * topk + r] = p[r].0 as u32;
w[tok * topk + r] = p[r].1 / wsum;
}
}
(ids, w)
}
#[cfg(test)]
mod tests {
use super::*;
fn max_rel(a: &[f32], b: &[f32]) -> f32 {
a.iter()
.zip(b)
.map(|(x, y)| (x - y).abs() / x.abs().max(1e-6))
.fold(0.0, f32::max)
}
fn gen_dp4a(rows: usize, k: usize) -> (Vec<i8>, Vec<i8>, Vec<f32>) {
let mut s = 0xA5A5_1234_9E37_79B9u64;
let mut nx = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
let wq: Vec<i8> = (0..rows * k).map(|_| (nx() % 251) as i8).collect();
let xq: Vec<i8> = (0..k).map(|_| (nx() % 251) as i8).collect();
let wd: Vec<f32> = (0..rows * (k / 32))
.map(|_| (nx() % 1000) as f32 / 8000.0 + 0.01)
.collect();
(wq, xq, wd)
}
#[cfg(feature = "cpu")]
#[test]
fn matvec_q8_dp4a_autotune_cpu_bit_exact_and_cached() {
use crate::tune::Tuner;
use cubecl::cpu::{CpuDevice, CpuRuntime};
let (rows, k) = (48usize, 256usize); let (wq, xq, wd) = gen_dp4a(rows, k);
let wq32: Vec<i32> = wq.iter().map(|&v| v as i32).collect();
let xq32: Vec<i32> = xq.iter().map(|&v| v as i32).collect();
let want = matvec_q8_dp4a_ref(&wq32, &xq32, &wd, rows, k);
let c = CpuRuntime::client(&CpuDevice::default());
let mut ref_bits: Option<Vec<u32>> = None;
for &(vw, block, nr) in &[
(1usize, 64u32, 1usize),
(4, 64, 1),
(2, 128, 1),
(8, 256, 1),
(1, 64, 2),
(4, 128, 2),
(2, 256, 4),
] {
let (got, _ms) = matvec_q8_dp4a_tuned_bench::<CpuRuntime>(
&c, &wq, &xq, &wd, rows, k, vw, block, nr, 1,
);
let gbits: Vec<u32> = got.iter().map(|v| v.to_bits()).collect();
match &ref_bits {
None => ref_bits = Some(gbits),
Some(rb) => assert_eq!(
&gbits, rb,
"wg={block} vw={vw} nr={nr} not byte-identical to b64_v1_r1"
),
}
let rel = max_rel(&want, &got);
eprintln!("[matvec_q8_dp4a_tuned CPU] wg={block} vw={vw} nr={nr} max_rel={rel:.2e}");
assert!(
rel < 2e-3,
"dp4a tuned wg={block} vw={vw} nr={nr} max_rel {rel} vs oracle"
);
}
let dir = std::env::temp_dir().join(format!(
"hk-tune-dp4a-{}-{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
let tuner = Tuner::new(&dir);
let p = matvec_q8_dp4a_tuned_set::<CpuRuntime>(&c, &wq, &xq, &wd, rows, k)
.pick_with(&tuner, "cpu");
assert!(
!p.from_cache && p.benched == 4,
"first call must benchmark all 4 variants"
);
assert!(
max_rel(&want, &p.output) < 2e-3,
"autotuned winner not within tolerance"
);
eprintln!(
"[matvec_q8_dp4a autotune CPU] winner={} timings={:?}",
p.winner, p.timings
);
let p2 = matvec_q8_dp4a_tuned_set::<CpuRuntime>(&c, &wq, &xq, &wd, rows, k)
.pick_with(&tuner, "cpu");
assert!(
p2.from_cache && p2.benched == 0,
"second call must hit the cache and skip timing"
);
assert_eq!(p2.winner, p.winner, "cache returned a different winner");
std::fs::remove_dir_all(&dir).ok();
}
#[cfg(feature = "cpu")]
#[test]
fn matvec_dp4a_nr_ragged_cpu() {
use cubecl::cpu::{CpuDevice, CpuRuntime};
let (rows, k) = (49usize, 256usize); let (wq, xq, wd) = gen_dp4a(rows, k);
let wq32: Vec<i32> = wq.iter().map(|&v| v as i32).collect();
let xq32: Vec<i32> = xq.iter().map(|&v| v as i32).collect();
let want = matvec_q8_dp4a_ref(&wq32, &xq32, &wd, rows, k);
let c = CpuRuntime::client(&CpuDevice::default());
let mut ref_bits: Option<Vec<u32>> = None;
for &(vw, block, nr) in &[
(2usize, 64u32, 1usize),
(2, 64, 2),
(2, 128, 4),
(4, 256, 4),
] {
let (got, _ms) = matvec_q8_dp4a_tuned_bench::<CpuRuntime>(
&c, &wq, &xq, &wd, rows, k, vw, block, nr, 1,
);
assert_eq!(got.len(), rows, "output truncated at nr={nr}");
let gbits: Vec<u32> = got.iter().map(|v| v.to_bits()).collect();
match &ref_bits {
None => ref_bits = Some(gbits),
Some(rb) => assert_eq!(&gbits, rb, "ragged nr={nr} not byte-identical to nr=1"),
}
assert!(
max_rel(&want, &got) < 2e-3,
"ragged nr={nr} diverged from the oracle"
);
}
}
#[cfg(feature = "cpu")]
#[test]
fn matvec_dp4a_space_search_cpu() {
use cubecl::cpu::{CpuDevice, CpuRuntime};
let (rows, k) = (64usize, 256usize); let (_, xq, wd) = gen_dp4a(rows, k);
let banks: Vec<Vec<i8>> = (0..3)
.map(|b| {
let mut s = 0x51ED_2701u64.wrapping_add(b).wrapping_mul(0x9E37_79B1) | 1;
(0..rows * k)
.map(|_| {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
(s % 251) as i8
})
.collect()
})
.collect();
let space = matvec_dp4a_space(k);
let feasible = space.enumerate();
assert!(!feasible.is_empty(), "empty feasible space");
assert!(feasible
.iter()
.all(|cfg| (k / 4) % (cfg.get(&space, "VW") as usize) == 0));
assert!(feasible.iter().all(|cfg| {
let w = cfg.get(&space, "WG");
w > 0 && w <= 1024
}));
let c = CpuRuntime::client(&CpuDevice::default());
let eval = MatvecDp4aEval::new(&c, &space, &banks, &xq, &wd, rows, k, 1);
let dir = std::env::temp_dir().join(format!(
"hk-dp4a-hunt-{}-{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
let tuner = Tuner::new(&dir);
let evo = Evolution::new()
.population(8)
.generations(3)
.measure_iters(1);
let r = matvec_dp4a_hunt(&tuner, "cpu", rows, k, &eval, &evo, 0xC0FFEE);
assert!(!r.from_cache, "first hunt must not be a cache hit");
let rep = r
.report
.as_ref()
.expect("a miss carries the evidence trail");
assert!(rep.best_ms.is_finite(), "no measurable winner");
let win = space
.parse(&r.winner)
.expect("winner is a valid config name");
assert!(space.feasible(&win), "winner {} is infeasible", r.winner);
assert!(
eval.worst_rel() < 2e-3,
"a measured variant diverged from the oracle: {:.2e}",
eval.worst_rel()
);
eprintln!(
"[dp4a hunt CPU] winner={} evaluated={} worst_rel={:.2e}",
r.winner,
rep.evaluated,
eval.worst_rel()
);
let r2 = matvec_dp4a_hunt(&tuner, "cpu", rows, k, &eval, &evo, 0xC0FFEE);
assert!(
r2.from_cache && r2.report.is_none(),
"second hunt must hit the cache"
);
assert_eq!(r2.winner, r.winner, "cache returned a different winner");
std::fs::remove_dir_all(&dir).ok();
}
#[cfg(feature = "cpu")]
#[test]
fn matvec_q4k_f32_blk_cpu_bit_exact() {
use cubecl::cpu::{CpuDevice, CpuRuntime};
let (rows, k) = (35usize, 512usize); let (wqs, wsc, wd, wdm, x) = gen_q4k(rows, k);
let want = matvec_q4k_ref(&wqs, &wsc, &wd, &wdm, &x, rows, k);
let maxref = want.iter().fold(0f32, |a, &v| a.max(v.abs())).max(1e-30);
let c = CpuRuntime::client(&CpuDevice::default());
for &(nt, nr) in &[
(16usize, 1usize),
(32, 1),
(64, 1),
(16, 2),
(32, 4),
(64, 2),
] {
let got = matvec_q4k_f32_blk_run::<CpuRuntime>(
&c, &wqs, &wsc, &wd, &wdm, &x, rows, k, nt, nr,
);
assert_eq!(got.len(), rows, "output truncated at nt={nt} nr={nr}");
let rel = got
.iter()
.zip(&want)
.map(|(a, b)| (a - b).abs())
.fold(0f32, f32::max)
/ maxref;
eprintln!("[matvec_q4k_f32_blk CPU] nt={nt} nr={nr} scale_rel={rel:.2e}");
assert!(
rel < 1e-4,
"f32-direct nt={nt} nr={nr} scale_rel {rel} vs Q4_K oracle"
);
}
}
#[cfg(feature = "cpu")]
#[test]
fn matvec_q4k_f32_search_cpu() {
use cubecl::cpu::{CpuDevice, CpuRuntime};
let (rows, k) = (48usize, 512usize);
let (wqs, wsc, wd, wdm, x) = gen_q4k(rows, k);
let space = matvec_q4k_f32_space();
let feasible = space.enumerate();
assert!(!feasible.is_empty());
assert!(feasible.iter().all(|cfg| {
let w = cfg.get(&space, "WG");
w >= 2 && (w & (w - 1)) == 0
}));
let c = CpuRuntime::client(&CpuDevice::default());
let eval = MatvecQ4kF32Eval::new(&c, &space, &wqs, &wsc, &wd, &wdm, &x, rows, k, 2, 1);
let dir = std::env::temp_dir().join(format!(
"hk-q4kf32-hunt-{}-{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
let tuner = Tuner::new(&dir);
let evo = Evolution::new()
.population(8)
.generations(3)
.measure_iters(1);
let r = matvec_q4k_f32_hunt(&tuner, "cpu", rows, k, &eval, &evo, 0xF32);
assert!(!r.from_cache, "first hunt must not be a cache hit");
let rep = r
.report
.as_ref()
.expect("a miss carries the evidence trail");
assert!(rep.best_ms.is_finite(), "no measurable winner");
let win = space
.parse(&r.winner)
.expect("winner is a valid config name");
assert!(space.feasible(&win), "winner {} is infeasible", r.winner);
assert!(
eval.worst_rel() < 1e-3,
"a measured variant diverged from the oracle: {:.2e}",
eval.worst_rel()
);
eprintln!(
"[q4k_f32 hunt CPU] winner={} evaluated={} worst_rel={:.2e}",
r.winner,
rep.evaluated,
eval.worst_rel()
);
let r2 = matvec_q4k_f32_hunt(&tuner, "cpu", rows, k, &eval, &evo, 0xF32);
assert!(
r2.from_cache && r2.report.is_none(),
"second hunt must hit the cache"
);
assert_eq!(r2.winner, r.winner, "cache returned a different winner");
std::fs::remove_dir_all(&dir).ok();
}
#[cfg(feature = "cpu")]
#[test]
fn moe_matvec_q4k_cpu_bit_exact() {
use cubecl::cpu::{CpuDevice, CpuRuntime};
let (e, n, slots, k) = (8usize, 12usize, 5usize, 512usize);
let (wqs, wsc, wd, wdm, x, ids) = gen_moe_q4k(e, n, slots, k);
let c = CpuRuntime::client(&CpuDevice::default());
let got =
moe_matvec_q4k_run::<CpuRuntime>(&c, &wqs, &wsc, &wd, &wdm, &x, &ids, slots, n, k);
let want = moe_matvec_q4k_ref(&wqs, &wsc, &wd, &wdm, &x, &ids, slots, n, k);
let gbits: Vec<u32> = got.iter().map(|v| v.to_bits()).collect();
let wbits: Vec<u32> = want.iter().map(|v| v.to_bits()).collect();
assert_eq!(
gbits, wbits,
"moe_matvec_q4k CPU kernel != oracle bit-exact"
);
}
#[cfg(feature = "cpu")]
#[test]
fn moe_matvec_q6k_cpu_bit_exact() {
use cubecl::cpu::{CpuDevice, CpuRuntime};
let (e, n, slots, k) = (8usize, 12usize, 5usize, 512usize);
let (wql, wqh, wsc, wd, x, ids) = gen_moe_q6k(e, n, slots, k);
let c = CpuRuntime::client(&CpuDevice::default());
let got =
moe_matvec_q6k_run::<CpuRuntime>(&c, &wql, &wqh, &wsc, &wd, &x, &ids, slots, n, k);
let want = moe_matvec_q6k_ref(&wql, &wqh, &wsc, &wd, &x, &ids, slots, n, k);
let gbits: Vec<u32> = got.iter().map(|v| v.to_bits()).collect();
let wbits: Vec<u32> = want.iter().map(|v| v.to_bits()).collect();
assert_eq!(
gbits, wbits,
"moe_matvec_q6k CPU kernel != oracle bit-exact"
);
}
}