use metal::MTLSize;
use crate::buffer::MlxBuffer;
use crate::device::MlxDevice;
use crate::encoder::{as_bytes, CapturedOpKind, CommandEncoder, KernelArg};
use crate::error::{MlxError, Result};
use crate::kernel_registry::KernelRegistry;
pub static FLASH_ATTN_VEC_TQ_HB_SHADER_SOURCE: &str =
include_str!("../shaders/flash_attn_vec_tq_hb.metal");
pub fn register(registry: &mut KernelRegistry) {
registry.register_source("flash_attn_vec_tq_hb_dk256", FLASH_ATTN_VEC_TQ_HB_SHADER_SOURCE);
registry.register_source("flash_attn_vec_tq_hb_dk512", FLASH_ATTN_VEC_TQ_HB_SHADER_SOURCE);
registry.register_source("flash_attn_vec_tq_hb_batched_dk256", FLASH_ATTN_VEC_TQ_HB_SHADER_SOURCE);
registry.register_source("flash_attn_vec_tq_hb_batched_dk512", FLASH_ATTN_VEC_TQ_HB_SHADER_SOURCE);
}
#[derive(Debug, Clone, Copy)]
pub struct FlashAttnVecTqHbParams {
pub num_heads: u32,
pub num_kv_heads: u32,
pub head_dim: u32,
pub kv_seq_len: u32,
pub kv_capacity: u32,
pub scale: f32,
pub mask_type: u32,
pub sliding_window: u32,
pub softcap: f32,
pub ring_start: u32,
pub scale_factor_d512: f32,
pub codebook_bits: u32,
pub fuse_fwht_pre: u32,
pub nsg: u32,
}
#[repr(C)]
#[derive(Debug, Clone, Copy, bytemuck::Pod, bytemuck::Zeroable)]
struct FlashAttnVecTqHbParamsGpu {
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
kv_seq_len: u32,
kv_capacity: u32,
scale: f32,
mask_type: u32,
sliding_window: u32,
softcap: f32,
nwg: u32,
ring_start: u32,
scale_factor_d512: f32,
codebook_bits: u32,
fuse_fwht_pre: u32,
nsg: u32,
}
#[repr(C)]
#[derive(Debug, Clone, Copy, bytemuck::Pod, bytemuck::Zeroable)]
struct FlashAttnVecReduceParamsGpu {
nrows: u32,
}
fn validate_params(params: &FlashAttnVecTqHbParams) -> Result<()> {
if params.head_dim != 256 && params.head_dim != 512 {
return Err(MlxError::InvalidArgument(format!(
"flash_attn_vec_tq_hb: head_dim must be 256 or 512, got {}",
params.head_dim
)));
}
if params.num_heads == 0 || params.num_kv_heads == 0 {
return Err(MlxError::InvalidArgument(
"flash_attn_vec_tq_hb: num_heads and num_kv_heads must be > 0".into(),
));
}
if params.num_heads % params.num_kv_heads != 0 {
return Err(MlxError::InvalidArgument(format!(
"flash_attn_vec_tq_hb: num_heads ({}) % num_kv_heads ({}) != 0",
params.num_heads, params.num_kv_heads
)));
}
if params.kv_seq_len == 0 {
return Err(MlxError::InvalidArgument(
"flash_attn_vec_tq_hb: kv_seq_len must be > 0".into(),
));
}
if params.kv_capacity < params.kv_seq_len {
return Err(MlxError::InvalidArgument(format!(
"flash_attn_vec_tq_hb: kv_capacity ({}) < kv_seq_len ({})",
params.kv_capacity, params.kv_seq_len
)));
}
if !matches!(params.codebook_bits, 5 | 6 | 8) {
return Err(MlxError::InvalidArgument(format!(
"flash_attn_vec_tq_hb: codebook_bits must be 5, 6, or 8, got {}",
params.codebook_bits
)));
}
if params.nsg == 0 || (params.nsg & (params.nsg - 1)) != 0 {
return Err(MlxError::InvalidArgument(format!(
"flash_attn_vec_tq_hb: nsg must be a power of 2 (1, 2, 4, ...), got {}",
params.nsg
)));
}
if params.nsg > 4 {
return Err(MlxError::InvalidArgument(format!(
"flash_attn_vec_tq_hb: nsg must be ≤ 4 (kernel reduce cap), got {}",
params.nsg
)));
}
Ok(())
}
#[doc(hidden)]
pub(crate) static CACHED_TQ_NSG: std::sync::atomic::AtomicI32 = std::sync::atomic::AtomicI32::new(-1);
pub fn compute_nsg(kv_seq_len: u32) -> u32 {
use std::sync::atomic::Ordering;
let mut v = CACHED_TQ_NSG.load(Ordering::Relaxed);
if v < 0 {
let parsed = std::env::var("HF2Q_TQ_NSG")
.ok()
.and_then(|s| s.parse::<u32>().ok())
.filter(|&n| n == 1 || n == 2 || n == 4)
.unwrap_or(0);
CACHED_TQ_NSG.store(parsed as i32, Ordering::Relaxed);
v = parsed as i32;
}
if v > 0 {
return v as u32;
}
if kv_seq_len > 1024 { 4 } else { 1 }
}
fn compute_nwg(kv_seq_len: u32) -> u32 {
use std::sync::atomic::{AtomicI32, Ordering};
static CACHED_TQ_NWG: AtomicI32 = AtomicI32::new(-1);
let mut v = CACHED_TQ_NWG.load(Ordering::Relaxed);
if v < 0 {
let parsed = std::env::var("HF2Q_TQ_NWG")
.ok()
.and_then(|s| s.parse::<u32>().ok())
.filter(|&n| n >= 1 && n <= 32)
.unwrap_or(0);
CACHED_TQ_NWG.store(parsed as i32, Ordering::Relaxed);
v = parsed as i32;
}
if v > 0 {
return v as u32;
}
if kv_seq_len > 512 { 32 } else { 16 }
}
#[allow(clippy::too_many_arguments)]
pub fn flash_attn_vec_tq_hb(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
q: &MlxBuffer,
k_packed: &MlxBuffer,
k_norms: &MlxBuffer,
v_packed: &MlxBuffer,
v_norms: &MlxBuffer,
output: &MlxBuffer,
tmp: &MlxBuffer,
params: &FlashAttnVecTqHbParams,
) -> Result<()> {
validate_params(params)?;
let head_dim = params.head_dim;
let nwg = compute_nwg(params.kv_seq_len);
let gpu_params = FlashAttnVecTqHbParamsGpu {
n_heads: params.num_heads,
n_kv_heads: params.num_kv_heads,
head_dim: params.head_dim,
kv_seq_len: params.kv_seq_len,
kv_capacity: params.kv_capacity,
scale: params.scale,
mask_type: params.mask_type,
sliding_window: params.sliding_window,
softcap: params.softcap,
nwg,
ring_start: params.ring_start,
scale_factor_d512: params.scale_factor_d512,
codebook_bits: params.codebook_bits,
fuse_fwht_pre: params.fuse_fwht_pre,
nsg: params.nsg,
};
let kernel_name = match head_dim {
256 => "flash_attn_vec_tq_hb_dk256",
512 => "flash_attn_vec_tq_hb_dk512",
_ => return Err(MlxError::InvalidArgument(format!(
"flash_attn_vec_tq_hb: unsupported head_dim {head_dim}"
))),
};
let cbits_const = (params.codebook_bits as i32, 50usize);
let pipeline = registry
.get_pipeline_with_constants(
kernel_name,
device.metal_device(),
&[],
&[(cbits_const.1, cbits_const.0)],
)?;
let pk = pad2(head_dim as usize, 128);
let pv = pad2(head_dim as usize, 128);
let sh = 4 * 32;
let nsg = params.nsg as usize;
let shmem_halfs = pk + nsg * (sh + 2 * pv);
let shmem_bytes = shmem_halfs * 2;
encoder.set_op_kind(CapturedOpKind::Sdpa);
let threadgroups = MTLSize::new(1, params.num_heads as u64, nwg as u64);
let threadgroup_size = MTLSize::new(32, params.nsg as u64, 1);
let dst_buf = if nwg == 1 { output } else { tmp };
encoder.encode_threadgroups_with_args_and_shared(
pipeline,
&[
(0, KernelArg::Bytes(as_bytes(&gpu_params))),
(1, KernelArg::Buffer(q)),
(2, KernelArg::Buffer(k_packed)),
(3, KernelArg::Buffer(k_norms)),
(4, KernelArg::Buffer(v_packed)),
(5, KernelArg::Buffer(v_norms)),
(6, KernelArg::Buffer(dst_buf)),
],
&[(0, shmem_bytes as u64)],
threadgroups,
threadgroup_size,
);
if nwg > 1 {
encoder.memory_barrier();
let reduce_params = FlashAttnVecReduceParamsGpu { nrows: params.num_heads };
let reduce_kernel = match head_dim {
256 => "flash_attn_vec_reduce_dk256",
512 => "flash_attn_vec_reduce_dk512",
_ => unreachable!(),
};
let reduce_pipeline = registry.get_pipeline(reduce_kernel, device.metal_device())?;
let reduce_tg = MTLSize::new(params.num_heads as u64, 1, 1);
let reduce_tg_size = MTLSize::new(32 * nwg as u64, 1, 1);
encoder.encode_threadgroups_with_args(
reduce_pipeline,
&[
(0, KernelArg::Bytes(as_bytes(&reduce_params))),
(1, KernelArg::Buffer(tmp)),
(2, KernelArg::Buffer(output)),
(3, KernelArg::Bytes(as_bytes(&nwg))),
],
reduce_tg,
reduce_tg_size,
);
}
Ok(())
}
pub fn tmp_buffer_bytes(num_heads: u32, head_dim: u32) -> usize {
let nrows = num_heads as usize;
let max_nwg = 32usize;
let dv = head_dim as usize;
(nrows * max_nwg * (dv + 2)) * std::mem::size_of::<f32>()
}
#[allow(clippy::too_many_arguments)]
pub fn flash_attn_vec_tq_hb_with_fused_undo(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
q: &MlxBuffer,
k_packed: &MlxBuffer,
k_norms: &MlxBuffer,
v_packed: &MlxBuffer,
v_norms: &MlxBuffer,
output: &MlxBuffer,
tmp: &MlxBuffer,
params: &FlashAttnVecTqHbParams,
) -> Result<()> {
validate_params(params)?;
let head_dim = params.head_dim;
let nwg = compute_nwg(params.kv_seq_len);
let gpu_params = FlashAttnVecTqHbParamsGpu {
n_heads: params.num_heads,
n_kv_heads: params.num_kv_heads,
head_dim: params.head_dim,
kv_seq_len: params.kv_seq_len,
kv_capacity: params.kv_capacity,
scale: params.scale,
mask_type: params.mask_type,
sliding_window: params.sliding_window,
softcap: params.softcap,
nwg,
ring_start: params.ring_start,
scale_factor_d512: params.scale_factor_d512,
codebook_bits: params.codebook_bits,
fuse_fwht_pre: params.fuse_fwht_pre,
nsg: params.nsg,
};
let kernel_name = match head_dim {
256 => "flash_attn_vec_tq_hb_dk256",
512 => "flash_attn_vec_tq_hb_dk512",
_ => return Err(MlxError::InvalidArgument(format!(
"flash_attn_vec_tq_hb_with_fused_undo: unsupported head_dim {head_dim}"
))),
};
let cbits_const = (params.codebook_bits as i32, 50usize);
let pipeline = registry
.get_pipeline_with_constants(
kernel_name,
device.metal_device(),
&[],
&[(cbits_const.1, cbits_const.0)],
)?;
let pk = pad2(head_dim as usize, 128);
let pv = pad2(head_dim as usize, 128);
let sh = 4 * 32;
let nsg = params.nsg as usize;
let shmem_halfs = pk + nsg * (sh + 2 * pv);
let shmem_bytes = shmem_halfs * 2;
encoder.set_op_kind(CapturedOpKind::Sdpa);
let threadgroups = MTLSize::new(1, params.num_heads as u64, nwg as u64);
let threadgroup_size = MTLSize::new(32, params.nsg as u64, 1);
let dst_buf = if nwg == 1 { output } else { tmp };
encoder.encode_threadgroups_with_args_and_shared(
pipeline,
&[
(0, KernelArg::Bytes(as_bytes(&gpu_params))),
(1, KernelArg::Buffer(q)),
(2, KernelArg::Buffer(k_packed)),
(3, KernelArg::Buffer(k_norms)),
(4, KernelArg::Buffer(v_packed)),
(5, KernelArg::Buffer(v_norms)),
(6, KernelArg::Buffer(dst_buf)),
],
&[(0, shmem_bytes as u64)],
threadgroups,
threadgroup_size,
);
if nwg > 1 {
encoder.memory_barrier();
crate::ops::flash_attn_vec_reduce_tq_hb_undo::dispatch_flash_attn_vec_reduce_tq_hb_undo(
encoder, registry, device,
tmp, output,
params.num_heads, head_dim, nwg,
)?;
} else {
encoder.memory_barrier();
crate::ops::fwht_standalone::dispatch_fwht_sign_undo_f32(
encoder, registry, device.metal_device(),
output, params.num_heads, head_dim,
)?;
}
Ok(())
}
#[repr(C)]
#[derive(Debug, Clone, Copy, bytemuck::Pod, bytemuck::Zeroable)]
struct FlashAttnVecTqHbBatchedParamsGpu {
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
kv_seq_len: u32,
kv_capacity: u32,
scale: f32,
mask_type: u32,
sliding_window: u32,
softcap: f32,
nwg: u32,
ring_start: u32,
scale_factor_d512: f32,
codebook_bits: u32,
fuse_fwht_pre: u32,
nsg: u32,
n_queries: u32,
}
#[allow(clippy::too_many_arguments)]
pub fn flash_attn_vec_tq_hb_batched(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
n_q: u32,
q: &MlxBuffer,
k_packed: &MlxBuffer,
k_norms: &MlxBuffer,
v_packed: &MlxBuffer,
v_norms: &MlxBuffer,
output: &MlxBuffer,
tmp: &MlxBuffer,
slot_id_arr: &MlxBuffer,
seq_pos_arr: &MlxBuffer,
params: &FlashAttnVecTqHbParams,
) -> Result<()> {
validate_params(params)?;
let head_dim = params.head_dim;
let nwg = compute_nwg(params.kv_seq_len);
let gpu_params = FlashAttnVecTqHbBatchedParamsGpu {
n_heads: params.num_heads,
n_kv_heads: params.num_kv_heads,
head_dim: params.head_dim,
kv_seq_len: params.kv_seq_len,
kv_capacity: params.kv_capacity,
scale: params.scale,
mask_type: params.mask_type,
sliding_window: params.sliding_window,
softcap: params.softcap,
nwg,
ring_start: params.ring_start,
scale_factor_d512: params.scale_factor_d512,
codebook_bits: params.codebook_bits,
fuse_fwht_pre: params.fuse_fwht_pre,
nsg: params.nsg,
n_queries: n_q,
};
let kernel_name = match head_dim {
256 => "flash_attn_vec_tq_hb_batched_dk256",
512 => "flash_attn_vec_tq_hb_batched_dk512",
_ => return Err(MlxError::InvalidArgument(format!(
"flash_attn_vec_tq_hb_batched: unsupported head_dim {head_dim}"
))),
};
let cbits_const = (params.codebook_bits as i32, 50usize);
let pipeline = registry
.get_pipeline_with_constants(
kernel_name,
device.metal_device(),
&[],
&[(cbits_const.1, cbits_const.0)],
)?;
let pk = pad2(head_dim as usize, 128);
let pv = pad2(head_dim as usize, 128);
let sh = 4 * 32;
let nsg = params.nsg as usize;
let shmem_halfs = pk + nsg * (sh + 2 * pv);
let shmem_bytes = shmem_halfs * 2;
encoder.set_op_kind(CapturedOpKind::Sdpa);
let threadgroups = MTLSize::new(n_q as u64, params.num_heads as u64, nwg as u64);
let threadgroup_size = MTLSize::new(32, params.nsg as u64, 1);
let dst_buf = if nwg == 1 { output } else { tmp };
encoder.encode_threadgroups_with_args_and_shared(
pipeline,
&[
(0, KernelArg::Bytes(as_bytes(&gpu_params))),
(1, KernelArg::Buffer(q)),
(2, KernelArg::Buffer(k_packed)),
(3, KernelArg::Buffer(k_norms)),
(4, KernelArg::Buffer(v_packed)),
(5, KernelArg::Buffer(v_norms)),
(6, KernelArg::Buffer(dst_buf)),
(7, KernelArg::Buffer(slot_id_arr)),
(8, KernelArg::Buffer(seq_pos_arr)),
],
&[(0, shmem_bytes as u64)],
threadgroups,
threadgroup_size,
);
if nwg > 1 {
encoder.memory_barrier();
crate::ops::flash_attn_vec_reduce_tq_hb_undo::dispatch_flash_attn_vec_reduce_tq_hb_undo(
encoder, registry, device,
tmp, output,
n_q * params.num_heads, head_dim, nwg,
)?;
} else {
encoder.memory_barrier();
crate::ops::fwht_standalone::dispatch_fwht_sign_undo_f32(
encoder, registry, device.metal_device(),
output, n_q * params.num_heads, head_dim,
)?;
}
Ok(())
}
fn pad2(x: usize, n: usize) -> usize {
(x + n - 1) & !(n - 1)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_gpu_params_size() {
assert_eq!(std::mem::size_of::<FlashAttnVecTqHbParamsGpu>(), 60);
}
#[test]
fn test_validate_bad_bits() {
let p = FlashAttnVecTqHbParams {
num_heads: 8,
num_kv_heads: 4,
head_dim: 256,
kv_seq_len: 64,
kv_capacity: 1024,
scale: 1.0,
mask_type: 0,
sliding_window: 0,
softcap: 0.0,
ring_start: 0,
scale_factor_d512: 1.0,
codebook_bits: 4, fuse_fwht_pre: 0,
nsg: 1,
};
assert!(validate_params(&p).is_err());
}
#[test]
fn test_validate_ok_8bit() {
let p = FlashAttnVecTqHbParams {
num_heads: 8,
num_kv_heads: 4,
head_dim: 256,
kv_seq_len: 64,
kv_capacity: 1024,
scale: 1.0,
mask_type: 0,
sliding_window: 0,
softcap: 0.0,
ring_start: 0,
scale_factor_d512: 1.0,
codebook_bits: 8,
fuse_fwht_pre: 0,
nsg: 1,
};
assert!(validate_params(&p).is_ok());
}
#[test]
fn test_validate_nsg_zero_rejected() {
let p = FlashAttnVecTqHbParams {
num_heads: 8, num_kv_heads: 4, head_dim: 256,
kv_seq_len: 64, kv_capacity: 1024, scale: 1.0, mask_type: 0,
sliding_window: 0, softcap: 0.0, ring_start: 0,
scale_factor_d512: 1.0, codebook_bits: 8, fuse_fwht_pre: 0,
nsg: 0,
};
assert!(validate_params(&p).is_err(), "nsg=0 must reject");
}
#[test]
fn test_validate_nsg_non_pow2_rejected() {
for nsg in [3u32, 5, 6, 7, 9, 16, 31, 33] {
let p = FlashAttnVecTqHbParams {
num_heads: 8, num_kv_heads: 4, head_dim: 256,
kv_seq_len: 64, kv_capacity: 1024, scale: 1.0, mask_type: 0,
sliding_window: 0, softcap: 0.0, ring_start: 0,
scale_factor_d512: 1.0, codebook_bits: 8, fuse_fwht_pre: 0,
nsg,
};
assert!(validate_params(&p).is_err(), "nsg={nsg} must reject (not pow-2 or > 4)");
}
}
#[test]
fn test_validate_nsg_pow2_accepted() {
for nsg in [1u32, 2, 4] {
let p = FlashAttnVecTqHbParams {
num_heads: 8, num_kv_heads: 4, head_dim: 256,
kv_seq_len: 64, kv_capacity: 1024, scale: 1.0, mask_type: 0,
sliding_window: 0, softcap: 0.0, ring_start: 0,
scale_factor_d512: 1.0, codebook_bits: 8, fuse_fwht_pre: 0,
nsg,
};
assert!(validate_params(&p).is_ok(), "nsg={nsg} (pow-2 ≤ 4) must accept");
}
}
static NSG_ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
#[test]
fn test_compute_nsg_adaptive_threshold() {
let _guard = NSG_ENV_LOCK.lock().unwrap();
std::env::remove_var("HF2Q_TQ_NSG");
for kl in [1u32, 64, 256, 1024] {
assert_eq!(compute_nsg(kl), 1, "compute_nsg({kl}) must be 1 (kL ≤ 1024)");
}
for kl in [1025u32, 1536, 2048, 4096, 8192, 16384] {
assert_eq!(compute_nsg(kl), 4, "compute_nsg({kl}) must be 4 (kL > 1024)");
}
}
#[test]
fn test_compute_nsg_env_override() {
use std::sync::atomic::Ordering;
let _guard = NSG_ENV_LOCK.lock().unwrap();
let invalidate = || CACHED_TQ_NSG.store(-1, Ordering::Relaxed);
invalidate();
std::env::set_var("HF2Q_TQ_NSG", "4");
assert_eq!(compute_nsg(64), 4);
invalidate();
std::env::set_var("HF2Q_TQ_NSG", "2");
assert_eq!(compute_nsg(64), 2);
invalidate();
std::env::set_var("HF2Q_TQ_NSG", "1");
assert_eq!(compute_nsg(64), 1);
invalidate();
std::env::set_var("HF2Q_TQ_NSG", "3");
assert_eq!(compute_nsg(64), 1);
invalidate();
std::env::remove_var("HF2Q_TQ_NSG");
}
}