use metal::MTLSize;
use crate::buffer::MlxBuffer;
use crate::device::MlxDevice;
use crate::encoder::{CommandEncoder, KernelArg, as_bytes};
use crate::error::{MlxError, Result};
use crate::kernel_registry::KernelRegistry;
use crate::DType;
pub static FLASH_ATTN_PREFILL_MASK_SHADER_SOURCE: &str =
include_str!("../shaders/flash_attn_prefill_mask.metal");
pub const K_FILL_BF16: &str = "flash_attn_prefill_mask_fill_bf16";
pub const K_FILL_F16: &str = "flash_attn_prefill_mask_fill_f16";
pub const K_FILL_BLOCKDIAG_BF16: &str = "flash_attn_prefill_mask_fill_blockdiag_bf16";
pub fn register(registry: &mut KernelRegistry) {
registry.register_source(K_FILL_BF16, FLASH_ATTN_PREFILL_MASK_SHADER_SOURCE);
registry.register_source(K_FILL_F16, FLASH_ATTN_PREFILL_MASK_SHADER_SOURCE);
registry.register_source(K_FILL_BLOCKDIAG_BF16, FLASH_ATTN_PREFILL_MASK_SHADER_SOURCE);
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SdpaMaskParams {
pub seq_len_q: u32,
pub seq_len_k: u32,
pub window_size: Option<u32>,
pub causal: bool,
pub q_abs_offset: u32,
}
#[repr(C)]
#[derive(Debug, Clone, Copy, bytemuck::Pod, bytemuck::Zeroable)]
struct MaskFillParamsGpu {
seq_len_k: u32,
q_abs_offset: u32,
n_swa: i32,
causal: u32,
}
pub fn build_sdpa_mask_bf16(
device: &MlxDevice,
registry: &mut KernelRegistry,
encoder: &mut CommandEncoder,
params: &SdpaMaskParams,
) -> Result<MlxBuffer> {
if params.seq_len_q == 0 {
return Err(MlxError::InvalidArgument(
"build_sdpa_mask_bf16: seq_len_q must be > 0".into(),
));
}
if params.seq_len_k == 0 {
return Err(MlxError::InvalidArgument(
"build_sdpa_mask_bf16: seq_len_k must be > 0".into(),
));
}
if let Some(0) = params.window_size {
return Err(MlxError::InvalidArgument(
"build_sdpa_mask_bf16: window_size=Some(0) is not allowed \
(llama.cpp treats n_swa=0 as undefined; pass None for \
no-window / causal-only)".into(),
));
}
let total_elems = (params.seq_len_q as u64)
.checked_mul(params.seq_len_k as u64)
.ok_or_else(|| {
MlxError::InvalidArgument(format!(
"build_sdpa_mask_bf16: seq_len_q ({}) * seq_len_k ({}) overflows u64",
params.seq_len_q, params.seq_len_k
))
})?;
let byte_len = (total_elems as usize)
.checked_mul(2)
.ok_or_else(|| {
MlxError::InvalidArgument(format!(
"build_sdpa_mask_bf16: mask size ({} elems × 2 B) overflows usize",
total_elems
))
})?;
let mask = device.alloc_buffer(
byte_len,
DType::BF16,
vec![params.seq_len_q as usize, params.seq_len_k as usize],
)?;
let fill_params = MaskFillParamsGpu {
seq_len_k: params.seq_len_k,
q_abs_offset: params.q_abs_offset,
n_swa: match params.window_size {
None => -1,
Some(w) => w.min(i32::MAX as u32) as i32,
},
causal: if params.causal { 1 } else { 0 },
};
let pipeline = registry.get_pipeline(K_FILL_BF16, device.metal_device())?;
let tg_x = {
let want = params.seq_len_k.next_power_of_two().max(32);
want.min(256)
};
let threadgroups = MTLSize::new(params.seq_len_q as u64, 1, 1);
let tg_size = MTLSize::new(tg_x as u64, 1, 1);
encoder.encode_threadgroups_with_args(
pipeline,
&[
(0, KernelArg::Buffer(&mask)),
(1, KernelArg::Bytes(as_bytes(&fill_params))),
],
threadgroups,
tg_size,
);
Ok(mask)
}
pub fn build_sdpa_mask_f16(
device: &MlxDevice,
registry: &mut KernelRegistry,
encoder: &mut CommandEncoder,
params: &SdpaMaskParams,
) -> Result<MlxBuffer> {
if params.seq_len_q == 0 || params.seq_len_k == 0 {
return Err(MlxError::InvalidArgument(
"build_sdpa_mask_f16: sequence lengths must be > 0".into(),
));
}
if params.window_size == Some(0) {
return Err(MlxError::InvalidArgument(
"build_sdpa_mask_f16: window_size=Some(0) is not allowed".into(),
));
}
let total = (params.seq_len_q as usize)
.checked_mul(params.seq_len_k as usize)
.ok_or_else(|| MlxError::InvalidArgument("build_sdpa_mask_f16: mask size overflow".into()))?;
let mask = device.alloc_buffer(
total
.checked_mul(2)
.ok_or_else(|| MlxError::InvalidArgument("build_sdpa_mask_f16: mask size overflow".into()))?,
DType::F16,
vec![params.seq_len_q as usize, params.seq_len_k as usize],
)?;
let fill_params = MaskFillParamsGpu {
seq_len_k: params.seq_len_k,
q_abs_offset: params.q_abs_offset,
n_swa: params
.window_size
.map(|w| w.min(i32::MAX as u32) as i32)
.unwrap_or(-1),
causal: u32::from(params.causal),
};
let pipeline = registry.get_pipeline(K_FILL_F16, device.metal_device())?;
let tg_x = params.seq_len_k.next_power_of_two().max(32).min(256);
encoder.encode_threadgroups_with_args(
pipeline,
&[
(0, KernelArg::Buffer(&mask)),
(1, KernelArg::Bytes(as_bytes(&fill_params))),
],
MTLSize::new(params.seq_len_q as u64, 1, 1),
MTLSize::new(tg_x as u64, 1, 1),
);
Ok(mask)
}
#[repr(C)]
#[derive(Debug, Clone, Copy, bytemuck::Pod, bytemuck::Zeroable)]
struct BlockDiagMaskParamsGpu {
seq_len: u32,
n_swa: i32,
causal: u32,
}
pub fn build_block_diagonal_sdpa_mask_bf16(
device: &MlxDevice,
registry: &mut KernelRegistry,
encoder: &mut CommandEncoder,
seq_id: &MlxBuffer,
local_pos: &MlxBuffer,
t: u32,
window_size: Option<u32>,
causal: bool,
) -> Result<MlxBuffer> {
if t == 0 {
return Err(MlxError::InvalidArgument(
"build_block_diagonal_sdpa_mask_bf16: t must be > 0".into(),
));
}
if let Some(0) = window_size {
return Err(MlxError::InvalidArgument(
"build_block_diagonal_sdpa_mask_bf16: window_size=Some(0) is not \
allowed (pass None for causal-only)".into(),
));
}
let byte_len = (t as usize)
.checked_mul(t as usize)
.and_then(|x| x.checked_mul(2))
.ok_or_else(|| {
MlxError::InvalidArgument(format!(
"build_block_diagonal_sdpa_mask_bf16: mask size (T={t}) overflows usize"
))
})?;
let mask = device.alloc_buffer(byte_len, DType::BF16, vec![t as usize, t as usize])?;
let fill_params = BlockDiagMaskParamsGpu {
seq_len: t,
n_swa: match window_size {
None => -1,
Some(w) => w.min(i32::MAX as u32) as i32,
},
causal: if causal { 1 } else { 0 },
};
let pipeline = registry.get_pipeline(K_FILL_BLOCKDIAG_BF16, device.metal_device())?;
let tg_x = t.next_power_of_two().max(32).min(256);
let threadgroups = MTLSize::new(t as u64, 1, 1);
let tg_size = MTLSize::new(tg_x as u64, 1, 1);
encoder.encode_threadgroups_with_args(
pipeline,
&[
(0, KernelArg::Buffer(&mask)),
(1, KernelArg::Bytes(as_bytes(&fill_params))),
(2, KernelArg::Buffer(seq_id)),
(3, KernelArg::Buffer(local_pos)),
],
threadgroups,
tg_size,
);
Ok(mask)
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
mod tests {
use super::*;
#[test]
fn test_mask_fill_params_gpu_size() {
assert_eq!(std::mem::size_of::<MaskFillParamsGpu>(), 16);
}
#[test]
fn test_mask_fill_params_encoding_global() {
let p = MaskFillParamsGpu {
seq_len_k: 2048,
q_abs_offset: 0,
n_swa: -1,
causal: 1,
};
assert_eq!(p.n_swa, -1, "global mask encodes n_swa=-1");
assert_eq!(p.causal, 1);
}
#[test]
fn test_mask_fill_params_encoding_sliding() {
let p = MaskFillParamsGpu {
seq_len_k: 2048,
q_abs_offset: 0,
n_swa: 1024,
causal: 1,
};
assert_eq!(p.n_swa, 1024, "sliding mask encodes n_swa>0");
}
#[test]
fn test_reject_zero_seq_len_q() {
let p = SdpaMaskParams {
seq_len_q: 0,
seq_len_k: 8,
window_size: None,
causal: true,
q_abs_offset: 0,
};
assert_eq!(p.seq_len_q, 0);
}
#[test]
fn test_register_adds_kernel_name() {
let mut registry = KernelRegistry::new();
register(&mut registry);
assert_eq!(K_FILL_BF16, "flash_attn_prefill_mask_fill_bf16");
}
#[test]
fn test_f16_mask_matches_bf16_predicate_exactly() {
let device = match MlxDevice::new() {
Ok(device) => device,
Err(error) => {
eprintln!("skipping: no Metal device: {error}");
return;
}
};
let mut registry = KernelRegistry::new();
register(&mut registry);
for params in [
SdpaMaskParams { seq_len_q: 5, seq_len_k: 13, window_size: None, causal: true, q_abs_offset: 7 },
SdpaMaskParams { seq_len_q: 6, seq_len_k: 17, window_size: Some(4), causal: true, q_abs_offset: 9 },
SdpaMaskParams { seq_len_q: 3, seq_len_k: 11, window_size: Some(5), causal: false, q_abs_offset: 4 },
] {
let mut encoder = device.command_encoder().expect("encoder");
let bf16 = build_sdpa_mask_bf16(&device, &mut registry, &mut encoder, ¶ms)
.expect("BF16 mask");
let f16 = build_sdpa_mask_f16(&device, &mut registry, &mut encoder, ¶ms)
.expect("F16 mask");
encoder.commit_and_wait().expect("commit");
let bf16_values = bf16.as_slice::<half::bf16>().expect("BF16 slice");
let f16_values = f16.as_slice::<half::f16>().expect("F16 slice");
assert_eq!(bf16_values.len(), f16_values.len());
for (index, (&bf16_value, &f16_value)) in bf16_values.iter().zip(f16_values).enumerate() {
let bf16_masked = bf16_value.to_f32().is_infinite() && bf16_value.to_f32().is_sign_negative();
let f16_masked = f16_value.to_f32().is_infinite() && f16_value.to_f32().is_sign_negative();
assert_eq!(f16_masked, bf16_masked, "params={params:?} index={index}");
if !f16_masked {
assert_eq!(f16_value.to_bits(), half::f16::ZERO.to_bits());
assert_eq!(bf16_value.to_bits(), half::bf16::ZERO.to_bits());
}
}
}
}
}