use metal::{Buffer, ComputeCommandEncoderRef, ComputePipelineState, MTLSize};
use crate::CeraError;
use crate::backend::metal::{MetalContext, MetalParams, TqAttnParams, TqParams, shaders};
use crate::kv_cache::checked_elems;
use crate::model::{BlockType, ModelConfig};
use crate::turboquant::{
CompressedKeyCache, CompressedValueCache, RotationState, TurboQuantConfig,
decode_compressed_keys, decode_compressed_values, encode_compressed_keys,
encode_compressed_values,
};
pub use crate::turboquant::{TqLayout, TqMode, head_dim_supported};
pub const TQ_THREADS: u64 = 128;
pub const TQ_ATTN_THREADS: u64 = 256;
struct TqPipelines {
encode_keys: ComputePipelineState,
encode_values: ComputePipelineState,
rotate_q: ComputePipelineState,
attention: ComputePipelineState,
}
pub struct TqLayerCache {
pub keys: Buffer,
pub values: Buffer,
pub n_kv_heads: usize,
}
pub struct TqMetalCache {
pub mode: TqMode,
pub layout: TqLayout,
pub layers: Vec<Option<TqLayerCache>>,
signs: Buffer,
qrot: Buffer,
q_cap: usize,
pub config: TurboQuantConfig,
pipelines: TqPipelines,
max_seq_len: usize,
}
impl TqMetalCache {
pub fn new(
ctx: &MetalContext,
config: &ModelConfig,
max_seq_len: usize,
q_cap: usize,
mode: TqMode,
) -> Result<Self, CeraError> {
let head_dim = config.head_dim;
if !head_dim_supported(head_dim) {
return Err(crate::turboquant::unsupported_head_dim(head_dim));
}
let layout = TqLayout::new(head_dim);
let n_layers = config.block_types.len();
let mut signs = Vec::with_capacity(n_layers * 2 * head_dim);
for layer_idx in 0..n_layers {
let rot = RotationState::try_from_seed(mode.seed ^ layer_idx as u64, head_dim)?;
signs.extend_from_slice(&rot.polar_signs);
signs.extend_from_slice(&rot.jl_signs);
}
let mut layers = Vec::with_capacity(n_layers);
for (i, bt) in config.block_types.iter().enumerate() {
match bt {
BlockType::Attention => {
let n_kv_heads = config.kv_heads_per_layer[i];
let vecs = checked_elems::<u32>(n_kv_heads, max_seq_len)?;
let k_bytes = TqLayout::words_to_bytes(layout.key_words(vecs)?);
let v_bytes = TqLayout::words_to_bytes(layout.value_words(vecs)?);
layers.push(Some(TqLayerCache {
keys: ctx.create_buffer(k_bytes),
values: ctx.create_buffer(v_bytes),
n_kv_heads,
}));
}
BlockType::GatedConv => layers.push(None),
}
}
let q_rows = checked_elems::<f32>(q_cap, config.n_heads)?;
let qrot_floats = checked_elems::<f32>(q_rows, 2 * head_dim + 1)?;
let pipelines = TqPipelines {
encode_keys: pipeline(ctx, shaders::TURBOQUANT, "tq_encode_keys")?,
encode_values: pipeline(ctx, shaders::TURBOQUANT, "tq_encode_values")?,
rotate_q: pipeline(ctx, shaders::TURBOQUANT, "tq_rotate_q")?,
attention: pipeline(ctx, shaders::FLASH_ATTENTION_TQ, "flash_attention_tq")?,
};
Ok(Self {
mode,
layout,
layers,
signs: ctx.upload_f32(&signs),
qrot: ctx.create_buffer(TqLayout::words_to_bytes(qrot_floats)),
q_cap,
config: TurboQuantConfig::for_head_dim(head_dim),
pipelines,
max_seq_len,
})
}
pub fn layer(&self, layer: usize) -> Option<&TqLayerCache> {
self.layers[layer].as_ref()
}
fn base_params(
&self,
n_tokens: usize,
n_heads: usize,
layer: usize,
start_pos: usize,
) -> TqParams {
let head_dim = self.layout.head_dim;
let c = &self.config.centroids;
let b = &self.config.boundaries;
TqParams {
n_tokens: n_tokens as u32,
n_heads: n_heads as u32,
head_dim: head_dim as u32,
src_stride: (n_heads * head_dim) as u32,
dst_pos: start_pos as u32,
max_seq_len: self.max_seq_len as u32,
sign_off: (layer * 2 * head_dim) as u32,
q_cap: self.q_cap as u32,
c0: c[0],
c1: c[1],
c2: c[2],
c3: c[3],
b0: b[0],
b1: b[1],
b2: b[2],
_pad: 0,
}
}
pub fn encode_kv(
&self,
enc: &ComputeCommandEncoderRef,
layer: usize,
k_src: &Buffer,
v_src: &Buffer,
n_tokens: usize,
start_pos: usize,
) {
let l = self
.layers
.get(layer)
.and_then(|l| l.as_ref())
.expect("encode_kv on a layer without a compressed cache");
assert!(
start_pos + n_tokens <= self.max_seq_len,
"encode_kv writes timesteps [{start_pos}, {}) past the cache capacity {}",
start_pos + n_tokens,
self.max_seq_len,
);
let params = self.base_params(n_tokens, l.n_kv_heads, layer, start_pos);
let groups = (n_tokens * l.n_kv_heads) as u64;
self.dispatch(
enc,
&self.pipelines.encode_keys,
k_src,
&l.keys,
¶ms,
groups,
);
self.dispatch(
enc,
&self.pipelines.encode_values,
v_src,
&l.values,
¶ms,
groups,
);
}
pub fn rotate_queries(
&self,
enc: &ComputeCommandEncoderRef,
layer: usize,
q_src: &Buffer,
n_tokens: usize,
n_heads: usize,
) {
assert!(
n_tokens <= self.q_cap,
"rotate_queries: {n_tokens} rows exceeds the qrot scratch capacity {}",
self.q_cap
);
let params = self.base_params(n_tokens, n_heads, layer, 0);
self.dispatch(
enc,
&self.pipelines.rotate_q,
q_src,
&self.qrot,
¶ms,
(n_tokens * n_heads) as u64,
);
}
fn dispatch(
&self,
enc: &ComputeCommandEncoderRef,
pipeline: &ComputePipelineState,
src: &Buffer,
dst: &Buffer,
params: &TqParams,
groups: u64,
) {
enc.set_compute_pipeline_state(pipeline);
enc.set_buffer(0, Some(src), 0);
enc.set_buffer(1, Some(dst), 0);
enc.set_buffer(2, Some(&self.signs), 0);
params.set(enc, 3);
enc.dispatch_thread_groups(MTLSize::new(groups, 1, 1), MTLSize::new(TQ_THREADS, 1, 1));
}
#[allow(clippy::too_many_arguments)]
pub fn attention(
&self,
enc: &ComputeCommandEncoderRef,
layer: usize,
out: &Buffer,
n_tokens: usize,
n_heads: usize,
start_pos: usize,
scale: f32,
) {
let l = self
.layers
.get(layer)
.and_then(|l| l.as_ref())
.expect("attention on a layer without a compressed cache");
let head_dim = self.layout.head_dim;
let params = TqAttnParams {
n_heads: n_heads as u32,
n_kv_heads: l.n_kv_heads as u32,
head_dim: head_dim as u32,
max_seq: (start_pos + n_tokens) as u32,
start_pos: start_pos as u32,
scale,
q_cap: self.q_cap as u32,
out_stride: (n_heads * head_dim) as u32,
qjl_scale: crate::turboquant::qjl_scale(head_dim),
sign_off: (layer * 2 * head_dim) as u32,
c0: self.config.centroids[0],
c1: self.config.centroids[1],
c2: self.config.centroids[2],
c3: self.config.centroids[3],
q_base: 0,
cache_cap: self.max_seq_len as u32,
};
enc.set_compute_pipeline_state(&self.pipelines.attention);
enc.set_buffer(0, Some(&self.qrot), 0);
enc.set_buffer(1, Some(&l.keys), 0);
enc.set_buffer(2, Some(&l.values), 0);
enc.set_buffer(3, Some(out), 0);
params.set(enc, 4);
enc.set_buffer(5, Some(&self.signs), 0);
enc.dispatch_thread_groups(
MTLSize::new(n_heads as u64, n_tokens as u64, 1),
MTLSize::new(TQ_ATTN_THREADS, 1, 1),
);
}
}
impl TqMetalCache {
fn region(buf: &Buffer, word_off: usize, words: usize) -> &[u32] {
let end = (word_off + words) * 4;
assert!(
end as u64 <= buf.length(),
"snapshot would read to byte {end} of a {}-byte buffer",
buf.length()
);
unsafe { std::slice::from_raw_parts((buf.contents() as *const u32).add(word_off), words) }
}
pub fn snapshot_layer(&self, layer: usize, seq_len: usize) -> (Vec<u8>, Vec<u8>) {
let l = self
.layer(layer)
.expect("snapshot_layer on a layer without a compressed cache");
assert!(
seq_len <= self.max_seq_len,
"snapshot_layer: seq_len {seq_len} exceeds the cache capacity {}",
self.max_seq_len
);
let head_dim = self.layout.head_dim;
let n = l.n_kv_heads;
let mut keys = CompressedKeyCache::new(n, head_dim, seq_len);
let mut values = CompressedValueCache::new(n, head_dim, seq_len);
if seq_len > 0 {
let pw = self.layout.polar_words;
let jw = self.layout.jl_words;
let cap = self.max_seq_len;
let (jl_off, norm_off) = self.layout.key_regions(n * cap);
let v_norm_off = self.layout.value_norm_offset(n * cap);
for h in 0..n {
let k_polar = Self::region(&l.keys, h * cap * pw, seq_len * pw);
let k_jl = Self::region(&l.keys, jl_off + h * cap * jw, seq_len * jw);
let k_norms = Self::region(&l.keys, norm_off + h * cap, seq_len);
let v_polar = Self::region(&l.values, h * cap * pw, seq_len * pw);
let v_norms = Self::region(&l.values, v_norm_off + h * cap, seq_len);
for t in 0..seq_len {
let nw = k_norms[t];
keys.append(
h,
bytemuck::cast_slice(&k_polar[t * pw..(t + 1) * pw]),
bytemuck::cast_slice(&k_jl[t * jw..(t + 1) * jw]),
(nw & 0xFFFF) as u16,
(nw >> 16) as u16,
);
values.append(
h,
bytemuck::cast_slice(&v_polar[t * pw..(t + 1) * pw]),
(v_norms[t] & 0xFFFF) as u16,
);
}
}
}
(
encode_compressed_keys(&keys),
encode_compressed_values(&values),
)
}
pub fn restore_layer(
&self,
layer: usize,
keys_blob: &[u8],
values_blob: &[u8],
) -> Option<usize> {
let l = self.layer(layer)?;
let keys = decode_compressed_keys(keys_blob)?;
let values = decode_compressed_values(values_blob)?;
let seq_len = self
.layout
.blobs_match(&keys, &values, l.n_kv_heads, self.max_seq_len)?;
if seq_len == 0 {
return Some(0);
}
let pw = self.layout.polar_words;
let jw = self.layout.jl_words;
let cap = self.max_seq_len;
let n = l.n_kv_heads;
let (jl_off, norm_off) = self.layout.key_regions(n * cap);
let v_norm_off = self.layout.value_norm_offset(n * cap);
let mut norm_words = Vec::with_capacity(seq_len);
for h in 0..n {
write_words(&l.keys, h * cap * pw, &keys.polar_data[h]);
write_words(&l.keys, jl_off + h * cap * jw, &keys.jl_data[h]);
norm_words.clear();
for t in 0..seq_len {
norm_words.push(
u32::from(keys.norms[h][t]) | (u32::from(keys.residual_norms[h][t]) << 16),
);
}
write_words(
&l.keys,
norm_off + h * cap,
bytemuck::cast_slice(&norm_words),
);
write_words(&l.values, h * cap * pw, &values.polar_data[h]);
norm_words.clear();
norm_words.extend(values.norms[h].iter().map(|&b| u32::from(b)));
write_words(
&l.values,
v_norm_off + h * cap,
bytemuck::cast_slice(&norm_words),
);
}
Some(seq_len)
}
}
fn pipeline(
ctx: &MetalContext,
src: &'static str,
entry: &str,
) -> Result<ComputePipelineState, CeraError> {
ctx.create_pipeline(src, entry)
.map_err(|e| CeraError::Backend(format!("TurboQuant kernel '{entry}': {e}")))
}
fn write_words(buf: &Buffer, word_off: usize, bytes: &[u8]) {
let end = word_off * 4 + bytes.len();
assert!(
end as u64 <= buf.length(),
"restore would write to byte {end} of a {}-byte buffer",
buf.length()
);
unsafe {
std::ptr::copy_nonoverlapping(
bytes.as_ptr(),
(buf.contents() as *mut u8).add(word_off * 4),
bytes.len(),
);
}
}