use cudarc::driver::sys::CUdeviceptr;
use cudarc::driver::{LaunchConfig, PushKernelArg};
use onnx_runtime_ep_api::{EpError, Result};
use crate::error::driver_err;
use crate::runtime::CudaRuntime;
const MODULE_KEY: &str = "gqa_decode_attention_f32_v6";
const ENTRY_PREFIX: &str = "gqa_decode_attention_f32_dpl";
const MERGE_ENTRY_PREFIX: &str = "gqa_decode_attention_f32_merge_dpl";
pub(super) const MAX_HEAD_DIM: usize = 512;
const MAX_DIMS_PER_LANE: usize = 16;
fn dims_per_lane(head_dim: usize) -> usize {
head_dim
.div_ceil(WARP_SIZE as usize)
.clamp(1, MAX_DIMS_PER_LANE)
}
const WARPS_PER_BLOCK: u32 = 4;
const WARP_SIZE: u32 = 32;
pub(super) const MAX_SPLITS: usize = 16;
pub(super) fn supported(query_seq: usize, head_dim: usize) -> bool {
query_seq == 1 && (1..=MAX_HEAD_DIM).contains(&head_dim)
}
pub(super) fn single_split_direct_flag() -> i32 {
#[cfg(test)]
{
let override_value = TEST_SINGLE_SPLIT_OVERRIDE.load(std::sync::atomic::Ordering::Relaxed);
if override_value >= 0 {
return override_value;
}
}
use std::sync::OnceLock;
static FLAG: OnceLock<i32> = OnceLock::new();
*FLAG.get_or_init(
|| match std::env::var("ONNX_GENAI_CUDA_GQA_DIRECT_SINGLE_SPLIT") {
Ok(value) if value == "0" => 0,
_ => 1,
},
)
}
pub(super) fn split_fill_override(max_splits: usize) -> Option<usize> {
use std::sync::OnceLock;
static VALUE: OnceLock<Option<usize>> = OnceLock::new();
*VALUE.get_or_init(|| {
let value = std::env::var("ONNX_GENAI_CUDA_GQA_SPLITS").ok()?;
let parsed = value.parse::<usize>().ok()?;
Some(parsed.clamp(1, max_splits))
})
}
#[cfg(test)]
pub(super) static TEST_SINGLE_SPLIT_OVERRIDE: std::sync::atomic::AtomicI32 =
std::sync::atomic::AtomicI32::new(-1);
#[cfg(test)]
pub(super) static TEST_SINGLE_SPLIT_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
const DECODE_SRC: &str = r#"
#define GQA_WARP_SIZE 32
// Dims-per-lane (DPL) is a compile-time template parameter: each of the 32 warp
// lanes owns `DPL` output dims, so the kernel covers head_dim <= 32 * DPL. The
// launcher instantiates the exact tier `DPL = ceil(head_dim / 32)` (1..=8) so a
// small head (e.g. 64 -> DPL 2, 128 -> DPL 4) keeps its original register
// footprint while head_dim=256 uses DPL 8 and head_dim=512 uses DPL 16 -- no
// register regression on small heads, no head256/head512 fallback to the serial
// reference kernel.
#define GQA_MAX_HEAD_SIZE 512
#define GQA_MAX_SPLITS 16
#define GQA_MAX_SCRATCH_ROWS 256
#define GQA_SCRATCH_STRIDE (GQA_MAX_HEAD_SIZE + 2)
__device__ __align__(16) float gqa_split_scratch[
GQA_MAX_SCRATCH_ROWS * GQA_MAX_SPLITS * GQA_SCRATCH_STRIDE];
__device__ __forceinline__ int gqa_active_splits(const int sequence_length) {
if (sequence_length <= 64) return 1;
if (sequence_length <= 128) return 2;
if (sequence_length <= 256) return 4;
if (sequence_length <= 512) return 8;
return GQA_MAX_SPLITS;
}
template<int DPL>
__device__ __forceinline__ void gqa_decode_f32_impl(
const float* __restrict__ query,
const float* __restrict__ key,
const float* __restrict__ value,
float* __restrict__ output,
const int* __restrict__ total_lengths,
const int batch,
const int query_heads,
const int kv_heads,
const int query_seq,
const int head_size,
const int cache_capacity,
const int group_size,
const float scale,
const int local_window,
const float softcap,
const int single_split_direct,
const float* __restrict__ head_sink)
{
extern __shared__ float smem[];
const int warps_per_block = blockDim.x / GQA_WARP_SIZE;
float* warp_max = smem;
float* warp_sum = warp_max + warps_per_block;
float* warp_acc = warp_sum + warps_per_block;
const int lane = threadIdx.x % GQA_WARP_SIZE;
const int warp_in_block = threadIdx.x / GQA_WARP_SIZE;
const int row = blockIdx.x / GQA_MAX_SPLITS;
const int split = blockIdx.x % GQA_MAX_SPLITS;
const int rows = batch * query_heads * query_seq;
if (row >= rows) return;
const int query_pos = row % query_seq;
const int query_head = (row / query_seq) % query_heads;
const int batch_index = row / (query_heads * query_seq);
const int kv_head = query_head / group_size;
const int total = total_lengths[batch_index];
const int causal_limit = total - query_seq + query_pos;
const int local_start =
(local_window > 0 && causal_limit + 1 > local_window)
? causal_limit + 1 - local_window
: 0;
const int sequence_length = max(0, causal_limit + 1 - local_start);
const int active_splits =
(row < GQA_MAX_SCRATCH_ROWS) ? gqa_active_splits(sequence_length) : 1;
if (split >= active_splits) return;
const int keys_per_split =
(sequence_length + active_splits - 1) / active_splits;
const int split_start = local_start + split * keys_per_split;
const int split_end = min(causal_limit + 1, split_start + keys_per_split);
const long q_base =
((long)(batch_index * query_heads + query_head) * query_seq + query_pos)
* (long)head_size;
const long kv_plane =
(long)(batch_index * kv_heads + kv_head) * (long)cache_capacity * (long)head_size;
float q_reg[DPL];
float acc[DPL];
#pragma unroll
for (int i = 0; i < DPL; ++i) {
const int d = lane + i * GQA_WARP_SIZE;
q_reg[i] = (d < head_size) ? query[q_base + d] : 0.0f;
acc[i] = 0.0f;
}
const float negative_infinity = __int_as_float(0xff800000);
float running_max = negative_infinity;
float running_sum = 0.0f;
for (int key_pos = split_start + warp_in_block; key_pos < split_end;
key_pos += warps_per_block) {
const long k_base = kv_plane + (long)key_pos * (long)head_size;
float partial = 0.0f;
#pragma unroll
for (int i = 0; i < DPL; ++i) {
const int d = lane + i * GQA_WARP_SIZE;
if (d < head_size) {
partial += q_reg[i] * key[k_base + d];
}
}
// Butterfly all-reduce: every lane ends with the full QK dot product.
#pragma unroll
for (int offset = GQA_WARP_SIZE / 2; offset > 0; offset >>= 1) {
partial += __shfl_xor_sync(0xffffffffu, partial, offset);
}
float score = partial * scale;
if (softcap != 0.0f) {
score = softcap * tanhf(score / softcap);
}
const float new_max = fmaxf(running_max, score);
const float correction = expf(running_max - new_max);
const float probability = expf(score - new_max);
running_sum = running_sum * correction + probability;
#pragma unroll
for (int i = 0; i < DPL; ++i) {
const int d = lane + i * GQA_WARP_SIZE;
const float v = (d < head_size) ? value[k_base + d] : 0.0f;
acc[i] = acc[i] * correction + probability * v;
}
running_max = new_max;
}
if (lane == 0) {
warp_max[warp_in_block] = running_max;
warp_sum[warp_in_block] = running_sum;
}
#pragma unroll
for (int i = 0; i < DPL; ++i) {
const int d = lane + i * GQA_WARP_SIZE;
if (d < head_size) {
warp_acc[warp_in_block * head_size + d] = acc[i];
}
}
__syncthreads();
if (warp_in_block != 0) return;
float global_max = negative_infinity;
for (int w = 0; w < warps_per_block; ++w) {
global_max = fmaxf(global_max, warp_max[w]);
}
// When `single_split_direct` is enabled and the on-device length selects a
// single split, the sole active CTA already owns the complete flash state,
// so it normalizes and writes the final output directly, skipping the
// scratch round-trip and the merge pass. Bit-identical to the two-step path
// (the merge with one split multiplies by exp(0)==1 and the same 1/denom).
const bool direct_output =
(row >= GQA_MAX_SCRATCH_ROWS) || (single_split_direct != 0 && active_splits == 1);
// Learned attention sink (gpt-oss family): a per-head logit that joins the
// softmax denominator once per row but contributes no value vector. It is
// folded in only on the FINAL normalization (direct-output here, or the
// merge kernel for multi-split rows), never per-split, so split states stay
// sink-free. When absent, `sink_val == -inf` makes the fold a no-op
// (byte-identical to the pre-sink path).
const float sink_val =
(head_sink != nullptr) ? head_sink[query_head] : negative_infinity;
if (direct_output) {
global_max = fmaxf(global_max, sink_val);
}
float denom = 0.0f;
for (int w = 0; w < warps_per_block; ++w) {
denom += warp_sum[w] * expf(warp_max[w] - global_max);
}
if (direct_output) {
denom += expf(sink_val - global_max);
}
float* split_state = direct_output
? nullptr
: gqa_split_scratch
+ (row * GQA_MAX_SPLITS + split) * GQA_SCRATCH_STRIDE;
if (lane == 0 && !direct_output) {
split_state[0] = global_max;
split_state[1] = denom;
}
const float inverse_sum =
(direct_output && denom > 0.0f) ? (1.0f / denom) : 0.0f;
#pragma unroll
for (int i = 0; i < DPL; ++i) {
const int d = lane + i * GQA_WARP_SIZE;
if (d < head_size) {
float out = 0.0f;
for (int w = 0; w < warps_per_block; ++w) {
out += warp_acc[w * head_size + d]
* expf(warp_max[w] - global_max);
}
if (direct_output) {
output[q_base + d] = out * inverse_sum;
} else {
split_state[2 + d] = out;
}
}
}
}
template<int DPL>
__device__ __forceinline__ void gqa_merge_f32_impl(
float* __restrict__ output,
const int* __restrict__ total_lengths,
const int batch,
const int query_heads,
const int query_seq,
const int head_size,
const int local_window,
const int single_split_direct,
const float* __restrict__ head_sink)
{
const int row = blockIdx.x;
const int rows = batch * query_heads * query_seq;
if (row >= rows || row >= GQA_MAX_SCRATCH_ROWS) return;
const int lane = threadIdx.x;
const int query_pos = row % query_seq;
const int batch_index = row / (query_heads * query_seq);
const int total = total_lengths[batch_index];
const int causal_limit = total - query_seq + query_pos;
const int local_start =
(local_window > 0 && causal_limit + 1 > local_window)
? causal_limit + 1 - local_window
: 0;
const int sequence_length = max(0, causal_limit + 1 - local_start);
const int active_splits = gqa_active_splits(sequence_length);
// Single-split rows were finalized in-place by the decode kernel.
if (single_split_direct != 0 && active_splits <= 1) return;
const int query_head = (row / query_seq) % query_heads;
float global_max = __int_as_float(0xff800000);
for (int split = 0; split < active_splits; ++split) {
const float* state = gqa_split_scratch
+ (row * GQA_MAX_SPLITS + split) * GQA_SCRATCH_STRIDE;
global_max = fmaxf(global_max, state[0]);
}
// Learned attention sink (see the decode kernel): folded into the softmax
// denominator once, on this final merge, for multi-split rows. `-inf` when
// absent => byte-identical no-op.
const float sink_val =
(head_sink != nullptr) ? head_sink[query_head] : __int_as_float(0xff800000);
global_max = fmaxf(global_max, sink_val);
float denom = 0.0f;
for (int split = 0; split < active_splits; ++split) {
const float* state = gqa_split_scratch
+ (row * GQA_MAX_SPLITS + split) * GQA_SCRATCH_STRIDE;
denom += state[1] * expf(state[0] - global_max);
}
denom += expf(sink_val - global_max);
const float inverse_sum = (denom > 0.0f) ? (1.0f / denom) : 0.0f;
const long q_base = (long)row * (long)head_size;
#pragma unroll
for (int i = 0; i < DPL; ++i) {
const int d = lane + i * GQA_WARP_SIZE;
if (d < head_size) {
float out = 0.0f;
for (int split = 0; split < active_splits; ++split) {
const float* state = gqa_split_scratch
+ (row * GQA_MAX_SPLITS + split) * GQA_SCRATCH_STRIDE;
out += state[2 + d] * expf(state[0] - global_max);
}
output[q_base + d] = out * inverse_sum;
}
}
}
// Emit the `extern "C"` launch entry points for one DPL tier. The launcher
// selects `gqa_decode_attention_f32_dpl{N}` / `..._merge_dpl{N}` for
// `N = ceil(head_dim / 32)`.
#define GQA_DECODE_ENTRIES(DPL) \
extern "C" __global__ void gqa_decode_attention_f32_dpl##DPL( \
const float* __restrict__ query, const float* __restrict__ key, \
const float* __restrict__ value, float* __restrict__ output, \
const int* __restrict__ total_lengths, const int batch, \
const int query_heads, const int kv_heads, const int query_seq, \
const int head_size, const int cache_capacity, const int group_size, \
const float scale, const int local_window, const float softcap, \
const int single_split_direct, \
const float* __restrict__ head_sink) { \
gqa_decode_f32_impl<DPL>(query, key, value, output, total_lengths, batch, \
query_heads, kv_heads, query_seq, head_size, cache_capacity, \
group_size, scale, local_window, softcap, single_split_direct, \
head_sink); \
} \
extern "C" __global__ void gqa_decode_attention_f32_merge_dpl##DPL( \
float* __restrict__ output, const int* __restrict__ total_lengths, \
const int batch, const int query_heads, const int query_seq, \
const int head_size, const int local_window, \
const int single_split_direct, \
const float* __restrict__ head_sink) { \
gqa_merge_f32_impl<DPL>(output, total_lengths, batch, query_heads, \
query_seq, head_size, local_window, single_split_direct, \
head_sink); \
}
GQA_DECODE_ENTRIES(1)
GQA_DECODE_ENTRIES(2)
GQA_DECODE_ENTRIES(3)
GQA_DECODE_ENTRIES(4)
GQA_DECODE_ENTRIES(5)
GQA_DECODE_ENTRIES(6)
GQA_DECODE_ENTRIES(7)
GQA_DECODE_ENTRIES(8)
GQA_DECODE_ENTRIES(9)
GQA_DECODE_ENTRIES(10)
GQA_DECODE_ENTRIES(11)
GQA_DECODE_ENTRIES(12)
GQA_DECODE_ENTRIES(13)
GQA_DECODE_ENTRIES(14)
GQA_DECODE_ENTRIES(15)
GQA_DECODE_ENTRIES(16)
"#;
#[allow(clippy::too_many_arguments)]
pub(super) fn run(
runtime: &CudaRuntime,
batch: usize,
num_heads: usize,
num_kv_heads: usize,
query_seq: usize,
head_dim: usize,
cache_capacity: usize,
group: usize,
scale: f32,
query: CUdeviceptr,
key: CUdeviceptr,
value: CUdeviceptr,
output: CUdeviceptr,
total_lengths: CUdeviceptr,
local_window: i32,
softcap: f32,
head_sink: CUdeviceptr,
) -> Result<()> {
let as_i32 = |name: &str, value: usize| {
i32::try_from(value).map_err(|_| {
EpError::KernelFailed(format!("cuda_ep GQA decode: {name} {value} exceeds i32"))
})
};
let batch_i = as_i32("batch", batch)?;
let heads_i = as_i32("num_heads", num_heads)?;
let kv_heads_i = as_i32("num_kv_heads", num_kv_heads)?;
let query_seq_i = as_i32("query_seq", query_seq)?;
let dim_i = as_i32("head_dim", head_dim)?;
let capacity_i = as_i32("cache_capacity", cache_capacity)?;
let group_i = as_i32("GQA group", group)?;
let rows = batch
.checked_mul(num_heads)
.and_then(|value| value.checked_mul(query_seq))
.ok_or_else(|| EpError::KernelFailed("cuda_ep GQA decode: row count overflow".into()))?;
let partial_blocks = rows
.checked_mul(MAX_SPLITS)
.ok_or_else(|| EpError::KernelFailed("cuda_ep GQA decode: split grid overflow".into()))?;
let grid_x = u32::try_from(partial_blocks.max(1)).map_err(|_| {
EpError::KernelFailed(format!(
"cuda_ep GQA decode: {partial_blocks} split blocks exceed CUDA grid.x"
))
})?;
let merge_grid_x = u32::try_from(rows.max(1)).map_err(|_| {
EpError::KernelFailed(format!(
"cuda_ep GQA decode: {rows} rows exceed CUDA grid.x"
))
})?;
let warps = WARPS_PER_BLOCK as usize;
let shared_floats = warps
.checked_mul(2)
.and_then(|base| warps.checked_mul(head_dim).map(|acc| base + acc))
.ok_or_else(|| EpError::KernelFailed("cuda_ep GQA decode: shared-mem overflow".into()))?;
let shared_mem_bytes =
u32::try_from(shared_floats * std::mem::size_of::<f32>()).map_err(|_| {
EpError::KernelFailed("cuda_ep GQA decode: shared-mem bytes exceed u32".into())
})?;
let dpl = dims_per_lane(head_dim);
let decode_entry = format!("{ENTRY_PREFIX}{dpl}");
let merge_entry = format!("{MERGE_ENTRY_PREFIX}{dpl}");
let function = runtime.nvrtc_function(MODULE_KEY, DECODE_SRC, &decode_entry)?;
let single_split_direct = single_split_direct_flag();
let mut builder = runtime.stream().launch_builder(&function);
builder
.arg(&query)
.arg(&key)
.arg(&value)
.arg(&output)
.arg(&total_lengths)
.arg(&batch_i)
.arg(&heads_i)
.arg(&kv_heads_i)
.arg(&query_seq_i)
.arg(&dim_i)
.arg(&capacity_i)
.arg(&group_i)
.arg(&scale)
.arg(&local_window)
.arg(&softcap)
.arg(&single_split_direct)
.arg(&head_sink);
unsafe {
builder.launch(LaunchConfig {
grid_dim: (grid_x, 1, 1),
block_dim: (WARPS_PER_BLOCK * WARP_SIZE, 1, 1),
shared_mem_bytes,
})
}
.map_err(|error| driver_err("launch GQA f32 split-K attention", error))?;
let merge_function = runtime.nvrtc_function(MODULE_KEY, DECODE_SRC, &merge_entry)?;
let mut merge_builder = runtime.stream().launch_builder(&merge_function);
merge_builder
.arg(&output)
.arg(&total_lengths)
.arg(&batch_i)
.arg(&heads_i)
.arg(&query_seq_i)
.arg(&dim_i)
.arg(&local_window)
.arg(&single_split_direct)
.arg(&head_sink);
unsafe {
merge_builder.launch(LaunchConfig {
grid_dim: (merge_grid_x, 1, 1),
block_dim: (WARP_SIZE, 1, 1),
shared_mem_bytes: 0,
})
}
.map_err(|error| driver_err("launch GQA f32 split-K merge", error))?;
Ok(())
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::*;
fn runtime() -> Option<Arc<CudaRuntime>> {
crate::test_support::maybe_runtime()
}
fn as_bytes<T: Copy>(values: &[T]) -> &[u8] {
unsafe {
std::slice::from_raw_parts(values.as_ptr().cast::<u8>(), std::mem::size_of_val(values))
}
}
fn as_bytes_mut<T: Copy>(values: &mut [T]) -> &mut [u8] {
unsafe {
std::slice::from_raw_parts_mut(
values.as_mut_ptr().cast::<u8>(),
std::mem::size_of_val(values),
)
}
}
#[allow(clippy::too_many_arguments)]
fn cpu_reference(
query: &[f32],
key: &[f32],
value: &[f32],
total: usize,
num_heads: usize,
num_kv_heads: usize,
head_dim: usize,
cache_capacity: usize,
scale: f32,
) -> Vec<f32> {
let group = num_heads / num_kv_heads;
let mut output = vec![0.0f32; num_heads * head_dim];
for h in 0..num_heads {
let kv_head = h / group;
let q_base = h * head_dim;
let mut scores = vec![0.0f64; total];
let mut maximum = f64::NEG_INFINITY;
for (key_pos, score_slot) in scores.iter_mut().enumerate() {
let k_base = (kv_head * cache_capacity + key_pos) * head_dim;
let mut dot = 0.0f64;
for d in 0..head_dim {
dot += query[q_base + d] as f64 * key[k_base + d] as f64;
}
let score = dot * scale as f64;
*score_slot = score;
maximum = maximum.max(score);
}
let mut denom = 0.0f64;
for score in scores.iter_mut() {
*score = (*score - maximum).exp();
denom += *score;
}
for d in 0..head_dim {
let mut acc = 0.0f64;
for (key_pos, prob) in scores.iter().enumerate() {
let v_index = (kv_head * cache_capacity + key_pos) * head_dim + d;
acc += prob / denom * value[v_index] as f64;
}
output[q_base + d] = acc as f32;
}
}
output
}
#[test]
fn decode_kernel_matches_reference_softmax() {
let Some(runtime) = runtime() else {
eprintln!("skipping CUDA GQA decode parity test: CUDA runtime unavailable");
return;
};
let batch = 1usize;
let num_heads = 14usize;
let num_kv_heads = 2usize;
let head_dim = 64usize;
let cache_capacity = 1024usize;
let group = num_heads / num_kv_heads;
let scale = 1.0f32 / (head_dim as f32).sqrt();
let mut state = 0x1234_5678u64;
let mut next = || {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((state >> 33) as f32 / u32::MAX as f32) * 2.0 - 1.0
};
let query: Vec<f32> = (0..num_heads * head_dim).map(|_| next()).collect();
let key: Vec<f32> = (0..num_kv_heads * cache_capacity * head_dim)
.map(|_| next())
.collect();
let value: Vec<f32> = (0..num_kv_heads * cache_capacity * head_dim)
.map(|_| next())
.collect();
let query_dev = runtime.alloc_raw(query.len() * 4).unwrap();
let key_dev = runtime.alloc_raw(key.len() * 4).unwrap();
let value_dev = runtime.alloc_raw(value.len() * 4).unwrap();
let output_dev = runtime.alloc_raw(num_heads * head_dim * 4).unwrap();
let totals_dev = runtime.alloc_raw(batch * 4).unwrap();
unsafe {
runtime.htod(as_bytes(&query), query_dev).unwrap();
runtime.htod(as_bytes(&key), key_dev).unwrap();
runtime.htod(as_bytes(&value), value_dev).unwrap();
}
let mut worst_abs = 0.0f32;
let mut worst_rel = 0.0f32;
for total in [1usize, 7, 64, 255, 1023] {
let totals = [total as i32];
unsafe {
runtime.htod(as_bytes(&totals), totals_dev).unwrap();
}
run(
&runtime,
batch,
num_heads,
num_kv_heads,
1,
head_dim,
cache_capacity,
group,
scale,
query_dev,
key_dev,
value_dev,
output_dev,
totals_dev,
0,
0.0,
0, )
.unwrap();
let mut got = vec![0.0f32; num_heads * head_dim];
unsafe {
runtime.dtoh(as_bytes_mut(&mut got), output_dev).unwrap();
}
let expected = cpu_reference(
&query,
&key,
&value,
total,
num_heads,
num_kv_heads,
head_dim,
cache_capacity,
scale,
);
for (g, e) in got.iter().zip(expected.iter()) {
let abs = (g - e).abs();
let rel = abs / e.abs().max(1e-4);
worst_abs = worst_abs.max(abs);
worst_rel = worst_rel.max(rel);
}
}
unsafe {
runtime.free_raw(query_dev).unwrap();
runtime.free_raw(key_dev).unwrap();
runtime.free_raw(value_dev).unwrap();
runtime.free_raw(output_dev).unwrap();
runtime.free_raw(totals_dev).unwrap();
}
eprintln!("GQA decode parity: max_abs={worst_abs:.3e} max_rel={worst_rel:.3e}");
assert!(
worst_abs < 1e-3,
"decode kernel diverged from reference softmax: max_abs={worst_abs:.3e}"
);
assert!(
worst_rel < 5e-3,
"decode kernel diverged from reference softmax: max_rel={worst_rel:.3e}"
);
}
fn parity_for_shape(
runtime: &CudaRuntime,
num_heads: usize,
num_kv_heads: usize,
head_dim: usize,
totals: &[usize],
) -> (f32, f32) {
let batch = 1usize;
let cache_capacity = 1024usize;
let group = num_heads / num_kv_heads;
let scale = 1.0f32 / (head_dim as f32).sqrt();
let mut state = 0x0BADC0DEu64 ^ ((head_dim as u64) << 17);
let mut next = || {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((state >> 33) as f32 / u32::MAX as f32) * 2.0 - 1.0
};
let query: Vec<f32> = (0..num_heads * head_dim).map(|_| next()).collect();
let key: Vec<f32> = (0..num_kv_heads * cache_capacity * head_dim)
.map(|_| next())
.collect();
let value: Vec<f32> = (0..num_kv_heads * cache_capacity * head_dim)
.map(|_| next())
.collect();
let query_dev = runtime.alloc_raw(query.len() * 4).unwrap();
let key_dev = runtime.alloc_raw(key.len() * 4).unwrap();
let value_dev = runtime.alloc_raw(value.len() * 4).unwrap();
let output_dev = runtime.alloc_raw(num_heads * head_dim * 4).unwrap();
let totals_dev = runtime.alloc_raw(batch * 4).unwrap();
unsafe {
runtime.htod(as_bytes(&query), query_dev).unwrap();
runtime.htod(as_bytes(&key), key_dev).unwrap();
runtime.htod(as_bytes(&value), value_dev).unwrap();
}
let mut worst_abs = 0.0f32;
let mut worst_rel = 0.0f32;
for &total in totals {
let totals_host = [total as i32];
unsafe {
runtime.htod(as_bytes(&totals_host), totals_dev).unwrap();
}
run(
runtime,
batch,
num_heads,
num_kv_heads,
1,
head_dim,
cache_capacity,
group,
scale,
query_dev,
key_dev,
value_dev,
output_dev,
totals_dev,
0,
0.0,
0, )
.unwrap();
let mut got = vec![0.0f32; num_heads * head_dim];
unsafe {
runtime.dtoh(as_bytes_mut(&mut got), output_dev).unwrap();
}
let expected = cpu_reference(
&query,
&key,
&value,
total,
num_heads,
num_kv_heads,
head_dim,
cache_capacity,
scale,
);
for (g, e) in got.iter().zip(expected.iter()) {
let abs = (g - e).abs();
let rel = abs / e.abs().max(1e-4);
worst_abs = worst_abs.max(abs);
worst_rel = worst_rel.max(rel);
}
}
unsafe {
runtime.free_raw(query_dev).unwrap();
runtime.free_raw(key_dev).unwrap();
runtime.free_raw(value_dev).unwrap();
runtime.free_raw(output_dev).unwrap();
runtime.free_raw(totals_dev).unwrap();
}
(worst_abs, worst_rel)
}
#[test]
fn decode_kernel_matches_reference_softmax_head256() {
let Some(runtime) = runtime() else {
eprintln!("skipping CUDA GQA decode head256 parity test: CUDA runtime unavailable");
return;
};
let totals = [1usize, 7, 64, 65, 128, 129, 255, 256, 257, 512, 1023];
let (worst_abs, worst_rel) = parity_for_shape(&runtime, 8, 2, 256, &totals);
eprintln!("GQA decode head256 parity: max_abs={worst_abs:.3e} max_rel={worst_rel:.3e}");
assert!(
worst_abs < 1e-3,
"head256 decode kernel diverged from reference softmax: max_abs={worst_abs:.3e}"
);
assert!(
worst_rel < 5e-3,
"head256 decode kernel diverged from reference softmax: max_rel={worst_rel:.3e}"
);
}
#[test]
fn decode_kernel_matches_reference_softmax_general_head_dims() {
let Some(runtime) = runtime() else {
eprintln!(
"skipping CUDA GQA decode general head-dim parity test: CUDA runtime unavailable"
);
return;
};
let totals = [1usize, 7, 64, 65, 128, 129, 255, 256, 300];
for head_dim in [64usize, 80, 96, 112, 128, 192, 256, 320, 384, 448, 512] {
let (worst_abs, worst_rel) = parity_for_shape(&runtime, 8, 2, head_dim, &totals);
eprintln!(
"GQA decode head{head_dim} parity (dpl={}): max_abs={worst_abs:.3e} max_rel={worst_rel:.3e}",
super::dims_per_lane(head_dim)
);
assert!(
worst_abs < 1e-3,
"head{head_dim} decode kernel diverged from reference softmax: max_abs={worst_abs:.3e}"
);
assert!(
worst_rel < 5e-3,
"head{head_dim} decode kernel diverged from reference softmax: max_rel={worst_rel:.3e}"
);
}
}
#[test]
fn support_gate_targets_single_token_decode() {
assert!(supported(1, 64));
assert!(supported(1, 128));
assert!(supported(1, 129));
assert!(supported(1, 256));
assert!(supported(1, 257));
assert!(supported(1, 512));
assert!(!supported(1, 513));
assert!(!supported(2, 64));
assert!(!supported(1, 0));
}
#[test]
fn single_split_direct_output_is_bit_exact_to_two_step_path_f32() {
use std::sync::atomic::Ordering;
let _serial = TEST_SINGLE_SPLIT_LOCK
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let Some(runtime) = runtime() else {
eprintln!("skipping GQA f32 single-split parity test: CUDA runtime unavailable");
return;
};
let batch = 1usize;
let num_heads = 8usize;
let num_kv_heads = 2usize;
let cache_capacity = 1024usize;
let group = num_heads / num_kv_heads;
let mut state = 0x51A9_7C3Du64;
let mut next = || {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((state >> 33) as f32 / u32::MAX as f32) * 2.0 - 1.0
};
for head_dim in [64usize, 128usize, 256usize, 512usize] {
let scale = 1.0f32 / (head_dim as f32).sqrt();
let mut q = vec![0.0f32; num_heads * head_dim];
for slot in q.iter_mut() {
*slot = next();
}
let kv_len = num_kv_heads * cache_capacity * head_dim;
let mut k = vec![0.0f32; kv_len];
let mut v = vec![0.0f32; kv_len];
for i in 0..kv_len {
k[i] = next();
v[i] = next();
}
let query_dev = runtime.alloc_raw(q.len() * 4).unwrap();
let key_dev = runtime.alloc_raw(k.len() * 4).unwrap();
let value_dev = runtime.alloc_raw(v.len() * 4).unwrap();
let output_dev = runtime.alloc_raw(num_heads * head_dim * 4).unwrap();
let totals_dev = runtime.alloc_raw(batch * 4).unwrap();
unsafe {
runtime.htod(as_bytes(&q), query_dev).unwrap();
runtime.htod(as_bytes(&k), key_dev).unwrap();
runtime.htod(as_bytes(&v), value_dev).unwrap();
}
let launch = |flag: i32, total: usize| -> Vec<f32> {
TEST_SINGLE_SPLIT_OVERRIDE.store(flag, Ordering::Relaxed);
let totals = [total as i32];
unsafe {
runtime.htod(as_bytes(&totals), totals_dev).unwrap();
}
run(
&runtime,
batch,
num_heads,
num_kv_heads,
1,
head_dim,
cache_capacity,
group,
scale,
query_dev,
key_dev,
value_dev,
output_dev,
totals_dev,
0,
0.0,
0, )
.unwrap();
let mut got = vec![0.0f32; num_heads * head_dim];
unsafe {
runtime.dtoh(as_bytes_mut(&mut got), output_dev).unwrap();
}
got
};
for total in [1usize, 8, 33, 64, 65, 129, 256, 257, 300] {
let direct = launch(1, total);
let two_step = launch(0, total);
assert_eq!(
direct.iter().map(|x| x.to_bits()).collect::<Vec<_>>(),
two_step.iter().map(|x| x.to_bits()).collect::<Vec<_>>(),
"single-split fast path diverged from two-step path at head_dim={head_dim} total={total}"
);
}
TEST_SINGLE_SPLIT_OVERRIDE.store(-1, Ordering::Relaxed);
unsafe {
runtime.free_raw(query_dev).unwrap();
runtime.free_raw(key_dev).unwrap();
runtime.free_raw(value_dev).unwrap();
runtime.free_raw(output_dev).unwrap();
runtime.free_raw(totals_dev).unwrap();
}
}
}
}