use crate::prelude::*;
use half::f16;
pub const BR: usize = 16;
pub const BC: usize = 16;
#[allow(clippy::too_many_arguments)]
#[kernel(targets(cuda, rocm, vulkan, metal, cpu), unchecked)]
pub fn flash_attn<F: Float>(
q: &Array<F>,
k: &Array<F>,
v: &Array<F>,
out: &mut Array<F>,
scale: &Array<F>,
meta: &Array<u32>,
#[comptime] d: usize,
#[comptime] br: usize,
#[comptime] bc: usize,
#[comptime] plane: usize,
#[comptime] target: Target,
) {
let cube = CUBE_POS as usize + meta[8] as usize; let lane = UNIT_POS as usize; let sc = scale[0];
let seq_q = meta[0] as usize;
let seq_k = meta[1] as usize;
let n_heads = meta[2] as usize;
let n_kv = meta[3] as usize;
let causal = meta[4];
let kv_batch_stride = meta[5] as usize;
let kv_head_stride = meta[6] as usize;
let key_stride = meta[7] as usize;
let n_qtiles = (seq_q + br - 1) / br;
let tiles_per_head = n_qtiles;
let hq = n_heads * tiles_per_head;
let b_i = cube / hq;
let rem = cube - b_i * hq;
let h = rem / tiles_per_head;
let qt = rem % tiles_per_head;
let kv = h / (n_heads / n_kv);
let q_head_base = ((b_i * n_heads + h) * seq_q) * d; let kvbase = b_i * kv_batch_stride + kv * kv_head_stride;
let per_s = br * bc / plane; let per_o = br * d / plane; let per_q = br * d / plane;
let mut qsh = SharedMemory::<f16>::new(br * d);
let mut sf = SharedMemory::<F>::new(br * bc);
let mut of = SharedMemory::<F>::new(br * d);
let mut mf = SharedMemory::<F>::new(br);
let mut lf = SharedMemory::<F>::new(br);
let mut emf = SharedMemory::<F>::new(br);
for e in 0..per_o {
of[lane * per_o + e] = F::new(0.0);
}
if lane < br {
mf[lane] = F::new(-3.4e38);
lf[lane] = F::new(0.0);
}
for e in 0..per_q {
let idx = lane * per_q + e;
let r = idx / d;
let dd = idx % d;
let qpos = qt * br + r;
let val = if qpos < seq_q {
q[q_head_base + qpos * d + dd]
} else {
F::new(0.0)
};
qsh[idx] = f16::cast_from(val);
}
sync_cube();
let n_jt = (seq_k + bc - 1) / bc;
let n_jt_eff = if causal == 1 {
let bound = (qt * br + br - 1) / bc + 1;
if bound < n_jt {
bound
} else {
n_jt
}
} else {
n_jt
};
for j in 0..n_jt_eff {
island! {
cuda | rocm | vulkan => {
let mut ksh = SharedMemory::<f16>::new(bc * d);
let per_k = bc * d / plane;
for e in 0..per_k {
let idx = lane * per_k + e;
let c = idx / d;
let dd = idx % d;
let kpos = j * bc + c;
let val = if kpos < seq_k { k[kvbase + kpos * key_stride + dd] } else { F::new(0.0) };
ksh[idx] = f16::cast_from(val);
}
sync_cube();
for cg in 0..bc / 16usize {
let cacc = cmma::Matrix::<F>::from_value(
cmma::MatrixIdent::Accumulator, 16usize, 16usize, 16usize, cmma::MatrixLayout::Undefined, F::new(0.0),
);
for dk in 0..d / 16usize {
let a = cmma::Matrix::<f16>::from_slice(
cmma::MatrixIdent::A, 16usize, 16usize, 16usize, cmma::MatrixLayout::RowMajor,
&qsh.to_slice().slice(dk * 16usize, br * d), d as u32,
);
let b = cmma::Matrix::<f16>::from_slice(
cmma::MatrixIdent::B, 16usize, 16usize, 16usize, cmma::MatrixLayout::ColMajor,
&ksh.to_slice().slice(cg * 16usize * d + dk * 16usize, bc * d), d as u32,
);
cmma::execute::<f16, f16, F, F>(&a, &b, &cacc, &cacc);
}
cmma::store(&mut sf.to_slice_mut().slice_mut(cg * 16usize, br * bc), &cacc, bc as u32, cmma::MatrixLayout::RowMajor);
}
sync_cube();
for e in 0..per_s {
let idx = lane * per_s + e;
let r = idx / bc;
let c = idx % bc;
let qpos = qt * br + r;
let kpos = j * bc + c;
let masked = causal == 1 && kpos > qpos;
if qpos < seq_q && kpos < seq_k && !masked {
sf[idx] = sf[idx] * sc;
} else {
sf[idx] = F::new(-3.4e38);
}
}
}
metal => {
let mut ksh = SharedMemory::<f16>::new(bc * d);
let per_k = bc * d / plane;
for e in 0..per_k {
let idx = lane * per_k + e;
let c = idx / d;
let dd = idx % d;
let kpos = j * bc + c;
let val = if kpos < seq_k { k[kvbase + kpos * key_stride + dd] } else { F::new(0.0) };
ksh[idx] = f16::cast_from(val);
}
sync_cube();
for ti in 0..br / 8usize {
for tj in 0..bc / 8usize {
let cacc = cmma::Matrix::<F>::from_value(
cmma::MatrixIdent::Accumulator, 8usize, 8usize, 8usize, cmma::MatrixLayout::Undefined, F::new(0.0),
);
for dk in 0..d / 8usize {
let a = cmma::Matrix::<f16>::from_slice(
cmma::MatrixIdent::A, 8usize, 8usize, 8usize, cmma::MatrixLayout::RowMajor,
&qsh.to_slice().slice(ti * 8usize * d + dk * 8usize, br * d), d as u32,
);
let b = cmma::Matrix::<f16>::from_slice(
cmma::MatrixIdent::B, 8usize, 8usize, 8usize, cmma::MatrixLayout::ColMajor,
&ksh.to_slice().slice(tj * 8usize * d + dk * 8usize, bc * d), d as u32,
);
cmma::execute::<f16, f16, F, F>(&a, &b, &cacc, &cacc);
}
cmma::store(
&mut sf.to_slice_mut().slice_mut(ti * 8usize * bc + tj * 8usize, br * bc),
&cacc, bc as u32, cmma::MatrixLayout::RowMajor,
);
}
}
sync_cube();
for e in 0..per_s {
let idx = lane * per_s + e;
let r = idx / bc;
let c = idx % bc;
let qpos = qt * br + r;
let kpos = j * bc + c;
let masked = causal == 1 && kpos > qpos;
if qpos < seq_q && kpos < seq_k && !masked {
sf[idx] = sf[idx] * sc;
} else {
sf[idx] = F::new(-3.4e38);
}
}
}
default => {
for e in 0..per_s {
let idx = lane * per_s + e;
let r = idx / bc;
let c = idx % bc;
let qpos = qt * br + r;
let kpos = j * bc + c;
let masked = causal == 1 && kpos > qpos;
if qpos < seq_q && kpos < seq_k && !masked {
let qb = q_head_base + qpos * d;
let kb = kvbase + kpos * key_stride;
let mut acc = F::new(0.0);
for dd in 0..d {
acc += q[qb + dd] * k[kb + dd];
}
sf[idx] = acc * sc;
} else {
sf[idx] = F::new(-3.4e38);
}
}
}
};
sync_cube();
if lane < br {
let r = lane;
let qpos = qt * br + r;
if qpos < seq_q {
let mut rowmax = F::new(-3.4e38);
for c in 0..bc {
let s = sf[r * bc + c];
if s > rowmax {
rowmax = s;
}
}
let m_old = mf[r];
let mut m_new = m_old;
if rowmax > m_new {
m_new = rowmax;
}
let em = (m_old - m_new).exp();
let mut psum = F::new(0.0);
for c in 0..bc {
let p = (sf[r * bc + c] - m_new).exp();
sf[r * bc + c] = p;
psum += p;
}
lf[r] = lf[r] * em + psum;
mf[r] = m_new;
emf[r] = em;
} else {
emf[r] = F::new(1.0);
}
}
sync_cube();
island! {
cuda | rocm | vulkan => {
let mut psh = SharedMemory::<f16>::new(br * bc);
let mut vsh = SharedMemory::<f16>::new(bc * d);
let mut ovt = SharedMemory::<F>::new(br * d);
let per_p = br * bc / plane;
let per_v = bc * d / plane;
for e in 0..per_p {
let idx = lane * per_p + e;
let r = idx / bc;
let qpos = qt * br + r;
let val = if qpos < seq_q { sf[idx] } else { F::new(0.0) };
psh[idx] = f16::cast_from(val);
}
for e in 0..per_v {
let idx = lane * per_v + e;
let c = idx / d;
let dd = idx % d;
let kpos = j * bc + c;
let val = if kpos < seq_k { v[kvbase + kpos * key_stride + dd] } else { F::new(0.0) };
vsh[idx] = f16::cast_from(val);
}
sync_cube();
for dn in 0..d / 16usize {
let cacc = cmma::Matrix::<F>::from_value(
cmma::MatrixIdent::Accumulator, 16usize, 16usize, 16usize, cmma::MatrixLayout::Undefined, F::new(0.0),
);
for kc in 0..bc / 16usize {
let a = cmma::Matrix::<f16>::from_slice(
cmma::MatrixIdent::A, 16usize, 16usize, 16usize, cmma::MatrixLayout::RowMajor,
&psh.to_slice().slice(kc * 16usize, br * bc), bc as u32,
);
let b = cmma::Matrix::<f16>::from_slice(
cmma::MatrixIdent::B, 16usize, 16usize, 16usize, cmma::MatrixLayout::RowMajor,
&vsh.to_slice().slice(kc * 16usize * d + dn * 16usize, bc * d), d as u32,
);
cmma::execute::<f16, f16, F, F>(&a, &b, &cacc, &cacc);
}
cmma::store(&mut ovt.to_slice_mut().slice_mut(dn * 16usize, br * d), &cacc, d as u32, cmma::MatrixLayout::RowMajor);
}
sync_cube();
for e in 0..per_o {
let idx = lane * per_o + e;
let r = idx / d;
of[idx] = of[idx] * emf[r] + ovt[idx];
}
}
metal => {
let mut psh = SharedMemory::<f16>::new(br * bc);
let mut vsh = SharedMemory::<f16>::new(bc * d);
let mut ovt = SharedMemory::<F>::new(br * d);
let per_p = br * bc / plane;
let per_v = bc * d / plane;
for e in 0..per_p {
let idx = lane * per_p + e;
let r = idx / bc;
let qpos = qt * br + r;
let val = if qpos < seq_q { sf[idx] } else { F::new(0.0) };
psh[idx] = f16::cast_from(val);
}
for e in 0..per_v {
let idx = lane * per_v + e;
let c = idx / d;
let dd = idx % d;
let kpos = j * bc + c;
let val = if kpos < seq_k { v[kvbase + kpos * key_stride + dd] } else { F::new(0.0) };
vsh[idx] = f16::cast_from(val);
}
sync_cube();
for ti in 0..br / 8usize {
for tn in 0..d / 8usize {
let cacc = cmma::Matrix::<F>::from_value(
cmma::MatrixIdent::Accumulator, 8usize, 8usize, 8usize, cmma::MatrixLayout::Undefined, F::new(0.0),
);
for k8 in 0..bc / 8usize {
let a = cmma::Matrix::<f16>::from_slice(
cmma::MatrixIdent::A, 8usize, 8usize, 8usize, cmma::MatrixLayout::RowMajor,
&psh.to_slice().slice(ti * 8usize * bc + k8 * 8usize, br * bc), bc as u32,
);
let b = cmma::Matrix::<f16>::from_slice(
cmma::MatrixIdent::B, 8usize, 8usize, 8usize, cmma::MatrixLayout::RowMajor,
&vsh.to_slice().slice(k8 * 8usize * d + tn * 8usize, bc * d), d as u32,
);
cmma::execute::<f16, f16, F, F>(&a, &b, &cacc, &cacc);
}
cmma::store(
&mut ovt.to_slice_mut().slice_mut(ti * 8usize * d + tn * 8usize, br * d),
&cacc, d as u32, cmma::MatrixLayout::RowMajor,
);
}
}
sync_cube();
for e in 0..per_o {
let idx = lane * per_o + e;
let r = idx / d;
of[idx] = of[idx] * emf[r] + ovt[idx];
}
}
default => {
for e in 0..per_o {
let idx = lane * per_o + e;
let r = idx / d;
let dd = idx % d;
let mut o = of[idx] * emf[r];
for c in 0..bc {
let kpos = j * bc + c;
if kpos < seq_k {
let kb = kvbase + kpos * key_stride;
o += sf[r * bc + c] * v[kb + dd];
}
}
of[idx] = o;
}
}
};
sync_cube();
}
for e in 0..per_o {
let idx = lane * per_o + e;
let r = idx / d;
let dd = idx % d;
let qpos = qt * br + r;
if qpos < seq_q {
out[q_head_base + qpos * d + dd] = of[idx] / lf[r];
}
}
}
pub fn flash_attn_cubes(b: usize, n_heads: usize, seq_q: usize) -> usize {
b * n_heads * seq_q.div_ceil(BR)
}
#[allow(clippy::too_many_arguments)]
pub fn flash_attn_run<R: Runtime>(
client: &ComputeClient<R>,
q: &[f32],
k: &[f32],
v: &[f32],
b: usize,
n_heads: usize,
n_kv: usize,
seq_q: usize,
seq_k: usize,
kv_seq_pad: usize,
d: usize,
causal: bool,
) -> Vec<f32> {
let plane = (client.properties().hardware.plane_size_max as usize).max(BR);
let cubes = flash_attn_cubes(b, n_heads, seq_q);
let target = Target::of(client);
flash_attn_launch(
client, q, k, v, b, n_heads, n_kv, seq_q, seq_k, kv_seq_pad, d, causal, plane, 0, cubes,
target,
)
}
#[allow(clippy::too_many_arguments)]
pub fn flash_attn_launch<R: Runtime>(
client: &ComputeClient<R>,
q: &[f32],
k: &[f32],
v: &[f32],
b: usize,
n_heads: usize,
n_kv: usize,
seq_q: usize,
seq_k: usize,
kv_seq_pad: usize,
d: usize,
causal: bool,
plane: usize,
cube_base: usize,
cube_count: usize,
target: Target,
) -> Vec<f32> {
let scale = 1.0f32 / (d as f32).sqrt();
let meta = [
seq_q as u32,
seq_k as u32,
n_heads as u32,
n_kv as u32,
causal as u32,
(n_kv * kv_seq_pad * d) as u32,
(kv_seq_pad * d) as u32,
d as u32,
cube_base as u32,
];
let qh = client.create_from_slice(f32::as_bytes(q));
let kh = client.create_from_slice(f32::as_bytes(k));
let vh = client.create_from_slice(f32::as_bytes(v));
let sh = client.create_from_slice(f32::as_bytes(&[scale]));
let mh = client.create_from_slice(u32::as_bytes(&meta));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; b * n_heads * seq_q * d]));
unsafe {
flash_attn::launch_unchecked::<f32, R>(
client,
Grid::Static(cube_count as u32, 1, 1),
Block::new_1d(plane as u32),
ArrayArg::from_raw_parts(qh.clone(), q.len()),
ArrayArg::from_raw_parts(kh.clone(), k.len()),
ArrayArg::from_raw_parts(vh.clone(), v.len()),
ArrayArg::from_raw_parts(oh.clone(), b * n_heads * seq_q * d),
ArrayArg::from_raw_parts(sh.clone(), 1),
ArrayArg::from_raw_parts(mh.clone(), 9),
d,
BR,
BC,
plane,
target,
);
}
f32::from_bytes(&client.read_one_unchecked(oh)).to_vec()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::attn::sdpa_ref;
fn rnd(n: usize, seed: u64) -> Vec<f32> {
let mut s = seed;
(0..n)
.map(|_| {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
(s % 2000) as f32 / 1000.0 - 1.0
})
.collect()
}
fn scale_rel(got: &[f32], want: &[f32]) -> f32 {
let refmax = want.iter().fold(0.0f32, |a, x| a.max(x.abs())).max(1e-6);
let maxd = got
.iter()
.zip(want)
.fold(0.0f32, |a, (g, w)| a.max((g - w).abs()));
maxd / refmax
}
#[allow(clippy::too_many_arguments)]
fn gate<R: Runtime>(
c: &ComputeClient<R>,
nh: usize,
nkv: usize,
sq: usize,
sk: usize,
d: usize,
causal: bool,
plane: usize,
tag: &str,
) {
let q = rnd(nh * sq * d, 0x1234_5678);
let k = rnd(nkv * sk * d, 0x9ABC_DEF0);
let v = rnd(nkv * sk * d, 0x0FED_CBA9);
let ncubes = flash_attn_cubes(1, nh, sq);
let mut got = vec![0.0f32; nh * sq * d];
for n in 0..ncubes {
let part = flash_attn_launch::<R>(
c,
&q,
&k,
&v,
1,
nh,
nkv,
sq,
sk,
sk,
d,
causal,
plane,
n,
1,
Target::Cpu,
);
for (g, p) in got.iter_mut().zip(&part) {
*g += *p;
}
}
let want = sdpa_ref(&q, &k, &v, nh, nkv, sq, sk, d, causal);
let rel = scale_rel(&got, &want);
eprintln!("[flash {tag}] nh{nh}/nkv{nkv} sq{sq} sk{sk} d{d} causal{causal} plane{plane} cubes{ncubes} scale_rel={rel:.2e}");
assert!(
rel < 2e-3,
"flash {tag}: scale_rel {rel} exceeds 2e-3 vs materialized sdpa_ref"
);
}
#[cfg(feature = "cpu")]
#[test]
fn flash_matches_materialized_ref_on_cpu() {
use cubecl::cpu::{CpuDevice, CpuRuntime};
let c = CpuRuntime::client(&CpuDevice::default());
gate::<CpuRuntime>(&c, 4, 2, 1, 1, 32, false, 32, "decode kv1");
gate::<CpuRuntime>(&c, 4, 2, 1, 17, 32, false, 32, "decode kv17 (tail)");
gate::<CpuRuntime>(&c, 4, 2, 1, 128, 32, false, 32, "decode kv128");
gate::<CpuRuntime>(&c, 8, 2, 1, 512, 64, false, 32, "decode kv512 GQA4");
gate::<CpuRuntime>(&c, 4, 2, 16, 128, 32, true, 32, "prefill qt1 causal GQA2");
gate::<CpuRuntime>(
&c,
4,
2,
16,
128,
32,
false,
32,
"prefill qt1 noncausal GQA2",
);
gate::<CpuRuntime>(&c, 4, 4, 16, 64, 64, true, 32, "prefill qt1 MHA causal d64");
gate::<CpuRuntime>(
&c,
4,
2,
48,
48,
32,
true,
32,
"prefill 3-tile causal GQA2 (aligned)",
);
gate::<CpuRuntime>(
&c,
6,
3,
40,
40,
64,
true,
32,
"prefill causal GQA2 tail sq40 d64",
);
gate::<CpuRuntime>(
&c,
2,
1,
512,
512,
128,
true,
32,
"prefill 512 causal MHA d128",
);
}
#[cfg(feature = "vulkan")]
#[test]
fn flash_coopmat_matches_ref_on_vulkan() {
use cubecl::wgpu::{WgpuDevice, WgpuRuntime};
let c = WgpuRuntime::client(&WgpuDevice::default());
let mut fails = 0usize;
let mut gate_gpu = |nh: usize,
nkv: usize,
sq: usize,
sk: usize,
d: usize,
causal: bool,
tag: &str| {
let q = rnd(nh * sq * d, 0x1234_5678);
let k = rnd(nkv * sk * d, 0x9ABC_DEF0);
let v = rnd(nkv * sk * d, 0x0FED_CBA9);
let got =
flash_attn_run::<WgpuRuntime>(&c, &q, &k, &v, 1, nh, nkv, sq, sk, sk, d, causal);
let want = sdpa_ref(&q, &k, &v, nh, nkv, sq, sk, d, causal);
let rel = scale_rel(&got, &want);
let nonfin = got.iter().filter(|x| !x.is_finite()).count();
let where_ = if rel < 2e-3 {
String::new()
} else {
got.iter()
.zip(&want)
.enumerate()
.find(|(_, (g, w))| (**g - **w).abs() > 1e-2 || !g.is_finite())
.map(|(i, (g, w))| {
let (h, rem) = (i / (sq * d), i % (sq * d));
format!(
" first-bad@[h{h},q{},d{}] got={g:.3e} want={w:.3e}",
rem / d,
rem % d
)
})
.unwrap_or_default()
};
eprintln!("[flash-vk {tag}] nh{nh}/nkv{nkv} sq{sq} sk{sk} d{d} causal{causal} scale_rel={rel:.2e} nonfinite={nonfin}{where_}");
if !(rel < 2e-2) {
fails += 1;
}
};
gate_gpu(4, 2, 1, 1, 32, false, "decode kv1");
gate_gpu(4, 2, 1, 17, 32, false, "decode kv17 (ragged tail)");
gate_gpu(4, 2, 1, 128, 32, false, "decode kv128");
gate_gpu(8, 2, 1, 512, 64, false, "decode kv512 GQA4");
gate_gpu(4, 2, 16, 128, 32, true, "prefill qt1 causal GQA2");
gate_gpu(4, 2, 16, 128, 32, false, "prefill qt1 noncausal GQA2");
gate_gpu(4, 4, 16, 64, 64, true, "prefill qt1 MHA causal d64");
gate_gpu(
4,
2,
48,
48,
32,
true,
"prefill 3-tile causal GQA2 (aligned)",
);
gate_gpu(6, 3, 40, 40, 64, true, "prefill causal GQA2 tail sq40 d64");
gate_gpu(2, 1, 512, 512, 128, true, "prefill 512 causal MHA d128");
assert_eq!(
fails, 0,
"flash coopmat: {fails} shape(s) exceeded 2e-2 vs materialized sdpa_ref"
);
}
#[cfg(feature = "cpu")]
#[test]
fn production_run_selects_scalar_oracle_on_cpu() {
use cubecl::cpu::{CpuDevice, CpuRuntime};
let c = CpuRuntime::client(&CpuDevice::default());
let (nh, nkv, sq, sk, d) = (1usize, 1usize, 16usize, 96usize, 64usize);
let q = rnd(nh * sq * d, 0x2222_3333);
let k = rnd(nkv * sk * d, 0x4444_5555);
let v = rnd(nkv * sk * d, 0x6666_7777);
let got = flash_attn_run::<CpuRuntime>(&c, &q, &k, &v, 1, nh, nkv, sq, sk, sk, d, true);
let want = sdpa_ref(&q, &k, &v, nh, nkv, sq, sk, d, true);
let rel = scale_rel(&got, &want);
eprintln!("[flash production] scalar-on-cpu scale_rel={rel:.2e}");
assert!(rel < 2e-3, "production run on cpu: scale_rel {rel}");
}
#[cfg(feature = "cuda")]
#[test]
fn flash_coopmat_matches_ref_on_cuda() {
use cubecl::cuda::{CudaDevice, CudaRuntime};
let c = CudaRuntime::client(&CudaDevice::default());
eprintln!(
"[flash-cuda] device plane_size_max={}",
c.properties().hardware.plane_size_max
);
let shapes: [(usize, usize, usize, usize, usize, bool, &str); 10] = [
(4, 2, 1, 1, 32, false, "decode kv1"),
(4, 2, 1, 17, 32, false, "decode kv17 (ragged tail)"),
(4, 2, 1, 128, 32, false, "decode kv128"),
(8, 2, 1, 512, 64, false, "decode kv512 GQA4"),
(4, 2, 16, 128, 32, true, "prefill qt1 causal GQA2"),
(4, 2, 16, 128, 32, false, "prefill qt1 noncausal GQA2"),
(4, 4, 16, 64, 64, true, "prefill qt1 MHA causal d64"),
(
4,
2,
48,
48,
32,
true,
"prefill 3-tile causal GQA2 (aligned)",
),
(6, 3, 40, 40, 64, true, "prefill causal GQA2 tail sq40 d64"),
(2, 1, 512, 512, 128, true, "prefill 512 causal MHA d128"),
];
let mut fails = 0usize;
for (nh, nkv, sq, sk, d, causal, tag) in shapes {
let q = rnd(nh * sq * d, 0x1234_5678);
let k = rnd(nkv * sk * d, 0x9ABC_DEF0);
let v = rnd(nkv * sk * d, 0x0FED_CBA9);
let got =
flash_attn_run::<CudaRuntime>(&c, &q, &k, &v, 1, nh, nkv, sq, sk, sk, d, causal);
let want = sdpa_ref(&q, &k, &v, nh, nkv, sq, sk, d, causal);
let rel = scale_rel(&got, &want);
let nonfin = got.iter().filter(|x| !x.is_finite()).count();
let where_ = if rel < 2e-3 {
String::new()
} else {
got.iter()
.zip(&want)
.enumerate()
.find(|(_, (g, w))| (**g - **w).abs() > 1e-2 || !g.is_finite())
.map(|(i, (g, w))| {
let (h, rem) = (i / (sq * d), i % (sq * d));
format!(
" first-bad@[h{h},q{},d{}] got={g:.3e} want={w:.3e}",
rem / d,
rem % d
)
})
.unwrap_or_default()
};
eprintln!("[flash-cuda {tag}] nh{nh}/nkv{nkv} sq{sq} sk{sk} d{d} causal{causal} scale_rel={rel:.2e} nonfinite={nonfin}{where_}");
if !(rel < 2e-2) {
fails += 1;
}
}
assert_eq!(
fails, 0,
"flash coopmat CUDA: {fails} shape(s) exceeded 2e-2 vs materialized sdpa_ref"
);
for (nh, nkv, sq, sk, d, causal, tag) in [
(
4usize,
2usize,
16usize,
128usize,
32usize,
true,
"prefill causal",
),
(8, 2, 1, 512, 64, false, "decode kv512"),
(2, 1, 512, 512, 128, true, "prefill 512 d128"),
] {
let q = rnd(nh * sq * d, 0x1234_5678);
let k = rnd(nkv * sk * d, 0x9ABC_DEF0);
let v = rnd(nkv * sk * d, 0x0FED_CBA9);
let cubes = flash_attn_cubes(1, nh, sq);
let p32 = flash_attn_launch::<CudaRuntime>(
&c,
&q,
&k,
&v,
1,
nh,
nkv,
sq,
sk,
sk,
d,
causal,
32,
0,
cubes,
Target::of(&c),
);
let p64 = flash_attn_launch::<CudaRuntime>(
&c,
&q,
&k,
&v,
1,
nh,
nkv,
sq,
sk,
sk,
d,
causal,
64,
0,
cubes,
Target::of(&c),
);
let agree = scale_rel(&p32, &p64);
eprintln!("[flash-cuda plane-agree {tag}] scale_rel(p32,p64)={agree:.2e}");
assert!(
agree < 1e-4,
"CUDA plane 32 vs 64 diverged ({agree}) -- coopmat no longer warp-scoped"
);
}
}
}