use crate::prelude::*;
use crate::tune::{Config, Evaluator, Evolution, Evolved, Space, Tuner, Verdict};
use cubecl::server::Handle;
#[kernel(targets(cuda, rocm, vulkan, metal), unchecked)]
pub fn wmma_hello_i8(a: &Array<i8>, b: &Array<i8>, out: &mut Array<i32>) {
let ma = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::A,
16usize,
16usize,
16usize,
cmma::MatrixLayout::RowMajor,
&a.to_slice(),
16,
);
let mb = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::B,
16usize,
16usize,
16usize,
cmma::MatrixLayout::ColMajor,
&b.to_slice(),
16,
);
let mc = cmma::Matrix::<i32>::from_value(
cmma::MatrixIdent::Accumulator,
16usize,
16usize,
16usize,
cmma::MatrixLayout::Undefined,
0i32,
);
cmma::execute::<i8, i8, i32, i32>(&ma, &mb, &mc, &mc);
cmma::store(
&mut out.to_slice_mut(),
&mc,
16,
cmma::MatrixLayout::RowMajor,
);
}
pub fn wmma_hello_i8_run<R: Runtime>(client: &ComputeClient<R>, a: &[i8], b: &[i8]) -> Vec<i32> {
let plane = client.properties().hardware.plane_size_max;
let ah = client.create_from_slice(i8::as_bytes(a));
let bh = client.create_from_slice(i8::as_bytes(b));
let oh = client.create_from_slice(i32::as_bytes(&vec![0i32; 256]));
unsafe {
wmma_hello_i8::launch_unchecked::<R>(
client,
Grid::Static(1, 1, 1),
Block::new_1d(plane),
ArrayArg::from_raw_parts(ah.clone(), 256),
ArrayArg::from_raw_parts(bh.clone(), 256),
ArrayArg::from_raw_parts(oh.clone(), 256),
);
}
i32::from_bytes(&client.read_one_unchecked(oh)).to_vec()
}
pub fn hello_i8_ref(a: &[i8], b: &[i8]) -> Vec<i32> {
let mut out = vec![0i32; 256];
for m in 0..16 {
for n in 0..16 {
let mut acc = 0i32;
for k in 0..16 {
acc += a[m * 16 + k] as i32 * b[n * 16 + k] as i32; }
out[m * 16 + n] = acc;
}
}
out
}
#[kernel(targets(cuda, rocm), unchecked)]
pub fn mma_hello_i8(
a: &Array<i8>, b: &Array<i8>, out: &mut Array<i32>, #[comptime] size_m: usize,
#[comptime] size_n: usize,
#[comptime] size_k: usize,
) {
let def = cmma::MmaDefinition::<i8, i8, i32>::new(size_m, size_n, size_k);
let lane_id = UNIT_POS_PLANE;
let vector_size_a = def.vector_size(cmma::MatrixIdent::A);
let size!(NA) = vector_size_a;
let vector_count_a = def.vectors_per_lane(cmma::MatrixIdent::A);
let mut registers_a = Array::<Vector<i8, NA>>::new(vector_count_a);
let vector_size_b = def.vector_size(cmma::MatrixIdent::B);
let size!(NB) = vector_size_b;
let vector_count_b = def.vectors_per_lane(cmma::MatrixIdent::B);
let mut registers_b = Array::<Vector<i8, NB>>::new(vector_count_b);
let vector_size_c = def.vector_size(cmma::MatrixIdent::Accumulator);
let size!(NC) = vector_size_c;
let vector_count_c = def.vectors_per_lane(cmma::MatrixIdent::Accumulator);
let mut registers_c = Array::<Vector<i32, NC>>::new(vector_count_c);
#[unroll]
for i in 0..vector_count_a {
let mut reg = Vector::<i8, NA>::empty();
#[unroll]
for kk in 0..vector_size_a {
let n_elem = i * vector_size_a + kk;
let (row, col) = def.position_of_nth(lane_id, n_elem as u32, cmma::MatrixIdent::A);
reg[kk] = a[(row * size_k as u32 + col) as usize];
}
registers_a[i] = reg;
}
#[unroll]
for i in 0..vector_count_b {
let mut reg = Vector::<i8, NB>::empty();
#[unroll]
for kk in 0..vector_size_b {
let n_elem = i * vector_size_b + kk;
let (row, col) = def.position_of_nth(lane_id, n_elem as u32, cmma::MatrixIdent::B);
reg[kk] = b[(row * size_n as u32 + col) as usize];
}
registers_b[i] = reg;
}
#[unroll]
for i in 0..vector_count_c {
let mut reg = Vector::<i32, NC>::empty();
#[unroll]
for kk in 0..vector_size_c {
reg[kk] = 0i32;
}
registers_c[i] = reg;
}
let registers_d = def.execute(®isters_a, ®isters_b, ®isters_c);
#[unroll]
for i in 0..vector_count_c {
let reg = registers_d[i];
#[unroll]
for kk in 0..vector_size_c {
let n_elem = i * vector_size_c + kk;
let (row, col) =
def.position_of_nth(lane_id, n_elem as u32, cmma::MatrixIdent::Accumulator);
out[(row * size_n as u32 + col) as usize] = reg[kk];
}
}
}
pub fn mma_hello_i8_run<R: Runtime>(
client: &ComputeClient<R>,
a: &[i8],
b: &[i8],
plane: u32,
) -> Vec<i32> {
let ah = client.create_from_slice(i8::as_bytes(a));
let bh = client.create_from_slice(i8::as_bytes(b));
let oh = client.create_from_slice(i32::as_bytes(&vec![0i32; 16 * 8]));
unsafe {
mma_hello_i8::launch_unchecked::<R>(
client,
Grid::Static(1, 1, 1),
Block::new_1d(plane),
ArrayArg::from_raw_parts(ah.clone(), 16 * 32),
ArrayArg::from_raw_parts(bh.clone(), 32 * 8),
ArrayArg::from_raw_parts(oh.clone(), 16 * 8),
16,
8,
32,
);
}
i32::from_bytes(&client.read_one_unchecked(oh)).to_vec()
}
pub fn mma_hello_ref(a: &[i8], b: &[i8]) -> Vec<i32> {
let (m, n, k) = (16usize, 8usize, 32usize);
let mut out = vec![0i32; m * n];
for i in 0..m {
for j in 0..n {
let mut acc = 0i32;
for l in 0..k {
acc += a[i * k + l] as i32 * b[l * n + j] as i32;
}
out[i * n + j] = acc;
}
}
out
}
#[kernel(targets(cuda, rocm, vulkan, metal, cpu), unchecked)]
pub fn mmq_q8_wmma(
xq: &Array<i8>,
xs: &Array<f32>,
wq: &Array<i8>,
wd: &Array<f32>,
out: &mut Array<f32>,
#[comptime] m: usize,
#[comptime] n: usize,
#[comptime] k: usize,
#[comptime] plane: usize,
#[comptime] target: Target,
) {
let nt = CUBE_POS_X as usize; let mt = CUBE_POS_Y as usize; let lane = UNIT_POS as usize; let kb_count = k / 32;
let mrow0 = mt * 16;
let ncol0 = nt * 16;
let per = 256 / plane;
let mut accf = SharedMemory::<f32>::new(256usize); let mut ci = SharedMemory::<i32>::new(256usize);
#[unroll]
for e in 0usize..per {
accf[lane * per + e] = 0.0f32;
}
sync_cube();
for kb in 0..kb_count {
let k0 = kb * 32;
island! {
cuda | rocm | vulkan | metal => {
let c = cmma::Matrix::<i32>::from_value(
cmma::MatrixIdent::Accumulator,
16usize,
16usize,
16usize,
cmma::MatrixLayout::Undefined,
0i32,
);
let a0 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::A,
16usize,
16usize,
16usize,
cmma::MatrixLayout::RowMajor,
&xq.slice(mrow0 * k + k0, m * k),
k as u32,
);
let b0 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::B,
16usize,
16usize,
16usize,
cmma::MatrixLayout::ColMajor,
&wq.slice(ncol0 * k + k0, n * k),
k as u32,
);
cmma::execute::<i8, i8, i32, i32>(&a0, &b0, &c, &c);
let a1 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::A,
16usize,
16usize,
16usize,
cmma::MatrixLayout::RowMajor,
&xq.slice(mrow0 * k + k0 + 16, m * k),
k as u32,
);
let b1 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::B,
16usize,
16usize,
16usize,
cmma::MatrixLayout::ColMajor,
&wq.slice(ncol0 * k + k0 + 16, n * k),
k as u32,
);
cmma::execute::<i8, i8, i32, i32>(&a1, &b1, &c, &c);
cmma::store(&mut ci.to_slice_mut(), &c, 16, cmma::MatrixLayout::RowMajor);
}
default => {
#[unroll]
for e in 0usize..per {
let idx = lane * per + e;
let mm = idx / 16;
let nn = idx % 16;
let mut isum = 0i32;
for l in 0usize..32 {
isum += i32::cast_from(xq[(mrow0 + mm) * k + k0 + l])
* i32::cast_from(wq[(ncol0 + nn) * k + k0 + l]);
}
ci[idx] = isum;
}
}
};
sync_cube();
#[unroll]
for e in 0usize..per {
let idx = lane * per + e;
let mm = idx / 16;
let nn = idx % 16;
let xsc = xs[(mrow0 + mm) * kb_count + kb];
let wsc = wd[(ncol0 + nn) * kb_count + kb];
accf[idx] += f32::cast_from(ci[idx]) * xsc * wsc;
}
sync_cube();
}
#[unroll]
for e in 0usize..per {
let idx = lane * per + e;
let mm = idx / 16;
let nn = idx % 16;
out[(mrow0 + mm) * n + (ncol0 + nn)] = accf[idx];
}
}
#[allow(clippy::too_many_arguments)]
pub fn mmq_q8_wmma_run<R: Runtime>(
client: &ComputeClient<R>,
xq: &[i8],
xs: &[f32],
wq: &[i8],
wd: &[f32],
m: usize,
n: usize,
k: usize,
iters: usize,
) -> (Vec<f32>, f64) {
mmq_q8_wmma_run_with(client, xq, xs, wq, wd, m, n, k, iters, Target::of(client))
}
#[allow(clippy::too_many_arguments)]
pub fn mmq_q8_wmma_run_with<R: Runtime>(
client: &ComputeClient<R>,
xq: &[i8],
xs: &[f32],
wq: &[i8],
wd: &[f32],
m: usize,
n: usize,
k: usize,
iters: usize,
target: Target,
) -> (Vec<f32>, f64) {
let plane = client.properties().hardware.plane_size_max;
let xqh = client.create_from_slice(i8::as_bytes(xq));
let xsh = client.create_from_slice(f32::as_bytes(xs));
let wqh = client.create_from_slice(i8::as_bytes(wq));
let wdh = client.create_from_slice(f32::as_bytes(wd));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; m * n]));
let grid = Grid::Static((n / 16) as u32, (m / 16) as u32, 1);
let launch = |c: &ComputeClient<R>| unsafe {
mmq_q8_wmma::launch_unchecked::<R>(
c,
grid.clone(),
Block::new_1d(plane),
ArrayArg::from_raw_parts(xqh.clone(), xq.len()),
ArrayArg::from_raw_parts(xsh.clone(), xs.len()),
ArrayArg::from_raw_parts(wqh.clone(), wq.len()),
ArrayArg::from_raw_parts(wdh.clone(), wd.len()),
ArrayArg::from_raw_parts(oh.clone(), m * n),
m,
n,
k,
plane as usize,
target,
);
};
launch(client);
let out = f32::from_bytes(&client.read_one_unchecked(oh.clone())).to_vec();
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);
let ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
(out, ms)
}
pub fn mmq_q8_ref(
xq: &[i8],
xs: &[f32],
wq: &[i8],
wd: &[f32],
m: usize,
n: usize,
k: usize,
) -> Vec<f32> {
let kb = k / 32;
let mut out = vec![0.0f32; m * n];
for i in 0..m {
for j in 0..n {
let mut acc = 0.0f32;
for b in 0..kb {
let mut isum = 0i32;
for l in 0..32 {
isum += xq[i * k + b * 32 + l] as i32 * wq[j * k + b * 32 + l] as i32;
}
acc += isum as f32 * xs[i * kb + b] * wd[j * kb + b];
}
out[i * n + j] = acc;
}
}
out
}
#[kernel(targets(cuda, rocm, vulkan, metal, cpu), unchecked)]
pub fn mmq_q8_wmma_blk(
xq: &Array<i8>,
xs: &Array<f32>,
wq: &Array<i8>,
wd: &Array<f32>,
out: &mut Array<f32>,
#[comptime] n: usize,
#[comptime] k: usize,
#[comptime] plane: usize,
#[comptime] target: Target,
) {
let tid = UNIT_POS as usize;
let warp = tid / plane; let lane = tid % plane; let wm = warp / 4; let wn = warp % 4; let mrow0 = CUBE_POS_Y as usize * 32;
let ncol0 = CUBE_POS_X as usize * 64;
let kb_count = k / 32;
let nthread = 8 * plane; let per = 256 / plane;
let mut sa = SharedMemory::<i8>::new(1024usize); let mut sb = SharedMemory::<i8>::new(2048usize); let mut ci = SharedMemory::<i32>::new(2048usize); let mut accf = SharedMemory::<f32>::new(2048usize);
for e in 0usize..(2048 / nthread) {
accf[tid * (2048 / nthread) + e] = 0.0f32;
}
sync_cube();
for kb in 0..kb_count {
let k0 = kb * 32;
for i in 0usize..(1024 / nthread) {
let idx = tid + i * nthread;
sa[idx] = xq[(mrow0 + idx / 32) * k + k0 + idx % 32];
}
for i in 0usize..(2048 / nthread) {
let idx = tid + i * nthread;
sb[idx] = wq[(ncol0 + idx / 32) * k + k0 + idx % 32];
}
sync_cube();
island! {
cuda | rocm | vulkan | metal => {
let c = cmma::Matrix::<i32>::from_value(
cmma::MatrixIdent::Accumulator,
16usize,
16usize,
16usize,
cmma::MatrixLayout::Undefined,
0i32,
);
let a0 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::A,
16usize,
16usize,
16usize,
cmma::MatrixLayout::RowMajor,
&sa.slice(wm * 512, 1024),
32,
);
let b0 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::B,
16usize,
16usize,
16usize,
cmma::MatrixLayout::ColMajor,
&sb.slice(wn * 512, 2048),
32,
);
cmma::execute::<i8, i8, i32, i32>(&a0, &b0, &c, &c);
let a1 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::A,
16usize,
16usize,
16usize,
cmma::MatrixLayout::RowMajor,
&sa.slice(wm * 512 + 16, 1024),
32,
);
let b1 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::B,
16usize,
16usize,
16usize,
cmma::MatrixLayout::ColMajor,
&sb.slice(wn * 512 + 16, 2048),
32,
);
cmma::execute::<i8, i8, i32, i32>(&a1, &b1, &c, &c);
cmma::store(
&mut ci.slice_mut(warp * 256, warp * 256 + 256),
&c,
16,
cmma::MatrixLayout::RowMajor,
);
}
default => {
for e in 0usize..per {
let p = lane * per + e;
let smm = p / 16;
let snn = p % 16;
let mut isum = 0i32;
for l in 0usize..32 {
isum += i32::cast_from(sa[(wm * 16 + smm) * 32 + l])
* i32::cast_from(sb[(wn * 16 + snn) * 32 + l]);
}
ci[warp * 256 + p] = isum;
}
}
};
sync_cube();
for e in 0usize..per {
let p = lane * per + e;
let smm = p / 16;
let snn = p % 16;
let gmm = wm * 16 + smm;
let gnn = wn * 16 + snn;
let xsc = xs[(mrow0 + gmm) * kb_count + kb];
let wsc = wd[(ncol0 + gnn) * kb_count + kb];
accf[gmm * 64 + gnn] += f32::cast_from(ci[warp * 256 + p]) * xsc * wsc;
}
sync_cube();
}
for e in 0usize..per {
let p = lane * per + e;
let gmm = wm * 16 + p / 16;
let gnn = wn * 16 + p % 16;
out[(mrow0 + gmm) * n + (ncol0 + gnn)] = accf[gmm * 64 + gnn];
}
}
#[allow(clippy::too_many_arguments)]
pub fn mmq_q8_wmma_blk_run<R: Runtime>(
client: &ComputeClient<R>,
xq: &[i8],
xs: &[f32],
wq: &[i8],
wd: &[f32],
m: usize,
n: usize,
k: usize,
iters: usize,
) -> (Vec<f32>, f64) {
mmq_q8_wmma_blk_run_with(client, xq, xs, wq, wd, m, n, k, iters, Target::of(client))
}
#[allow(clippy::too_many_arguments)]
pub fn mmq_q8_wmma_blk_run_with<R: Runtime>(
client: &ComputeClient<R>,
xq: &[i8],
xs: &[f32],
wq: &[i8],
wd: &[f32],
m: usize,
n: usize,
k: usize,
iters: usize,
target: Target,
) -> (Vec<f32>, f64) {
let plane = client.properties().hardware.plane_size_max;
let xqh = client.create_from_slice(i8::as_bytes(xq));
let xsh = client.create_from_slice(f32::as_bytes(xs));
let wqh = client.create_from_slice(i8::as_bytes(wq));
let wdh = client.create_from_slice(f32::as_bytes(wd));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; m * n]));
let grid = Grid::Static((n / 64) as u32, (m / 32) as u32, 1);
let launch = |c: &ComputeClient<R>| unsafe {
mmq_q8_wmma_blk::launch_unchecked::<R>(
c,
grid.clone(),
Block::new_1d(8 * plane),
ArrayArg::from_raw_parts(xqh.clone(), xq.len()),
ArrayArg::from_raw_parts(xsh.clone(), xs.len()),
ArrayArg::from_raw_parts(wqh.clone(), wq.len()),
ArrayArg::from_raw_parts(wdh.clone(), wd.len()),
ArrayArg::from_raw_parts(oh.clone(), m * n),
n,
k,
plane as usize,
target,
);
};
launch(client);
let out = f32::from_bytes(&client.read_one_unchecked(oh.clone())).to_vec();
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);
let ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
(out, ms)
}
pub fn gen_mmq(m: usize, n: usize, k: usize) -> (Vec<i8>, Vec<f32>, Vec<i8>, Vec<f32>) {
let kb = k / 32;
let mut s = 0xD1B54A32D192ED03u64;
let mut next = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
let xq: Vec<i8> = (0..m * k)
.map(|_| ((next() % 255) as i64 - 127) as i8)
.collect();
let wq: Vec<i8> = (0..n * k)
.map(|_| ((next() % 255) as i64 - 127) as i8)
.collect();
let xs: Vec<f32> = (0..m * kb)
.map(|_| (next() % 1000) as f32 / 50000.0 + 0.002)
.collect();
let wd: Vec<f32> = (0..n * kb)
.map(|_| (next() % 1000) as f32 / 50000.0 + 0.002)
.collect();
(xq, xs, wq, wd)
}
#[device]
fn q4k_byte(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 = q4k_byte(wsc, scbase, j) & 63;
if j >= 4 {
r = (q4k_byte(wsc, scbase, j + 4) & 15) | ((q4k_byte(wsc, scbase, j - 4) >> 6) << 4);
}
r
}
#[device]
fn q4k_m(wsc: &Array<u32>, scbase: usize, j: usize) -> u32 {
let mut r = q4k_byte(wsc, scbase, j + 4) & 63;
if j >= 4 {
r = (q4k_byte(wsc, scbase, j + 4) >> 4) | ((q4k_byte(wsc, scbase, j) >> 6) << 4);
}
r
}
#[kernel(targets(cuda, rocm, vulkan, metal, cpu), unchecked)]
pub fn mmq_q4k_wmma_blk(
xq: &Array<i8>,
xs: &Array<f32>,
xsum: &Array<f32>, wqs: &Array<u32>, wsc: &Array<u32>, wd: &Array<f32>, wdm: &Array<f32>, out: &mut Array<f32>,
#[comptime] n: usize,
#[comptime] k: usize,
#[comptime] plane: usize,
#[comptime] target: Target,
) {
let tid = UNIT_POS as usize;
let warp = tid / plane;
let lane = tid % plane;
let wm = warp / 4; let wn = warp % 4; let mrow0 = CUBE_POS_Y as usize * 32;
let ncol0 = CUBE_POS_X as usize * 64;
let kb_count = k / 32;
let nsb = k / 256; let nthread = 8 * plane;
let per = 256 / plane;
let mut sa = SharedMemory::<i8>::new(1024usize);
let mut sb = SharedMemory::<i8>::new(2048usize);
let mut ci = SharedMemory::<i32>::new(2048usize);
let mut accf = SharedMemory::<f32>::new(2048usize);
for e in 0usize..(2048 / nthread) {
accf[tid * (2048 / nthread) + e] = 0.0f32;
}
sync_cube();
for kb in 0..kb_count {
let k0 = kb * 32;
for i in 0usize..(1024 / nthread) {
let idx = tid + i * nthread;
sa[idx] = xq[(mrow0 + idx / 32) * k + k0 + idx % 32];
}
let is = kb % 8;
let g = is / 2;
let sbk = kb / 8; for i in 0usize..(2048 / nthread) {
let idx = tid + i * nthread;
let nrow = ncol0 + idx / 32;
let qi = idx % 32;
let qb = q4k_byte(wqs, (nrow * nsb + sbk) * 32, g * 32 + qi);
let nib = (qb >> (4u32 * (is % 2) as u32)) & 15;
sb[idx] = i8::cast_from(nib);
}
sync_cube();
island! {
cuda | rocm | vulkan | metal => {
let c = cmma::Matrix::<i32>::from_value(
cmma::MatrixIdent::Accumulator, 16usize, 16usize, 16usize,
cmma::MatrixLayout::Undefined, 0i32,
);
let a0 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::A, 16usize, 16usize, 16usize,
cmma::MatrixLayout::RowMajor, &sa.slice(wm * 512, 1024), 32,
);
let b0 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::B, 16usize, 16usize, 16usize,
cmma::MatrixLayout::ColMajor, &sb.slice(wn * 512, 2048), 32,
);
cmma::execute::<i8, i8, i32, i32>(&a0, &b0, &c, &c);
let a1 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::A, 16usize, 16usize, 16usize,
cmma::MatrixLayout::RowMajor, &sa.slice(wm * 512 + 16, 1024), 32,
);
let b1 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::B, 16usize, 16usize, 16usize,
cmma::MatrixLayout::ColMajor, &sb.slice(wn * 512 + 16, 2048), 32,
);
cmma::execute::<i8, i8, i32, i32>(&a1, &b1, &c, &c);
cmma::store(&mut ci.slice_mut(warp * 256, warp * 256 + 256), &c, 16, cmma::MatrixLayout::RowMajor);
}
default => {
for e in 0usize..per {
let p = lane * per + e;
let smm = p / 16;
let snn = p % 16;
let mut isum = 0i32;
for l in 0usize..32 {
isum += i32::cast_from(sa[(wm * 16 + smm) * 32 + l])
* i32::cast_from(sb[(wn * 16 + snn) * 32 + l]);
}
ci[warp * 256 + p] = isum;
}
}
};
sync_cube();
for e in 0usize..per {
let p = lane * per + e;
let smm = p / 16;
let snn = p % 16;
let gmm = wm * 16 + smm;
let gnn = wn * 16 + snn;
let nrow = ncol0 + gnn;
let blk = nrow * nsb + sbk;
let scbase = blk * 3;
let wdd = wd[blk] * f32::cast_from(q4k_sc(wsc, scbase, is));
let wmm = wdm[blk] * f32::cast_from(q4k_m(wsc, scbase, is));
let xsc = xs[(mrow0 + gmm) * kb_count + kb];
let xsm = xsum[(mrow0 + gmm) * kb_count + kb];
accf[gmm * 64 + gnn] += xsc * wdd * f32::cast_from(ci[warp * 256 + p]) - wmm * xsm;
}
sync_cube();
}
for e in 0usize..per {
let p = lane * per + e;
let gmm = wm * 16 + p / 16;
let gnn = wn * 16 + p % 16;
out[(mrow0 + gmm) * n + (ncol0 + gnn)] = accf[gmm * 64 + gnn];
}
}
#[allow(clippy::too_many_arguments)]
pub fn mmq_q4k_wmma_blk_run<R: Runtime>(
client: &ComputeClient<R>,
xq: &[i8],
xs: &[f32],
xsum: &[f32],
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
m: usize,
n: usize,
k: usize,
iters: usize,
) -> (Vec<f32>, f64) {
let target = Target::of(client);
let plane = client.properties().hardware.plane_size_max;
let xqh = client.create_from_slice(i8::as_bytes(xq));
let xsh = client.create_from_slice(f32::as_bytes(xs));
let xsumh = client.create_from_slice(f32::as_bytes(xsum));
let wqsh = client.create_from_slice(u32::as_bytes(wqs));
let wsch = client.create_from_slice(u32::as_bytes(wsc));
let wdh = client.create_from_slice(f32::as_bytes(wd));
let wdmh = client.create_from_slice(f32::as_bytes(wdm));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; m * n]));
let grid = Grid::Static((n / 64) as u32, (m / 32) as u32, 1);
let launch = |c: &ComputeClient<R>| unsafe {
mmq_q4k_wmma_blk::launch_unchecked::<R>(
c,
grid.clone(),
Block::new_1d(8 * plane),
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(wqsh.clone(), wqs.len()),
ArrayArg::from_raw_parts(wsch.clone(), wsc.len()),
ArrayArg::from_raw_parts(wdh.clone(), wd.len()),
ArrayArg::from_raw_parts(wdmh.clone(), wdm.len()),
ArrayArg::from_raw_parts(oh.clone(), m * n),
n,
k,
plane as usize,
target,
);
};
launch(client);
let out = f32::from_bytes(&client.read_one_unchecked(oh.clone())).to_vec();
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);
let ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
(out, ms)
}
#[kernel(targets(cuda, rocm, vulkan, metal, cpu), unchecked)]
pub fn mmq_q4k_wmma_rt(
xq: &Array<i8>,
xs: &Array<f32>,
xsum: &Array<f32>,
wqs: &Array<u32>,
wsc: &Array<u32>,
wd: &Array<f32>,
wdm: &Array<f32>,
out: &mut Array<f32>,
meta: &Array<u32>, #[comptime] plane: usize,
#[comptime] target: Target,
) {
let m = meta[0] as usize;
let n = meta[1] as usize;
let k = meta[2] as usize;
let tid = UNIT_POS as usize;
let warp = tid / plane;
let lane = tid % plane;
let wm = warp / 4;
let wn = warp % 4;
let mrow0 = CUBE_POS_Y as usize * 32;
let ncol0 = CUBE_POS_X as usize * 64;
let kb_count = k / 32;
let nsb = k / 256;
let nthread = 8 * plane;
let per = 256 / plane;
let mut sa = SharedMemory::<i8>::new(1024usize);
let mut sb = SharedMemory::<i8>::new(2048usize);
let mut ci = SharedMemory::<i32>::new(2048usize);
let mut accf = SharedMemory::<f32>::new(2048usize);
for e in 0usize..(2048 / nthread) {
accf[tid * (2048 / nthread) + e] = 0.0f32;
}
sync_cube();
for kb in 0..kb_count {
let k0 = kb * 32;
for i in 0usize..(1024 / nthread) {
let idx = tid + i * nthread;
let arow = mrow0 + idx / 32;
sa[idx] = 0i8;
if arow < m {
sa[idx] = xq[arow * k + k0 + idx % 32];
}
}
let is = kb % 8;
let g = is / 2;
let sbk = kb / 8;
for i in 0usize..(2048 / nthread) {
let idx = tid + i * nthread;
let nrow = ncol0 + idx / 32;
let qi = idx % 32;
sb[idx] = 0i8;
if nrow < n {
let qb = q4k_byte(wqs, (nrow * nsb + sbk) * 32, g * 32 + qi);
let nib = (qb >> (4u32 * (is % 2) as u32)) & 15;
sb[idx] = i8::cast_from(nib);
}
}
sync_cube();
island! {
cuda | rocm | vulkan | metal => {
let c = cmma::Matrix::<i32>::from_value(
cmma::MatrixIdent::Accumulator, 16usize, 16usize, 16usize,
cmma::MatrixLayout::Undefined, 0i32,
);
let a0 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::A, 16usize, 16usize, 16usize,
cmma::MatrixLayout::RowMajor, &sa.slice(wm * 512, 1024), 32,
);
let b0 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::B, 16usize, 16usize, 16usize,
cmma::MatrixLayout::ColMajor, &sb.slice(wn * 512, 2048), 32,
);
cmma::execute::<i8, i8, i32, i32>(&a0, &b0, &c, &c);
let a1 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::A, 16usize, 16usize, 16usize,
cmma::MatrixLayout::RowMajor, &sa.slice(wm * 512 + 16, 1024), 32,
);
let b1 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::B, 16usize, 16usize, 16usize,
cmma::MatrixLayout::ColMajor, &sb.slice(wn * 512 + 16, 2048), 32,
);
cmma::execute::<i8, i8, i32, i32>(&a1, &b1, &c, &c);
cmma::store(&mut ci.slice_mut(warp * 256, warp * 256 + 256), &c, 16, cmma::MatrixLayout::RowMajor);
}
default => {
for e in 0usize..per {
let p = lane * per + e;
let smm = p / 16;
let snn = p % 16;
let mut isum = 0i32;
for l in 0usize..32 {
isum += i32::cast_from(sa[(wm * 16 + smm) * 32 + l])
* i32::cast_from(sb[(wn * 16 + snn) * 32 + l]);
}
ci[warp * 256 + p] = isum;
}
}
};
sync_cube();
for e in 0usize..per {
let p = lane * per + e;
let gmm = wm * 16 + p / 16;
let gnn = wn * 16 + p % 16;
let arow = mrow0 + gmm;
let nrow = ncol0 + gnn;
if arow < m && nrow < n {
let blk = nrow * nsb + sbk;
let scbase = blk * 3;
let wdd = wd[blk] * f32::cast_from(q4k_sc(wsc, scbase, is));
let wmm = wdm[blk] * f32::cast_from(q4k_m(wsc, scbase, is));
let xsc = xs[arow * kb_count + kb];
let xsm = xsum[arow * kb_count + kb];
accf[gmm * 64 + gnn] += xsc * wdd * f32::cast_from(ci[warp * 256 + p]) - wmm * xsm;
}
}
sync_cube();
}
for e in 0usize..per {
let p = lane * per + e;
let gmm = wm * 16 + p / 16;
let gnn = wn * 16 + p % 16;
let arow = mrow0 + gmm;
let nrow = ncol0 + gnn;
if arow < m && nrow < n {
out[arow * n + nrow] = accf[gmm * 64 + gnn];
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn mmq_q4k_wmma_rt_run<R: Runtime>(
client: &ComputeClient<R>,
xq: &[i8],
xs: &[f32],
xsum: &[f32],
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
m: usize,
n: usize,
k: usize,
iters: usize,
) -> (Vec<f32>, f64) {
let target = Target::of(client);
let plane = client.properties().hardware.plane_size_max;
let meta = [m as u32, n as u32, k as u32];
let xqh = client.create_from_slice(i8::as_bytes(xq));
let xsh = client.create_from_slice(f32::as_bytes(xs));
let xsumh = client.create_from_slice(f32::as_bytes(xsum));
let wqsh = client.create_from_slice(u32::as_bytes(wqs));
let wsch = client.create_from_slice(u32::as_bytes(wsc));
let wdh = client.create_from_slice(f32::as_bytes(wd));
let wdmh = client.create_from_slice(f32::as_bytes(wdm));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; m * n]));
let mh = client.create_from_slice(u32::as_bytes(&meta));
let grid = Grid::Static(n.div_ceil(64) as u32, m.div_ceil(32) as u32, 1);
let launch = |c: &ComputeClient<R>| unsafe {
mmq_q4k_wmma_rt::launch_unchecked::<R>(
c,
grid.clone(),
Block::new_1d(8 * plane),
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(wqsh.clone(), wqs.len()),
ArrayArg::from_raw_parts(wsch.clone(), wsc.len()),
ArrayArg::from_raw_parts(wdh.clone(), wd.len()),
ArrayArg::from_raw_parts(wdmh.clone(), wdm.len()),
ArrayArg::from_raw_parts(oh.clone(), m * n),
ArrayArg::from_raw_parts(mh.clone(), meta.len()),
plane as usize,
target,
);
};
launch(client);
let out = f32::from_bytes(&client.read_one_unchecked(oh.clone())).to_vec();
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);
let ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
(out, ms)
}
#[kernel(targets(cuda, rocm, vulkan, metal, cpu), unchecked)]
pub fn moe_expert_rows(
ids: &Array<u32>,
counts: &mut Array<u32>,
rows: &mut Array<u32>,
meta: &Array<u32>, #[comptime] nt: usize,
) {
let nslots = meta[0] as usize;
let ne = meta[1] as usize;
let cap = meta[2] as usize;
let tid = UNIT_POS as usize;
let mut i = tid;
while i < ne {
counts[i] = 0u32;
i += nt;
}
sync_cube();
if tid == 0 {
let mut s = 0usize;
while s < nslots {
let e = ids[s] as usize;
let c = counts[e] as usize;
if c < cap {
rows[e * cap + c] = s as u32;
counts[e] = (c + 1) as u32;
}
s += 1;
}
}
}
#[allow(clippy::too_many_arguments)]
#[kernel(targets(cuda, rocm, vulkan, metal, cpu), unchecked)]
pub fn mmq_q4k_id(
xq: &Array<i8>,
xs: &Array<f32>,
xsum: &Array<f32>,
wqs: &Array<u32>,
wsc: &Array<u32>,
wd: &Array<f32>,
wdm: &Array<f32>,
rows: &Array<u32>, counts: &Array<u32>, out: &mut Array<f32>,
meta: &Array<u32>, #[comptime] plane: usize,
#[comptime] target: Target,
) {
let n = meta[0] as usize;
let k = meta[1] as usize;
let cap = meta[2] as usize;
let expert = CUBE_POS_Z as usize;
let ecount = counts[expert] as usize;
let mut mtile = CUBE_POS_Y as usize;
let mtile_stride = CUBE_COUNT_Y as usize;
while mtile * 32 < ecount {
let mrow0 = mtile * 32;
let tid = UNIT_POS as usize;
let warp = tid / plane;
let lane = tid % plane;
let wm = warp / 4;
let wn = warp % 4;
let ncol0 = CUBE_POS_X as usize * 64;
let kb_count = k / 32;
let nsb = k / 256;
let nthread = 8 * plane;
let per = 256 / plane;
let wbase = expert * n;
let mut srow = SharedMemory::<u32>::new(32usize);
let mut sa = SharedMemory::<i8>::new(1024usize);
let mut sb = SharedMemory::<i8>::new(2048usize);
let mut ci = SharedMemory::<i32>::new(2048usize);
let mut accf = SharedMemory::<f32>::new(2048usize);
let mut r = tid;
while r < 32 {
let lr = mrow0 + r;
let mut s = 0u32;
if lr < ecount {
s = rows[expert * cap + lr];
}
srow[r] = s;
r += nthread;
}
for e in 0usize..(2048 / nthread) {
accf[tid * (2048 / nthread) + e] = 0.0f32;
}
sync_cube();
for kb in 0..kb_count {
let k0 = kb * 32;
for i in 0usize..(1024 / nthread) {
let idx = tid + i * nthread;
let lr = mrow0 + idx / 32;
sa[idx] = 0i8;
if lr < ecount {
let slot = srow[idx / 32] as usize;
sa[idx] = xq[slot * k + k0 + idx % 32];
}
}
let is = kb % 8;
let g = is / 2;
let sbk = kb / 8;
for i in 0usize..(2048 / nthread) {
let idx = tid + i * nthread;
let nrow = ncol0 + idx / 32;
let qi = idx % 32;
sb[idx] = 0i8;
if nrow < n {
let qb = q4k_byte(wqs, ((wbase + nrow) * nsb + sbk) * 32, g * 32 + qi);
let nib = (qb >> (4u32 * (is % 2) as u32)) & 15;
sb[idx] = i8::cast_from(nib);
}
}
sync_cube();
island! {
cuda | rocm | vulkan | metal => {
let c = cmma::Matrix::<i32>::from_value(
cmma::MatrixIdent::Accumulator, 16usize, 16usize, 16usize,
cmma::MatrixLayout::Undefined, 0i32,
);
let a0 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::A, 16usize, 16usize, 16usize,
cmma::MatrixLayout::RowMajor, &sa.slice(wm * 512, 1024), 32,
);
let b0 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::B, 16usize, 16usize, 16usize,
cmma::MatrixLayout::ColMajor, &sb.slice(wn * 512, 2048), 32,
);
cmma::execute::<i8, i8, i32, i32>(&a0, &b0, &c, &c);
let a1 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::A, 16usize, 16usize, 16usize,
cmma::MatrixLayout::RowMajor, &sa.slice(wm * 512 + 16, 1024), 32,
);
let b1 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::B, 16usize, 16usize, 16usize,
cmma::MatrixLayout::ColMajor, &sb.slice(wn * 512 + 16, 2048), 32,
);
cmma::execute::<i8, i8, i32, i32>(&a1, &b1, &c, &c);
cmma::store(&mut ci.slice_mut(warp * 256, warp * 256 + 256), &c, 16, cmma::MatrixLayout::RowMajor);
}
default => {
for e in 0usize..per {
let p = lane * per + e;
let smm = p / 16;
let snn = p % 16;
let mut isum = 0i32;
for l in 0usize..32 {
isum += i32::cast_from(sa[(wm * 16 + smm) * 32 + l])
* i32::cast_from(sb[(wn * 16 + snn) * 32 + l]);
}
ci[warp * 256 + p] = isum;
}
}
};
sync_cube();
for e in 0usize..per {
let p = lane * per + e;
let gmm = wm * 16 + p / 16;
let gnn = wn * 16 + p % 16;
let lr = mrow0 + gmm;
let nrow = ncol0 + gnn;
if lr < ecount && nrow < n {
let slot = srow[gmm] as usize;
let blk = (wbase + nrow) * nsb + sbk;
let scbase = blk * 3;
let wdd = wd[blk] * f32::cast_from(q4k_sc(wsc, scbase, is));
let wmm = wdm[blk] * f32::cast_from(q4k_m(wsc, scbase, is));
let xsc = xs[slot * kb_count + kb];
let xsm = xsum[slot * kb_count + kb];
accf[gmm * 64 + gnn] +=
xsc * wdd * f32::cast_from(ci[warp * 256 + p]) - wmm * xsm;
}
}
sync_cube();
}
for e in 0usize..per {
let p = lane * per + e;
let gmm = wm * 16 + p / 16;
let gnn = wn * 16 + p % 16;
let lr = mrow0 + gmm;
let nrow = ncol0 + gnn;
if lr < ecount && nrow < n {
let slot = srow[gmm] as usize;
out[slot * n + nrow] = accf[gmm * 64 + gnn];
}
}
sync_cube();
mtile += mtile_stride;
}
}
pub fn moe_expert_rows_run<R: Runtime>(
client: &ComputeClient<R>,
ids: &[u32],
n_experts: usize,
cap: usize,
) -> (Vec<u32>, Vec<u32>) {
let idsh = client.create_from_slice(u32::as_bytes(ids));
let cntsh = client.create_from_slice(u32::as_bytes(&vec![0u32; n_experts]));
let rowsh = client.create_from_slice(u32::as_bytes(&vec![0u32; n_experts * cap]));
let meta = [ids.len() as u32, n_experts as u32, cap as u32];
let mh = client.create_from_slice(u32::as_bytes(&meta));
unsafe {
moe_expert_rows::launch_unchecked::<R>(
client,
Grid::Static(1, 1, 1),
Block::new_1d(64),
ArrayArg::from_raw_parts(idsh, ids.len()),
ArrayArg::from_raw_parts(cntsh.clone(), n_experts),
ArrayArg::from_raw_parts(rowsh.clone(), n_experts * cap),
ArrayArg::from_raw_parts(mh, meta.len()),
64usize,
);
}
let counts = u32::from_bytes(&client.read_one_unchecked(cntsh)).to_vec();
let rows = u32::from_bytes(&client.read_one_unchecked(rowsh)).to_vec();
(counts, rows)
}
#[allow(clippy::too_many_arguments)]
pub fn mmq_q4k_id_run<R: Runtime>(
client: &ComputeClient<R>,
xq: &[i8],
xs: &[f32],
xsum: &[f32],
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
ids: &[u32],
nslots: usize,
n_experts: usize,
cap: usize,
n: usize,
k: usize,
ytiles: usize,
iters: usize,
) -> (Vec<f32>, f64) {
let target = Target::of(client);
let plane = client.properties().hardware.plane_size_max;
let xqh = client.create_from_slice(i8::as_bytes(xq));
let xsh = client.create_from_slice(f32::as_bytes(xs));
let xsumh = client.create_from_slice(f32::as_bytes(xsum));
let wqsh = client.create_from_slice(u32::as_bytes(wqs));
let wsch = client.create_from_slice(u32::as_bytes(wsc));
let wdh = client.create_from_slice(f32::as_bytes(wd));
let wdmh = client.create_from_slice(f32::as_bytes(wdm));
let idsh = client.create_from_slice(u32::as_bytes(ids));
let cntsh = client.create_from_slice(u32::as_bytes(&vec![0u32; n_experts]));
let rowsh = client.create_from_slice(u32::as_bytes(&vec![0u32; n_experts * cap]));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; nslots * n]));
let gmeta = [nslots as u32, n_experts as u32, cap as u32];
let gmh = client.create_from_slice(u32::as_bytes(&gmeta));
let meta = [n as u32, k as u32, cap as u32];
let mh = client.create_from_slice(u32::as_bytes(&meta));
unsafe {
moe_expert_rows::launch_unchecked::<R>(
client,
Grid::Static(1, 1, 1),
Block::new_1d(64),
ArrayArg::from_raw_parts(idsh.clone(), ids.len()),
ArrayArg::from_raw_parts(cntsh.clone(), n_experts),
ArrayArg::from_raw_parts(rowsh.clone(), n_experts * cap),
ArrayArg::from_raw_parts(gmh.clone(), gmeta.len()),
64usize,
);
}
let grid = Grid::Static(
n.div_ceil(64) as u32,
ytiles.max(1) as u32,
n_experts as u32,
);
let launch = |c: &ComputeClient<R>| unsafe {
mmq_q4k_id::launch_unchecked::<R>(
c,
grid.clone(),
Block::new_1d(8 * plane),
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(wqsh.clone(), wqs.len()),
ArrayArg::from_raw_parts(wsch.clone(), wsc.len()),
ArrayArg::from_raw_parts(wdh.clone(), wd.len()),
ArrayArg::from_raw_parts(wdmh.clone(), wdm.len()),
ArrayArg::from_raw_parts(rowsh.clone(), n_experts * cap),
ArrayArg::from_raw_parts(cntsh.clone(), n_experts),
ArrayArg::from_raw_parts(oh.clone(), nslots * n),
ArrayArg::from_raw_parts(mh.clone(), meta.len()),
plane as usize,
target,
);
};
launch(client);
let out = f32::from_bytes(&client.read_one_unchecked(oh.clone())).to_vec();
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);
let ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
(out, ms)
}
#[allow(clippy::too_many_arguments)]
pub fn mmq_q4k_id_ref(
xq: &[i8],
xs: &[f32],
xsum: &[f32],
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
ids: &[u32],
nslots: usize,
n: usize,
k: usize,
) -> Vec<f32> {
let kb = k / 32;
let nsb = k / 256;
let mut out = vec![0.0f32; nslots * n];
for s in 0..nslots {
let wbase = ids[s] as usize * n;
for j in 0..n {
let mut acc = 0.0f32;
for b in 0..kb {
let is = b % 8;
let g = is / 2;
let blk = (wbase + j) * nsb + b / 8;
let mut isum = 0i32;
for qi in 0..32 {
let qbyte = cpu_q4k_byte(wqs, blk * 32, g * 32 + qi);
let nib = ((qbyte >> (4 * (is % 2))) & 15) as i32;
isum += xq[s * k + b * 32 + qi] as i32 * nib;
}
let dd = wd[blk] * cpu_q4k_sc(wsc, blk * 3, is) as f32;
let mm = wdm[blk] * cpu_q4k_m(wsc, blk * 3, is) as f32;
acc += xs[s * kb + b] * dd * isum as f32 - mm * xsum[s * kb + b];
}
out[s * n + j] = acc;
}
}
out
}
fn cpu_q4k_byte(a: &[u32], base: usize, i: usize) -> u32 {
(a[base + i / 4] >> (8 * (i % 4))) & 255
}
fn cpu_q4k_sc(wsc: &[u32], sb: usize, j: usize) -> u32 {
if j < 4 {
cpu_q4k_byte(wsc, sb, j) & 63
} else {
(cpu_q4k_byte(wsc, sb, j + 4) & 15) | ((cpu_q4k_byte(wsc, sb, j - 4) >> 6) << 4)
}
}
fn cpu_q4k_m(wsc: &[u32], sb: usize, j: usize) -> u32 {
if j < 4 {
cpu_q4k_byte(wsc, sb, j + 4) & 63
} else {
(cpu_q4k_byte(wsc, sb, j + 4) >> 4) | ((cpu_q4k_byte(wsc, sb, j) >> 6) << 4)
}
}
#[allow(clippy::too_many_arguments)]
pub fn mmq_q4k_ref(
xq: &[i8],
xs: &[f32],
xsum: &[f32],
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
m: usize,
n: usize,
k: usize,
) -> Vec<f32> {
let kb = k / 32;
let nsb = k / 256;
let mut out = vec![0.0f32; m * n];
for i in 0..m {
for j in 0..n {
let mut acc = 0.0f32;
for b in 0..kb {
let is = b % 8;
let g = is / 2;
let blk = j * nsb + b / 8;
let mut isum = 0i32;
for qi in 0..32 {
let qbyte = cpu_q4k_byte(wqs, blk * 32, g * 32 + qi);
let nib = ((qbyte >> (4 * (is % 2))) & 15) as i32;
isum += xq[i * k + b * 32 + qi] as i32 * nib;
}
let dd = wd[blk] * cpu_q4k_sc(wsc, blk * 3, is) as f32;
let mm = wdm[blk] * cpu_q4k_m(wsc, blk * 3, is) as f32;
acc += xs[i * kb + b] * dd * isum as f32 - mm * xsum[i * kb + b];
}
out[i * n + j] = acc;
}
}
out
}
pub fn gen_mmq_q4k(
m: usize,
n: usize,
k: usize,
) -> (
Vec<i8>,
Vec<f32>,
Vec<f32>,
Vec<u32>,
Vec<u32>,
Vec<f32>,
Vec<f32>,
) {
let kb = k / 32;
let nsb = k / 256;
let mut s = 0x243F6A8885A308D3u64;
let mut next = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
let xq: Vec<i8> = (0..m * k)
.map(|_| ((next() % 255) as i64 - 127) as i8)
.collect();
let wqs: Vec<u32> = (0..n * nsb * 32).map(|_| next() as u32).collect(); let wsc: Vec<u32> = (0..n * nsb * 3).map(|_| next() as u32).collect(); let wd: Vec<f32> = (0..n * nsb)
.map(|_| (next() % 1000) as f32 / 20000.0 + 0.002)
.collect();
let wdm: Vec<f32> = (0..n * nsb)
.map(|_| (next() % 1000) as f32 / 40000.0)
.collect();
let xs: Vec<f32> = (0..m * kb)
.map(|_| (next() % 1000) as f32 / 50000.0 + 0.002)
.collect();
let mut xsum = vec![0.0f32; m * kb];
for i in 0..m {
for b in 0..kb {
let mut acc = 0i32;
for l in 0..32 {
acc += xq[i * k + b * 32 + l] as i32;
}
xsum[i * kb + b] = xs[i * kb + b] * acc as f32;
}
}
(xq, xs, xsum, wqs, wsc, wd, wdm)
}
#[kernel(targets(cuda, rocm, vulkan, metal, cpu), unchecked)]
#[allow(clippy::too_many_arguments)]
pub fn mmq_q4k_wmma_tile(
xq: &Array<i8>,
xs: &Array<f32>,
xsum: &Array<f32>,
wqs: &Array<u32>,
wsc: &Array<u32>,
wd: &Array<f32>,
wdm: &Array<f32>,
out: &mut Array<f32>,
#[comptime] m: usize,
#[comptime] n: usize,
#[comptime] k: usize,
#[comptime] wm: usize,
#[comptime] wn: usize,
#[comptime] rm: usize,
#[comptime] rn: usize,
#[comptime] plane: usize,
#[comptime] target: Target,
) {
let nwarp = wm * wn;
let bm = wm * rm * 16;
let bn = wn * rn * 16;
let tid = UNIT_POS as usize;
let warp = tid / plane;
let lane = tid % plane;
let warp_m = warp / wn; let warp_n = warp % wn; let mrow0 = CUBE_POS_Y as usize * bm;
let ncol0 = CUBE_POS_X as usize * bn;
let kb_count = k / 32;
let nsb = k / 256; let nthread = nwarp * plane;
let per = 256 / plane;
let mut sa = SharedMemory::<i8>::new(bm * 32); let mut sb = SharedMemory::<i8>::new(bn * 32); let mut ci = SharedMemory::<i32>::new(bm * bn);
let mut acc = Array::<f32>::new(rm * rn * per);
for a in 0..(rm * rn * per) {
acc[a] = 0.0f32;
}
let a_elems = bm * 32;
let a_iters = (bm * 32 + nthread - 1) / nthread;
let b_elems = bn * 32;
let b_iters = (bn * 32 + nthread - 1) / nthread;
for kb in 0..kb_count {
let k0 = kb * 32;
for i in 0..a_iters {
let idx = tid + i * nthread;
if idx < a_elems {
let arow = mrow0 + idx / 32;
let mut v = 0i8;
if arow < m {
v = xq[arow * k + k0 + idx % 32];
}
sa[idx] = v;
}
}
let is = kb % 8;
let g = is / 2;
let sbk = kb / 8;
for i in 0..b_iters {
let idx = tid + i * nthread;
if idx < b_elems {
let nrow = ncol0 + idx / 32;
let qi = idx % 32;
let mut v = 0i8;
if nrow < n {
let qb = q4k_byte(wqs, (nrow * nsb + sbk) * 32, g * 32 + qi);
let nib = (qb >> (4u32 * (is % 2) as u32)) & 15;
v = i8::cast_from(nib);
}
sb[idx] = v;
}
}
sync_cube();
#[unroll]
for im in 0..rm {
#[unroll]
for jn in 0..rn {
let sm = warp_m * rm + im; let sn = warp_n * rn + jn; let cbase = (warp * rm * rn + im * rn + jn) * 256;
island! {
cuda | rocm | vulkan | metal => {
let c = cmma::Matrix::<i32>::from_value(
cmma::MatrixIdent::Accumulator, 16usize, 16usize, 16usize,
cmma::MatrixLayout::Undefined, 0i32,
);
let a0 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::A, 16usize, 16usize, 16usize,
cmma::MatrixLayout::RowMajor, &sa.slice(sm * 512, bm * 32), 32,
);
let b0 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::B, 16usize, 16usize, 16usize,
cmma::MatrixLayout::ColMajor, &sb.slice(sn * 512, bn * 32), 32,
);
cmma::execute::<i8, i8, i32, i32>(&a0, &b0, &c, &c);
let a1 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::A, 16usize, 16usize, 16usize,
cmma::MatrixLayout::RowMajor, &sa.slice(sm * 512 + 16, bm * 32), 32,
);
let b1 = cmma::Matrix::<i8>::from_slice(
cmma::MatrixIdent::B, 16usize, 16usize, 16usize,
cmma::MatrixLayout::ColMajor, &sb.slice(sn * 512 + 16, bn * 32), 32,
);
cmma::execute::<i8, i8, i32, i32>(&a1, &b1, &c, &c);
cmma::store(&mut ci.slice_mut(cbase, cbase + 256), &c, 16, cmma::MatrixLayout::RowMajor);
}
default => {
for e in 0usize..per {
let p = lane * per + e;
let smm = p / 16;
let snn = p % 16;
let mut isum = 0i32;
for l in 0usize..32 {
isum += i32::cast_from(sa[(sm * 16 + smm) * 32 + l])
* i32::cast_from(sb[(sn * 16 + snn) * 32 + l]);
}
ci[cbase + p] = isum;
}
}
};
}
}
sync_cube();
#[unroll]
for im in 0..rm {
#[unroll]
for jn in 0..rn {
let sm = warp_m * rm + im;
let sn = warp_n * rn + jn;
let cbase = (warp * rm * rn + im * rn + jn) * 256;
for e in 0usize..per {
let p = lane * per + e;
let gmm = sm * 16 + p / 16;
let gnn = sn * 16 + p % 16;
let arow = mrow0 + gmm;
let nrow = ncol0 + gnn;
if arow < m && nrow < n {
let blk = nrow * nsb + sbk;
let scbase = blk * 3;
let wdd = wd[blk] * f32::cast_from(q4k_sc(wsc, scbase, is));
let wmm = wdm[blk] * f32::cast_from(q4k_m(wsc, scbase, is));
let xsc = xs[arow * kb_count + kb];
let xsm = xsum[arow * kb_count + kb];
acc[(im * rn + jn) * per + e] +=
xsc * wdd * f32::cast_from(ci[cbase + p]) - wmm * xsm;
}
}
}
}
sync_cube();
}
#[unroll]
for im in 0..rm {
#[unroll]
for jn in 0..rn {
let sm = warp_m * rm + im;
let sn = warp_n * rn + jn;
for e in 0usize..per {
let p = lane * per + e;
let gmm = sm * 16 + p / 16;
let gnn = sn * 16 + p % 16;
let arow = mrow0 + gmm;
let nrow = ncol0 + gnn;
if arow < m && nrow < n {
out[arow * n + nrow] = acc[(im * rn + jn) * per + e];
}
}
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn mmq_q4k_wmma_tile_run<R: Runtime>(
client: &ComputeClient<R>,
xq: &[i8],
xs: &[f32],
xsum: &[f32],
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
m: usize,
n: usize,
k: usize,
wm: usize,
wn: usize,
rm: usize,
rn: usize,
iters: usize,
) -> (Vec<f32>, f64) {
let target = Target::of(client);
let plane = client.properties().hardware.plane_size_max as usize;
let bm = wm * rm * 16;
let bn = wn * rn * 16;
let xqh = client.create_from_slice(i8::as_bytes(xq));
let xsh = client.create_from_slice(f32::as_bytes(xs));
let xsumh = client.create_from_slice(f32::as_bytes(xsum));
let wqsh = client.create_from_slice(u32::as_bytes(wqs));
let wsch = client.create_from_slice(u32::as_bytes(wsc));
let wdh = client.create_from_slice(f32::as_bytes(wd));
let wdmh = client.create_from_slice(f32::as_bytes(wdm));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; m * n]));
let grid = Grid::Static(n.div_ceil(bn) as u32, m.div_ceil(bm) as u32, 1);
let block = Block::new_1d((wm * wn * plane) as u32);
let launch = |c: &ComputeClient<R>| unsafe {
mmq_q4k_wmma_tile::launch_unchecked::<R>(
c,
grid.clone(),
block,
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(wqsh.clone(), wqs.len()),
ArrayArg::from_raw_parts(wsch.clone(), wsc.len()),
ArrayArg::from_raw_parts(wdh.clone(), wd.len()),
ArrayArg::from_raw_parts(wdmh.clone(), wdm.len()),
ArrayArg::from_raw_parts(oh.clone(), m * n),
m,
n,
k,
wm,
wn,
rm,
rn,
plane,
target,
);
};
launch(client);
let out = f32::from_bytes(&client.read_one_unchecked(oh.clone())).to_vec();
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);
let ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
(out, ms)
}
#[kernel(targets(cuda, rocm, vulkan, metal, cpu), unchecked)]
#[allow(clippy::too_many_arguments)]
pub fn mmq_q4k_coopmat_tile<F: Float>(
x: &Array<F>, wqs: &Array<u32>, wsc: &Array<u32>, wd: &Array<F>, wdm: &Array<F>, out: &mut Array<F>,
#[comptime] m: usize,
#[comptime] n: usize,
#[comptime] k: usize,
#[comptime] wm: usize,
#[comptime] wn: usize,
#[comptime] rm: usize,
#[comptime] rn: usize,
#[comptime] plane: usize,
#[comptime] target: Target,
) {
let nwarp = wm * wn;
let bm = wm * rm * 16;
let bn = wn * rn * 16;
let tid = UNIT_POS as usize;
let warp = tid / plane;
let lane = tid % plane;
let warp_m = warp / wn;
let warp_n = warp % wn;
let mrow0 = CUBE_POS_Y as usize * bm;
let ncol0 = CUBE_POS_X as usize * bn;
let kb_count = k / 32;
let nsb = k / 256;
let nthread = nwarp * plane;
let per = 256 / plane;
let a_iters = (bm * 32 + nthread - 1) / nthread;
let b_iters = (bn * 32 + nthread - 1) / nthread;
island! {
cuda | rocm | vulkan | metal => {
let mut sa = SharedMemory::<half::f16>::new(bm * 32);
let mut sb = SharedMemory::<half::f16>::new(bn * 32);
let mut acc = Sequence::<cmma::Matrix<F>>::new();
#[unroll]
for _i in 0..(rm * rn) {
acc.push(cmma::Matrix::<F>::from_value(
cmma::MatrixIdent::Accumulator, 16usize, 16usize, 16usize, cmma::MatrixLayout::Undefined, F::new(0.0),
));
}
for kb in 0..kb_count {
let k0 = kb * 32;
for i in 0..a_iters {
let idx = tid + i * nthread;
if idx < bm * 32 {
let arow = mrow0 + idx / 32;
let mut v = F::new(0.0);
if arow < m {
v = x[arow * k + k0 + idx % 32];
}
sa[idx] = half::f16::cast_from(v);
}
}
let is = kb % 8;
let g = is / 2;
let sbk = kb / 8;
let shift = 4u32 * (is % 2) as u32;
for i in 0..b_iters {
let idx = tid + i * nthread;
if idx < bn * 32 {
let nrow = ncol0 + idx / 32;
let qi = idx % 32;
let mut wv = F::new(0.0);
if nrow < n {
let blk = nrow * nsb + sbk;
let ds = wd[blk] * F::cast_from(q4k_sc(wsc, blk * 3, is));
let ms = wdm[blk] * F::cast_from(q4k_m(wsc, blk * 3, is));
let qb = q4k_byte(wqs, blk * 32, g * 32 + qi);
let nib = (qb >> shift) & 15;
wv = ds * F::cast_from(nib) - ms;
}
sb[idx] = half::f16::cast_from(wv);
}
}
sync_cube();
#[unroll]
for kc in 0..2usize {
#[unroll]
for im in 0..rm {
let sm = warp_m * rm + im;
let ma = cmma::Matrix::<half::f16>::from_slice(
cmma::MatrixIdent::A, 16usize, 16usize, 16usize, cmma::MatrixLayout::RowMajor,
&sa.slice(sm * 512 + kc * 16, bm * 32), 32,
);
#[unroll]
for jn in 0..rn {
let sn = warp_n * rn + jn;
let mb = cmma::Matrix::<half::f16>::from_slice(
cmma::MatrixIdent::B, 16usize, 16usize, 16usize, cmma::MatrixLayout::ColMajor,
&sb.slice(sn * 512 + kc * 16, bn * 32), 32,
);
let c = acc.index(im * rn + jn);
cmma::execute::<half::f16, half::f16, F, F>(&ma, &mb, c, c);
}
}
}
sync_cube();
}
#[unroll]
for im in 0..rm {
#[unroll]
for jn in 0..rn {
let sm = warp_m * rm + im;
let sn = warp_n * rn + jn;
let orow = mrow0 + sm * 16;
let ocol = ncol0 + sn * 16;
if orow < m && ocol < n {
cmma::store(
&mut out.slice_mut(orow * n + ocol, out.len()),
acc.index(im * rn + jn),
n as u32,
cmma::MatrixLayout::RowMajor,
);
}
}
}
}
default => {
let mut sa = SharedMemory::<half::f16>::new(bm * 32);
let mut sb = SharedMemory::<half::f16>::new(bn * 32);
let mut acc = Array::<F>::new(rm * rn * per);
for a in 0..(rm * rn * per) {
acc[a] = F::new(0.0);
}
for kb in 0..kb_count {
let k0 = kb * 32;
for i in 0..a_iters {
let idx = tid + i * nthread;
if idx < bm * 32 {
let arow = mrow0 + idx / 32;
let mut v = F::new(0.0);
if arow < m {
v = x[arow * k + k0 + idx % 32];
}
sa[idx] = half::f16::cast_from(v);
}
}
let is = kb % 8;
let g = is / 2;
let sbk = kb / 8;
let shift = 4u32 * (is % 2) as u32;
for i in 0..b_iters {
let idx = tid + i * nthread;
if idx < bn * 32 {
let nrow = ncol0 + idx / 32;
let qi = idx % 32;
let mut wv = F::new(0.0);
if nrow < n {
let blk = nrow * nsb + sbk;
let ds = wd[blk] * F::cast_from(q4k_sc(wsc, blk * 3, is));
let ms = wdm[blk] * F::cast_from(q4k_m(wsc, blk * 3, is));
let qb = q4k_byte(wqs, blk * 32, g * 32 + qi);
let nib = (qb >> shift) & 15;
wv = ds * F::cast_from(nib) - ms;
}
sb[idx] = half::f16::cast_from(wv);
}
}
sync_cube();
#[unroll]
for im in 0..rm {
#[unroll]
for jn in 0..rn {
let sm = warp_m * rm + im;
let sn = warp_n * rn + jn;
for e in 0..per {
let p = lane * per + e;
let smm = p / 16;
let snn = p % 16;
let mut s = F::new(0.0);
for l in 0..32 {
s += F::cast_from(sa[(sm * 16 + smm) * 32 + l])
* F::cast_from(sb[(sn * 16 + snn) * 32 + l]);
}
acc[(im * rn + jn) * per + e] += s;
}
}
}
sync_cube();
}
#[unroll]
for im in 0..rm {
#[unroll]
for jn in 0..rn {
let sm = warp_m * rm + im;
let sn = warp_n * rn + jn;
for e in 0..per {
let p = lane * per + e;
let orow = mrow0 + sm * 16 + p / 16;
let ocol = ncol0 + sn * 16 + p % 16;
if orow < m && ocol < n {
out[orow * n + ocol] = acc[(im * rn + jn) * per + e];
}
}
}
}
}
};
}
#[allow(clippy::too_many_arguments)]
pub fn mmq_q4k_f16_ref(
x: &[f32],
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
m: usize,
n: usize,
k: usize,
) -> Vec<f32> {
let kb = k / 32;
let nsb = k / 256;
let mut out = vec![0.0f32; m * n];
for i in 0..m {
for j in 0..n {
let mut acc = 0.0f32;
for b in 0..kb {
let is = b % 8;
let g = is / 2;
let blk = j * nsb + b / 8;
let dd = wd[blk] * cpu_q4k_sc(wsc, blk * 3, is) as f32;
let mm = wdm[blk] * cpu_q4k_m(wsc, blk * 3, is) as f32;
for qi in 0..32 {
let qbyte = cpu_q4k_byte(wqs, blk * 32, g * 32 + qi);
let nib = ((qbyte >> (4 * (is % 2))) & 15) as f32;
acc += x[i * k + b * 32 + qi] * (dd * nib - mm);
}
}
out[i * n + j] = acc;
}
}
out
}
#[allow(clippy::too_many_arguments)]
pub fn mmq_q4k_coopmat_tile_run<R: Runtime>(
client: &ComputeClient<R>,
x: &[f32],
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
m: usize,
n: usize,
k: usize,
wm: usize,
wn: usize,
rm: usize,
rn: usize,
iters: usize,
) -> (Vec<f32>, f64) {
let target = Target::of(client);
let plane = client.properties().hardware.plane_size_max as usize;
let bm = wm * rm * 16;
let bn = wn * rn * 16;
let xh = client.create_from_slice(f32::as_bytes(x));
let wqsh = client.create_from_slice(u32::as_bytes(wqs));
let wsch = client.create_from_slice(u32::as_bytes(wsc));
let wdh = client.create_from_slice(f32::as_bytes(wd));
let wdmh = client.create_from_slice(f32::as_bytes(wdm));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; m * n]));
let grid = Grid::Static(n.div_ceil(bn) as u32, m.div_ceil(bm) as u32, 1);
let block = Block::new_1d((wm * wn * plane) as u32);
let launch = |c: &ComputeClient<R>| unsafe {
mmq_q4k_coopmat_tile::launch_unchecked::<f32, R>(
c,
grid.clone(),
block,
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(wqsh.clone(), wqs.len()),
ArrayArg::from_raw_parts(wsch.clone(), wsc.len()),
ArrayArg::from_raw_parts(wdh.clone(), wd.len()),
ArrayArg::from_raw_parts(wdmh.clone(), wdm.len()),
ArrayArg::from_raw_parts(oh.clone(), m * n),
m,
n,
k,
wm,
wn,
rm,
rn,
plane,
target,
);
};
launch(client);
let out = f32::from_bytes(&client.read_one_unchecked(oh.clone())).to_vec();
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);
let ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
(out, ms)
}
pub fn mmq_q4k_space() -> Space {
Space::new()
.param("WM", [1, 2, 4])
.param("WN", [1, 2, 4])
.param("RM", [1, 2, 4])
.param("RN", [1, 2, 4, 8])
.constraint(|c, s| {
let w = c.get(s, "WM") * c.get(s, "WN");
(1..=16).contains(&w)
})
}
fn mmq_tile_cfg(c: &Config, s: &Space) -> (usize, usize, usize, usize) {
(
c.get(s, "WM") as usize,
c.get(s, "WN") as usize,
c.get(s, "RM") as usize,
c.get(s, "RN") as usize,
)
}
fn mmq_geom(c: &Config, s: &Space) -> (usize, usize, usize) {
let (wm, wn, rm, rn) = mmq_tile_cfg(c, s);
(wm * rm * 16, wn * rn * 16, wm * wn)
}
fn mmq_lds_bytes(c: &Config, s: &Space) -> usize {
let (bm, bn, _) = mmq_geom(c, s);
bm * 32 + bn * 32 + bm * bn * 4
}
pub fn mmq_q4k_incumbent() -> (usize, usize, usize, usize) {
(2, 4, 1, 1) }
pub struct MmqQ4kEval<'a, R: Runtime> {
client: &'a ComputeClient<R>,
space: &'a Space,
banks: Vec<Handle>, xqh: Handle,
xsh: Handle,
xsumh: Handle,
wsch: Handle,
wdh: Handle,
wdmh: Handle,
outh: Handle,
wqs_len: usize,
xq_len: usize,
xs_len: usize,
xsum_len: usize,
wsc_len: usize,
wd_len: usize,
wdm_len: usize,
m: usize,
n: usize,
k: usize,
plane: usize,
lds_budget: usize,
reg_budget: usize,
oracle: Vec<f32>,
maxref: f32,
repeats: usize,
worst_rel: std::cell::Cell<f32>,
}
impl<'a, R: Runtime> MmqQ4kEval<'a, R> {
#[allow(clippy::too_many_arguments)]
pub fn new(
client: &'a ComputeClient<R>,
space: &'a Space,
wqs_banks: &[Vec<u32>],
xq: &[i8],
xs: &[f32],
xsum: &[f32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
m: usize,
n: usize,
k: usize,
repeats: usize,
) -> Self {
assert!(
!wqs_banks.is_empty(),
"MmqQ4kEval needs at least one weight bank"
);
let banks: Vec<Handle> = wqs_banks
.iter()
.map(|w| client.create_from_slice(u32::as_bytes(w)))
.collect();
let xqh = client.create_from_slice(i8::as_bytes(xq));
let xsh = client.create_from_slice(f32::as_bytes(xs));
let xsumh = client.create_from_slice(f32::as_bytes(xsum));
let wsch = client.create_from_slice(u32::as_bytes(wsc));
let wdh = client.create_from_slice(f32::as_bytes(wd));
let wdmh = client.create_from_slice(f32::as_bytes(wdm));
let outh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; m * n]));
let oracle = mmq_q4k_ref(xq, xs, xsum, &wqs_banks[0], wsc, wd, wdm, m, n, k);
let maxref = oracle.iter().fold(0f32, |a, &v| a.max(v.abs())).max(1e-30);
let plane = client.properties().hardware.plane_size_max as usize;
Self {
client,
space,
banks,
xqh,
xsh,
xsumh,
wsch,
wdh,
wdmh,
outh,
wqs_len: wqs_banks[0].len(),
xq_len: xq.len(),
xs_len: xs.len(),
xsum_len: xsum.len(),
wsc_len: wsc.len(),
wd_len: wd.len(),
wdm_len: wdm.len(),
m,
n,
k,
plane,
lds_budget: 48 * 1024, reg_budget: 16, 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, wm: usize, wn: usize, rm: usize, rn: usize) {
let bm = wm * rm * 16;
let bn = wn * rn * 16;
let grid = Grid::Static(self.n.div_ceil(bn) as u32, self.m.div_ceil(bm) as u32, 1);
let block = Block::new_1d((wm * wn * self.plane) as u32);
let target = Target::of(self.client);
unsafe {
mmq_q4k_wmma_tile::launch_unchecked::<R>(
self.client,
grid,
block,
ArrayArg::from_raw_parts(self.xqh.clone(), self.xq_len),
ArrayArg::from_raw_parts(self.xsh.clone(), self.xs_len),
ArrayArg::from_raw_parts(self.xsumh.clone(), self.xsum_len),
ArrayArg::from_raw_parts(bank.clone(), self.wqs_len),
ArrayArg::from_raw_parts(self.wsch.clone(), self.wsc_len),
ArrayArg::from_raw_parts(self.wdh.clone(), self.wd_len),
ArrayArg::from_raw_parts(self.wdmh.clone(), self.wdm_len),
ArrayArg::from_raw_parts(self.outh.clone(), self.m * self.n),
self.m,
self.n,
self.k,
wm,
wn,
rm,
rn,
self.plane,
target,
);
}
}
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 MmqQ4kEval<'a, R> {
fn static_check(&self, cfg: &Config) -> Verdict {
let (bm, bn, nwarp) = mmq_geom(cfg, self.space);
let wg = nwarp * self.plane;
if wg > 1024 {
return Verdict::Reject(format!(
"workgroup width {wg} = nwarp {nwarp} x plane {} > 1024",
self.plane
));
}
let lds = mmq_lds_bytes(cfg, self.space);
if lds > self.lds_budget {
return Verdict::Reject(format!(
"LDS {}KB (BM {bm} x BN {bn}) > {}KB budget (occupancy)",
lds / 1024,
self.lds_budget / 1024
));
}
let (wm, wn, rm, rn) = mmq_tile_cfg(cfg, self.space);
let _ = (wm, wn, bm, bn);
let frags = rm * rn;
if frags > self.reg_budget {
return Verdict::Reject(format!(
"register tile {frags} fragments/warp (rm {rm} x rn {rn}) > {} budget",
self.reg_budget
));
}
Verdict::Pass
}
fn measure(&self, cfg: &Config, iters: usize) -> f64 {
let (wm, wn, rm, rn) = mmq_tile_cfg(cfg, self.space);
self.dispatch(&self.banks[0], wm, wn, rm, rn);
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], wm, wn, rm, rn);
}
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], wm, wn, rm, rn);
}
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 mmq_q4k_hunt<R: Runtime>(
tuner: &Tuner,
device: &str,
m: usize,
n: usize,
k: usize,
eval: &MmqQ4kEval<R>,
evo: &Evolution,
seed: u64,
) -> Evolved {
tuner.evolve(
device,
"mmq_q4k_tile",
&format!("m={m},n={n},k={k}"),
eval.space(),
eval,
evo,
seed,
)
}
#[allow(clippy::too_many_arguments)]
pub fn mmq_q4k_autokernel<R: Runtime>(
client: &ComputeClient<R>,
tuner: &Tuner,
xq: &[i8],
xs: &[f32],
xsum: &[f32],
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
m: usize,
n: usize,
k: usize,
) -> (Vec<f32>, String) {
let device = crate::tune::device_id(client);
let space = mmq_q4k_space();
let key = format!("m={m},n={n},k={k}");
let (wm, wn, rm, rn, name) = match tuner.cached_winner(&device, "mmq_q4k_tile", &key) {
Some(w) => match space.parse(&w) {
Some(cfg) => {
let (a, b, c, d) = mmq_tile_cfg(&cfg, &space);
(a, b, c, d, w)
}
None => {
let (a, b, c, d) = mmq_q4k_incumbent();
(a, b, c, d, "incumbent".to_string())
}
},
None => {
let (a, b, c, d) = mmq_q4k_incumbent();
(a, b, c, d, "incumbent".to_string())
}
};
let (out, _ms) = mmq_q4k_wmma_tile_run(
client, xq, xs, xsum, wqs, wsc, wd, wdm, m, n, k, wm, wn, rm, rn, 1,
);
(out, name)
}
fn mmq_coopmat_lds_bytes(c: &Config, s: &Space) -> usize {
let (bm, bn, _) = mmq_geom(c, s);
(bm * 32 + bn * 32) * 2
}
pub struct CoopmatF16Eval<'a, R: Runtime> {
client: &'a ComputeClient<R>,
space: &'a Space,
banks: Vec<Handle>, xh: Handle,
wsch: Handle,
wdh: Handle,
wdmh: Handle,
outh: Handle,
wqs_len: usize,
x_len: usize,
wsc_len: usize,
wd_len: usize,
wdm_len: usize,
m: usize,
n: usize,
k: usize,
plane: usize,
lds_budget: usize,
reg_budget: usize,
oracle: Vec<f32>,
maxref: f32,
repeats: usize,
worst_rel: std::cell::Cell<f32>,
}
impl<'a, R: Runtime> CoopmatF16Eval<'a, R> {
#[allow(clippy::too_many_arguments)]
pub fn new(
client: &'a ComputeClient<R>,
space: &'a Space,
wqs_banks: &[Vec<u32>],
x: &[f32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
m: usize,
n: usize,
k: usize,
repeats: usize,
) -> Self {
assert!(
!wqs_banks.is_empty(),
"CoopmatF16Eval needs at least one weight bank"
);
let banks: Vec<Handle> = wqs_banks
.iter()
.map(|w| client.create_from_slice(u32::as_bytes(w)))
.collect();
let xh = client.create_from_slice(f32::as_bytes(x));
let wsch = client.create_from_slice(u32::as_bytes(wsc));
let wdh = client.create_from_slice(f32::as_bytes(wd));
let wdmh = client.create_from_slice(f32::as_bytes(wdm));
let outh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; m * n]));
let oracle = mmq_q4k_f16_ref(x, &wqs_banks[0], wsc, wd, wdm, m, n, k);
let maxref = oracle.iter().fold(0f32, |a, &v| a.max(v.abs())).max(1e-30);
let plane = client.properties().hardware.plane_size_max as usize;
Self {
client,
space,
banks,
xh,
wsch,
wdh,
wdmh,
outh,
wqs_len: wqs_banks[0].len(),
x_len: x.len(),
wsc_len: wsc.len(),
wd_len: wd.len(),
wdm_len: wdm.len(),
m,
n,
k,
plane,
lds_budget: 48 * 1024,
reg_budget: 16,
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, wm: usize, wn: usize, rm: usize, rn: usize) {
let bm = wm * rm * 16;
let bn = wn * rn * 16;
let grid = Grid::Static(self.n.div_ceil(bn) as u32, self.m.div_ceil(bm) as u32, 1);
let block = Block::new_1d((wm * wn * self.plane) as u32);
let target = Target::of(self.client);
unsafe {
mmq_q4k_coopmat_tile::launch_unchecked::<f32, R>(
self.client,
grid,
block,
ArrayArg::from_raw_parts(self.xh.clone(), self.x_len),
ArrayArg::from_raw_parts(bank.clone(), self.wqs_len),
ArrayArg::from_raw_parts(self.wsch.clone(), self.wsc_len),
ArrayArg::from_raw_parts(self.wdh.clone(), self.wd_len),
ArrayArg::from_raw_parts(self.wdmh.clone(), self.wdm_len),
ArrayArg::from_raw_parts(self.outh.clone(), self.m * self.n),
self.m,
self.n,
self.k,
wm,
wn,
rm,
rn,
self.plane,
target,
);
}
}
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 CoopmatF16Eval<'a, R> {
fn static_check(&self, cfg: &Config) -> Verdict {
let (bm, bn, nwarp) = mmq_geom(cfg, self.space);
let wg = nwarp * self.plane;
if wg > 1024 {
return Verdict::Reject(format!(
"workgroup width {wg} = nwarp {nwarp} x plane {} > 1024",
self.plane
));
}
let lds = mmq_coopmat_lds_bytes(cfg, self.space);
if lds > self.lds_budget {
return Verdict::Reject(format!(
"LDS {}KB (BM {bm} x BN {bn}) > {}KB budget",
lds / 1024,
self.lds_budget / 1024
));
}
let (_wm, _wn, rm, rn) = mmq_tile_cfg(cfg, self.space);
let frags = rm * rn;
if frags > self.reg_budget {
return Verdict::Reject(format!(
"register tile {frags} fragments/warp (rm {rm} x rn {rn}) > {} budget",
self.reg_budget
));
}
Verdict::Pass
}
fn measure(&self, cfg: &Config, iters: usize) -> f64 {
let (wm, wn, rm, rn) = mmq_tile_cfg(cfg, self.space);
self.dispatch(&self.banks[0], wm, wn, rm, rn);
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 > 5e-2 {
return f64::INFINITY;
}
let nb = self.banks.len();
for i in 0..(2 * nb) {
self.dispatch(&self.banks[i % nb], wm, wn, rm, rn);
}
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], wm, wn, rm, rn);
}
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 mmq_q4k_coopmat_hunt<R: Runtime>(
tuner: &Tuner,
device: &str,
m: usize,
n: usize,
k: usize,
eval: &CoopmatF16Eval<R>,
evo: &Evolution,
seed: u64,
) -> Evolved {
tuner.evolve(
device,
"mmq_q4k_coopmat",
&format!("m={m},n={n},k={k}"),
eval.space(),
eval,
evo,
seed,
)
}
#[allow(clippy::too_many_arguments)]
pub fn mmq_q4k_coopmat_autokernel<R: Runtime>(
client: &ComputeClient<R>,
tuner: &Tuner,
x: &[f32],
wqs: &[u32],
wsc: &[u32],
wd: &[f32],
wdm: &[f32],
m: usize,
n: usize,
k: usize,
) -> (Vec<f32>, String) {
let device = crate::tune::device_id(client);
let space = mmq_q4k_space();
let key = format!("m={m},n={n},k={k}");
let (wm, wn, rm, rn, name) = match tuner.cached_winner(&device, "mmq_q4k_coopmat", &key) {
Some(w) => match space.parse(&w) {
Some(cfg) => {
let (a, b, c, d) = mmq_tile_cfg(&cfg, &space);
(a, b, c, d, w)
}
None => {
let (a, b, c, d) = mmq_q4k_incumbent();
(a, b, c, d, "incumbent".to_string())
}
},
None => {
let (a, b, c, d) = mmq_q4k_incumbent();
(a, b, c, d, "incumbent".to_string())
}
};
let (out, _ms) =
mmq_q4k_coopmat_tile_run(client, x, wqs, wsc, wd, wdm, m, n, k, wm, wn, rm, rn, 1);
(out, name)
}
#[cfg(all(test, feature = "cpu"))]
mod tests {
use super::*;
use cubecl::cpu::{CpuDevice, CpuRuntime};
fn cpu_client() -> ComputeClient<CpuRuntime> {
CpuRuntime::client(&CpuDevice)
}
fn rel_to_max(got: &[f32], want: &[f32]) -> f32 {
let mut maxabs = 0f32;
let mut refmax = 1e-9f32;
for (g, w) in got.iter().zip(want) {
maxabs = maxabs.max((g - w).abs());
refmax = refmax.max(w.abs());
}
maxabs / refmax
}
#[test]
fn mmq_default_arm_is_bit_exact_on_cpu() {
let client = cpu_client();
for (m, n, k) in [(16usize, 16usize, 64usize), (32, 16, 128), (16, 32, 256)] {
let (xq, xs, wq, wd) = gen_mmq(m, n, k);
let (got, _) = mmq_q8_wmma_run_with::<CpuRuntime>(
&client,
&xq,
&xs,
&wq,
&wd,
m,
n,
k,
1,
Target::Cpu,
);
let want = mmq_q8_ref(&xq, &xs, &wq, &wd, m, n, k);
let gb: Vec<u32> = got.iter().map(|x| x.to_bits()).collect();
let wb: Vec<u32> = want.iter().map(|x| x.to_bits()).collect();
assert_eq!(
gb, wb,
"MMQ {m}x{n}x{k}: default arm != mmq_q8_ref (bit-exact gate)"
);
}
}
#[test]
fn mmq_blk_default_arm_matches_ref_on_cpu() {
let client = cpu_client();
let (m, n, k) = (32usize, 64usize, 128usize);
let (xq, xs, wq, wd) = gen_mmq(m, n, k);
let (got, _) = mmq_q8_wmma_blk_run_with::<CpuRuntime>(
&client,
&xq,
&xs,
&wq,
&wd,
m,
n,
k,
1,
Target::Cpu,
);
let want = mmq_q8_ref(&xq, &xs, &wq, &wd, m, n, k);
let rel = rel_to_max(&got, &want);
assert!(
rel < 1e-6,
"tiled MMQ {m}x{n}x{k}: rel_to_max={rel:.3e} vs mmq_q8_ref"
);
}
#[test]
fn mmq_q4k_affine_matches_ref_on_cpu() {
let client = cpu_client();
for (m, n, k) in [(32usize, 64usize, 256usize), (32, 64, 512)] {
let (xq, xs, xsum, wqs, wsc, wd, wdm) = gen_mmq_q4k(m, n, k);
let (got, _) = mmq_q4k_wmma_blk_run::<CpuRuntime>(
&client, &xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, m, n, k, 1,
);
let want = mmq_q4k_ref(&xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, m, n, k);
let rel = rel_to_max(&got, &want);
assert!(
rel < 1e-6,
"affine Q4_K MMQ {m}x{n}x{k}: rel_to_max={rel:.3e} vs mmq_q4k_ref"
);
}
}
#[test]
fn mmq_q4k_rt_matches_ref_with_tails_on_cpu() {
let client = cpu_client();
for (m, n, k) in [
(32usize, 64usize, 256usize),
(32, 64, 512),
(17, 64, 256),
(32, 50, 256),
(17, 50, 512),
(1, 64, 256),
] {
let (xq, xs, xsum, wqs, wsc, wd, wdm) = gen_mmq_q4k(m, n, k);
let (got, _) = mmq_q4k_wmma_rt_run::<CpuRuntime>(
&client, &xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, m, n, k, 1,
);
let want = mmq_q4k_ref(&xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, m, n, k);
let rel = rel_to_max(&got, &want);
assert!(
rel < 1e-6,
"runtime Q4_K MMQ {m}x{n}x{k}: rel_to_max={rel:.3e} vs mmq_q4k_ref"
);
}
}
#[test]
fn moe_expert_rows_groups_slots_stably_on_cpu() {
let client = cpu_client();
let n_experts = 5usize;
let cap = 8usize;
let ids: Vec<u32> = vec![0, 2, 0, 4, 1, 0, 2, 4, 0, 1, 2, 0];
let (counts, rows) = moe_expert_rows_run::<CpuRuntime>(&client, &ids, n_experts, cap);
let mut want: Vec<Vec<u32>> = vec![Vec::new(); n_experts];
for (s, &e) in ids.iter().enumerate() {
want[e as usize].push(s as u32);
}
assert_eq!(
counts.iter().sum::<u32>() as usize,
ids.len(),
"counts must cover every slot"
);
for e in 0..n_experts {
assert_eq!(counts[e] as usize, want[e].len(), "expert {e} count");
let got = &rows[e * cap..e * cap + want[e].len()];
assert_eq!(
got,
&want[e][..],
"expert {e} slot list (must be ascending slot order)"
);
}
}
#[test]
fn mmq_q4k_id_matches_ref_with_tails_on_cpu() {
let client = cpu_client();
for (nslots, n, k) in [
(32usize, 64usize, 256usize),
(32, 64, 512),
(17, 64, 256),
(32, 50, 256),
(17, 50, 512),
(1, 64, 256),
] {
let (xq, xs, xsum, wqs, wsc, wd, wdm) = gen_mmq_q4k(nslots, n, k);
let ids = vec![0u32; nslots];
let (got, _) = mmq_q4k_id_run::<CpuRuntime>(
&client,
&xq,
&xs,
&xsum,
&wqs,
&wsc,
&wd,
&wdm,
&ids,
nslots,
1,
nslots,
n,
k,
nslots.div_ceil(32),
1,
);
let want = mmq_q4k_id_ref(&xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, &ids, nslots, n, k);
let rel = rel_to_max(&got, &want);
assert!(
rel < 1e-6,
"expert-indexed Q4_K MMQ {nslots}x{n}x{k}: rel_to_max={rel:.3e} vs mmq_q4k_id_ref"
);
}
}
#[test]
fn mmq_q4k_id_is_bit_identical_to_rt_for_one_expert_on_cpu() {
let client = cpu_client();
let (m, n, k) = (32usize, 64usize, 512usize);
let (xq, xs, xsum, wqs, wsc, wd, wdm) = gen_mmq_q4k(m, n, k);
let ids = vec![0u32; m];
let (got, _) = mmq_q4k_id_run::<CpuRuntime>(
&client,
&xq,
&xs,
&xsum,
&wqs,
&wsc,
&wd,
&wdm,
&ids,
m,
1,
m,
n,
k,
m.div_ceil(32),
1,
);
let (want, _) = mmq_q4k_wmma_rt_run::<CpuRuntime>(
&client, &xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, m, n, k, 1,
);
let gb: Vec<u32> = got.iter().map(|x| x.to_bits()).collect();
let wb: Vec<u32> = want.iter().map(|x| x.to_bits()).collect();
assert_eq!(
gb, wb,
"expert-indexed MMQ must be bit-identical to mmq_q4k_wmma_rt at E=1"
);
}
#[test]
fn mmq_q4k_id_grid_stride_covers_every_row_on_cpu() {
let client = cpu_client();
let (nslots, n, k) = (128usize, 64usize, 256usize);
let (xq, xs, xsum, wqs, wsc, wd, wdm) = gen_mmq_q4k(nslots, n, k);
let ids = vec![0u32; nslots];
let want = mmq_q4k_id_ref(&xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, &ids, nslots, n, k);
for ytiles in [1usize, 2, 3, 4] {
let (got, _) = mmq_q4k_id_run::<CpuRuntime>(
&client, &xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, &ids, nslots, 1, nslots, n, k,
ytiles, 1,
);
let rel = rel_to_max(&got, &want);
assert!(
rel < 1e-6,
"grid-stride ytiles={ytiles}: rel_to_max={rel:.3e} vs mmq_q4k_id_ref"
);
}
}
#[test]
fn production_run_selects_oracle_on_cpu() {
let client = cpu_client();
let (m, n, k) = (16usize, 16usize, 64usize);
let (xq, xs, wq, wd) = gen_mmq(m, n, k);
let (got, _) = mmq_q8_wmma_run::<CpuRuntime>(&client, &xq, &xs, &wq, &wd, m, n, k, 1);
let want = mmq_q8_ref(&xq, &xs, &wq, &wd, m, n, k);
let gb: Vec<u32> = got.iter().map(|x| x.to_bits()).collect();
let wb: Vec<u32> = want.iter().map(|x| x.to_bits()).collect();
assert_eq!(gb, wb);
}
#[test]
fn mmq_q4k_tile_matches_ref_on_cpu() {
let client = cpu_client();
for (wm, wn, rm, rn) in [
(1, 1, 1, 1),
(2, 2, 1, 1),
(1, 1, 2, 2),
(2, 1, 1, 2),
(1, 2, 2, 1),
(2, 2, 2, 2),
] {
let bm = wm * rm * 16;
let bn = wn * rn * 16;
for k in [256usize, 512] {
let (xq, xs, xsum, wqs, wsc, wd, wdm) = gen_mmq_q4k(bm, bn, k);
let (got, _) = mmq_q4k_wmma_tile_run::<CpuRuntime>(
&client, &xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, bm, bn, k, wm, wn, rm, rn, 1,
);
let want = mmq_q4k_ref(&xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, bm, bn, k);
let rel = rel_to_max(&got, &want);
assert!(
rel < 1e-6,
"tile wm{wm} wn{wn} rm{rm} rn{rn} {bm}x{bn}x{k}: rel_to_max={rel:.3e}"
);
}
}
}
#[test]
fn mmq_q4k_coopmat_tile_matches_f16_ref_on_cpu() {
let client = cpu_client();
for (wm, wn, rm, rn) in [
(1, 1, 1, 1),
(2, 2, 1, 1),
(1, 1, 2, 2),
(2, 2, 2, 2),
(2, 2, 4, 4),
] {
let bm = wm * rm * 16;
let bn = wn * rn * 16;
for k in [256usize, 512] {
let (_xq, _xs, _xsum, wqs, wsc, wd, wdm) = gen_mmq_q4k(bm, bn, k);
let mut s = 0x1234_5678_9ABC_DEF0u64;
let x: Vec<f32> = (0..bm * k)
.map(|_| {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
(s % 2000) as f32 / 1000.0 - 1.0
})
.collect();
let (got, _) = mmq_q4k_coopmat_tile_run::<CpuRuntime>(
&client, &x, &wqs, &wsc, &wd, &wdm, bm, bn, k, wm, wn, rm, rn, 1,
);
let want = mmq_q4k_f16_ref(&x, &wqs, &wsc, &wd, &wdm, bm, bn, k);
let rel = rel_to_max(&got, &want);
eprintln!(
"[f16 coopmat] wm{wm} wn{wn} rm{rm} rn{rn} {bm}x{bn}x{k}: rel_to_max={rel:.3e}"
);
assert!(
rel < 5e-2,
"f16 coopmat wm{wm} wn{wn} rm{rm} rn{rn} {bm}x{bn}x{k}: rel_to_max={rel:.3e}"
);
}
}
}
#[cfg(feature = "cpu")]
#[test]
fn mmq_q4k_coopmat_space_search_cpu() {
let (m, n, k) = (16usize, 16usize, 256usize);
let (_xq, _xs, _xsum, wqs, wsc, wd, wdm) = gen_mmq_q4k(m, n, k);
let mut s = 0xABCD_1234u64;
let x: Vec<f32> = (0..m * k)
.map(|_| {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
(s % 2000) as f32 / 1000.0 - 1.0
})
.collect();
let mut wqs2 = wqs.clone();
for (i, v) in wqs2.iter_mut().enumerate() {
*v ^= 0x9E37_79B1u32.wrapping_mul(i as u32 + 1);
}
let banks = vec![wqs.clone(), wqs2];
let space = mmq_q4k_space();
let client = cpu_client();
let eval = CoopmatF16Eval::new(&client, &space, &banks, &x, &wsc, &wd, &wdm, m, n, k, 1);
let dir = std::env::temp_dir().join(format!(
"hk-coopmat-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 = mmq_q4k_coopmat_hunt(&tuner, "cpu", m, n, 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() < 5e-2,
"a measured tile diverged from the oracle: {:.2e}",
eval.worst_rel()
);
eprintln!(
"[coopmat f16 hunt CPU] winner={} evaluated={} measured={} rejected={} worst_rel={:.2e}",
r.winner, rep.evaluated, rep.measured.len(), rep.rejected.len(), eval.worst_rel()
);
let r2 = mmq_q4k_coopmat_hunt(&tuner, "cpu", m, n, 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 mmq_q4k_space_search_cpu() {
let (m, n, k) = (16usize, 16usize, 256usize);
let (xq, xs, xsum, wqs, wsc, wd, wdm) = gen_mmq_q4k(m, n, k);
let mut wqs2 = wqs.clone();
for (i, v) in wqs2.iter_mut().enumerate() {
*v ^= 0x9E37_79B1u32.wrapping_mul(i as u32 + 1);
}
let banks = vec![wqs.clone(), wqs2];
let space = mmq_q4k_space();
let feasible = space.enumerate();
assert!(!feasible.is_empty(), "empty feasible space");
let client = cpu_client();
let eval = MmqQ4kEval::new(
&client, &space, &banks, &xq, &xs, &xsum, &wsc, &wd, &wdm, m, n, k, 1,
);
let dir = std::env::temp_dir().join(format!(
"hk-mmq-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 = mmq_q4k_hunt(&tuner, "cpu", m, n, 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() < 1e-3,
"a measured tile diverged from the oracle: {:.2e}",
eval.worst_rel()
);
eprintln!(
"[mmq hunt CPU] winner={} evaluated={} measured={} rejected={} worst_rel={:.2e}",
r.winner,
rep.evaluated,
rep.measured.len(),
rep.rejected.len(),
eval.worst_rel()
);
let r2 = mmq_q4k_hunt(&tuner, "cpu", m, n, 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 mmq_q4k_autokernel_prefers_cached_genome_cpu() {
let (m, n, k) = (16usize, 16usize, 256usize);
let (xq, xs, xsum, wqs, wsc, wd, wdm) = gen_mmq_q4k(m, n, k);
let want = mmq_q4k_ref(&xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, m, n, k);
let client = cpu_client();
let device = crate::tune::device_id(&client);
let dir = std::env::temp_dir().join(format!(
"hk-mmq-auto-{}-{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
let tuner = Tuner::new(&dir);
let (out0, name0) = mmq_q4k_autokernel(
&client, &tuner, &xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, m, n, k,
);
assert_eq!(
name0, "incumbent",
"a cold cache must serve the incumbent tile"
);
assert!(
rel_to_max(&out0, &want) < 1e-6,
"incumbent tile diverged from the oracle"
);
let space = mmq_q4k_space();
let banks = vec![wqs.clone()];
let eval = MmqQ4kEval::new(
&client, &space, &banks, &xq, &xs, &xsum, &wsc, &wd, &wdm, m, n, k, 1,
);
let evo = Evolution::new()
.population(8)
.generations(3)
.measure_iters(1);
let r = mmq_q4k_hunt(&tuner, &device, m, n, k, &eval, &evo, 0xC0FFEE);
assert!(!r.from_cache, "the seeding hunt must actually run");
assert!(
r.report
.as_ref()
.expect("a miss carries the trail")
.best_ms
.is_finite(),
"the seeding hunt found no measurable winner"
);
let (out1, name1) = mmq_q4k_autokernel(
&client, &tuner, &xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, m, n, k,
);
assert_eq!(
name1, r.winner,
"the autokernel must replay the hunted winner, not the incumbent"
);
assert_ne!(
name1, "incumbent",
"the tuned dispatch must differ from the cold-start label"
);
assert!(
rel_to_max(&out1, &want) < 1e-6,
"the tuned tile diverged from the oracle"
);
std::fs::remove_dir_all(&dir).ok();
}
}