use crate::CeraError;
use crate::backend::wgpu::GpuContext;
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(crate) use crate::turboquant::describe_kv_mode;
pub use crate::turboquant::{TqLayout, TqMode, head_dim_supported};
const SLOTS_PER_LAYER: usize = 4;
const SLOT_ENCODE_KEYS: usize = 0;
const SLOT_ENCODE_VALUES: usize = 1;
const SLOT_ROTATE_Q: usize = 2;
const SLOT_ATTENTION: usize = 3;
const PARAMS_BYTES: usize = 64;
#[repr(C)]
#[derive(Clone, Copy, Default, bytemuck::Pod, bytemuck::Zeroable)]
pub struct TqParams {
pub n_tokens: u32,
pub n_heads: u32,
pub head_dim: u32,
pub src_stride: u32,
pub dst_pos: u32,
pub max_seq_len: u32,
pub sign_off: u32,
pub q_cap: u32,
pub c0: f32,
pub c1: f32,
pub c2: f32,
pub c3: f32,
pub b0: f32,
pub b1: f32,
pub b2: f32,
pub _pad: u32,
}
const _: () = assert!(size_of::<TqParams>() == PARAMS_BYTES);
impl TqParams {
pub fn with_quant_config(mut self, config: &TurboQuantConfig) -> Self {
let c = &config.centroids;
let b = &config.boundaries;
(self.c0, self.c1, self.c2, self.c3) = (c[0], c[1], c[2], c[3]);
(self.b0, self.b1, self.b2) = (b[0], b[1], b[2]);
self
}
}
#[derive(Clone, Copy)]
pub struct TqAttnParams {
pub n_heads: u32,
pub n_kv_heads: u32,
pub head_dim: u32,
pub max_seq: u32,
pub start_pos: u32,
pub scale: f32,
pub q_cap: u32,
pub out_stride: u32,
pub qjl_scale: f32,
pub sign_off: u32,
pub centroids: [f32; 4],
pub q_base: u32,
pub cache_cap: u32,
}
impl TqAttnParams {
pub fn to_u32_array(self) -> [u32; 16] {
[
self.n_heads,
self.n_kv_heads,
self.head_dim,
self.max_seq,
self.start_pos,
self.scale.to_bits(),
self.q_cap,
self.out_stride,
self.qjl_scale.to_bits(),
self.sign_off,
self.centroids[0].to_bits(),
self.centroids[1].to_bits(),
self.centroids[2].to_bits(),
self.centroids[3].to_bits(),
self.q_base,
self.cache_cap,
]
}
pub fn qjl_scale_for(head_dim: usize) -> f32 {
crate::turboquant::qjl_scale(head_dim)
}
}
struct TqPipelines {
encode_keys: wgpu::ComputePipeline,
encode_values: wgpu::ComputePipeline,
rotate_q: wgpu::ComputePipeline,
attention: wgpu::ComputePipeline,
}
pub struct TqLayerCache {
pub keys: wgpu::Buffer,
pub values: wgpu::Buffer,
pub n_kv_heads: usize,
}
pub struct TqGpuCache {
pub mode: TqMode,
pub layout: TqLayout,
pub layers: Vec<Option<TqLayerCache>>,
signs: wgpu::Buffer,
qrot: wgpu::Buffer,
q_cap: usize,
params: wgpu::Buffer,
slot_stride: usize,
pub config: TurboQuantConfig,
pipelines: TqPipelines,
max_seq_len: usize,
}
impl TqGpuCache {
pub fn new(
ctx: &GpuContext,
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_storage_rw(k_bytes, &format!("l{i}.tq_keys")),
values: ctx.create_storage_rw(v_bytes, &format!("l{i}.tq_values")),
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 qrot = ctx.create_storage_rw(TqLayout::words_to_bytes(qrot_floats), "tq_qrot");
let align = ctx.min_storage_buffer_offset_alignment.max(4) as usize;
let slot_stride = PARAMS_BYTES.next_multiple_of(align);
let params = ctx.create_storage_rw(
(n_layers * SLOTS_PER_LAYER * slot_stride) as u64,
"tq_params",
);
let pipelines = TqPipelines {
encode_keys: ctx.create_pipeline(
crate::backend::wgpu::shaders::TURBOQUANT,
"tq_encode_keys",
"tq_encode_keys",
),
encode_values: ctx.create_pipeline(
crate::backend::wgpu::shaders::TURBOQUANT,
"tq_encode_values",
"tq_encode_values",
),
rotate_q: ctx.create_pipeline(
crate::backend::wgpu::shaders::TURBOQUANT,
"tq_rotate_q",
"tq_rotate_q",
),
attention: ctx.create_pipeline(
crate::backend::wgpu::shaders::FLASH_ATTENTION_TQ,
"flash_attention_tq",
"flash_attention_tq",
),
};
Ok(Self {
mode,
layout,
layers,
signs: ctx.upload_f32(&signs, "tq_signs"),
qrot,
q_cap,
params,
slot_stride,
config: TurboQuantConfig::for_head_dim(head_dim),
pipelines,
max_seq_len,
})
}
pub fn layer(&self, layer: usize) -> Option<&TqLayerCache> {
self.layers[layer].as_ref()
}
pub fn write_params(
&self,
ctx: &GpuContext,
config: &ModelConfig,
n_tokens: usize,
start_pos: usize,
scale: f32,
) {
let head_dim = self.layout.head_dim;
let q_dim = config.n_heads * head_dim;
let mut slab = vec![0u8; self.layers.len() * SLOTS_PER_LAYER * self.slot_stride];
for (i, layer) in self.layers.iter().enumerate() {
let Some(l) = layer else { continue };
let sign_off = (i * 2 * head_dim) as u32;
let base = TqParams {
n_tokens: n_tokens as u32,
n_heads: l.n_kv_heads as u32,
head_dim: head_dim as u32,
src_stride: (l.n_kv_heads * head_dim) as u32,
dst_pos: start_pos as u32,
max_seq_len: self.max_seq_len as u32,
sign_off,
q_cap: self.q_cap as u32,
..Default::default()
}
.with_quant_config(&self.config);
self.put(&mut slab, i, SLOT_ENCODE_KEYS, bytemuck::bytes_of(&base));
self.put(&mut slab, i, SLOT_ENCODE_VALUES, bytemuck::bytes_of(&base));
let rot = TqParams {
n_heads: config.n_heads as u32,
src_stride: q_dim as u32,
..base
};
self.put(&mut slab, i, SLOT_ROTATE_Q, bytemuck::bytes_of(&rot));
let attn = TqAttnParams {
n_heads: config.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: q_dim as u32,
qjl_scale: TqAttnParams::qjl_scale_for(head_dim),
sign_off,
centroids: self.config.centroids,
q_base: 0,
cache_cap: self.max_seq_len as u32,
};
self.put(
&mut slab,
i,
SLOT_ATTENTION,
bytemuck::cast_slice(&attn.to_u32_array()),
);
}
ctx.queue.write_buffer(&self.params, 0, &slab);
}
fn put(&self, slab: &mut [u8], layer: usize, slot: usize, bytes: &[u8]) {
let off = (layer * SLOTS_PER_LAYER + slot) * self.slot_stride;
slab[off..off + bytes.len()].copy_from_slice(bytes);
}
fn params_binding(&self, layer: usize, slot: usize) -> wgpu::BindingResource<'_> {
let off = ((layer * SLOTS_PER_LAYER + slot) * self.slot_stride) as u64;
wgpu::BindingResource::Buffer(wgpu::BufferBinding {
buffer: &self.params,
offset: off,
size: Some(std::num::NonZeroU64::new(PARAMS_BYTES as u64).unwrap()),
})
}
pub fn encode_kv(
&self,
ctx: &GpuContext,
enc: &mut wgpu::CommandEncoder,
layer: usize,
k_src: &wgpu::Buffer,
v_src: &wgpu::Buffer,
n_tokens: usize,
) {
let l = self
.layers
.get(layer)
.and_then(|l| l.as_ref())
.expect("encode_kv on a layer without a compressed cache");
let groups = n_tokens * l.n_kv_heads;
self.dispatch(
ctx,
enc,
&self.pipelines.encode_keys,
k_src,
&l.keys,
layer,
SLOT_ENCODE_KEYS,
groups,
"tq_encode_keys",
);
self.dispatch(
ctx,
enc,
&self.pipelines.encode_values,
v_src,
&l.values,
layer,
SLOT_ENCODE_VALUES,
groups,
"tq_encode_values",
);
}
pub fn rotate_queries(
&self,
ctx: &GpuContext,
enc: &mut wgpu::CommandEncoder,
layer: usize,
q_src: &wgpu::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
);
self.dispatch(
ctx,
enc,
&self.pipelines.rotate_q,
q_src,
&self.qrot,
layer,
SLOT_ROTATE_Q,
n_tokens * n_heads,
"tq_rotate_q",
);
}
#[allow(clippy::too_many_arguments)]
fn dispatch(
&self,
ctx: &GpuContext,
enc: &mut wgpu::CommandEncoder,
pipeline: &wgpu::ComputePipeline,
src: &wgpu::Buffer,
dst: &wgpu::Buffer,
layer: usize,
slot: usize,
groups: usize,
label: &str,
) {
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some(label),
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: src.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: dst.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: self.signs.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: self.params_binding(layer, slot),
},
],
});
let grid = crate::backend::wgpu::gemv_row_workgroups(groups as u32);
let mut pass = ctx.begin_pass(enc, label);
pass.set_pipeline(pipeline);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(grid.0, grid.1, grid.2);
}
#[allow(clippy::too_many_arguments)]
pub fn attention(
&self,
ctx: &GpuContext,
enc: &mut wgpu::CommandEncoder,
layer: usize,
out: &wgpu::Buffer,
n_tokens: usize,
n_heads: usize,
) {
let l = self
.layers
.get(layer)
.and_then(|l| l.as_ref())
.expect("attention on a layer without a compressed cache");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("flash_attention_tq"),
layout: &self.pipelines.attention.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.qrot.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: l.keys.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: l.values.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: out.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: self.params_binding(layer, SLOT_ATTENTION),
},
wgpu::BindGroupEntry {
binding: 5,
resource: self.signs.as_entire_binding(),
},
],
});
let mut pass = ctx.begin_pass(enc, "flash_attention_tq");
pass.set_pipeline(&self.pipelines.attention);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(n_heads as u32, n_tokens as u32, 1);
}
}
impl TqGpuCache {
pub fn snapshot_layer(
&self,
ctx: &GpuContext,
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;
if seq_len == 0 {
let keys = CompressedKeyCache::new(n, head_dim, 0);
let values = CompressedValueCache::new(n, head_dim, 0);
return (
encode_compressed_keys(&keys),
encode_compressed_values(&values),
);
}
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);
let k_polar = seq_len * pw;
let k_jl = seq_len * jw;
let total = n * (k_polar + k_jl + seq_len) + n * (k_polar + seq_len);
let scratch = ctx.create_storage_rw(TqLayout::words_to_bytes(total), "tq_snapshot_gather");
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("tq_snapshot_gather"),
});
let mut dst = 0usize;
let mut gather = |src: &wgpu::Buffer, src_word: usize, words: usize, dst: &mut usize| {
enc.copy_buffer_to_buffer(
src,
TqLayout::words_to_bytes(src_word),
&scratch,
TqLayout::words_to_bytes(*dst),
TqLayout::words_to_bytes(words),
);
*dst += words;
};
for h in 0..n {
gather(&l.keys, h * cap * pw, k_polar, &mut dst);
}
for h in 0..n {
gather(&l.keys, jl_off + h * cap * jw, k_jl, &mut dst);
}
for h in 0..n {
gather(&l.keys, norm_off + h * cap, seq_len, &mut dst);
}
for h in 0..n {
gather(&l.values, h * cap * pw, k_polar, &mut dst);
}
for h in 0..n {
gather(&l.values, v_norm_off + h * cap, seq_len, &mut dst);
}
debug_assert_eq!(dst, total, "snapshot gather did not fill the scratch");
ctx.queue.submit(Some(enc.finish()));
let words = ctx.download_u32(&scratch, total);
let mut keys = CompressedKeyCache::new(n, head_dim, seq_len);
let mut values = CompressedValueCache::new(n, head_dim, seq_len);
let kp = 0;
let kj = kp + n * k_polar;
let kn = kj + n * k_jl;
let vp = kn + n * seq_len;
let vn = vp + n * k_polar;
for h in 0..n {
for t in 0..seq_len {
let polar = &words[kp + h * k_polar + t * pw..][..pw];
let jl = &words[kj + h * k_jl + t * jw..][..jw];
let nw = words[kn + h * seq_len + t];
keys.append(
h,
bytemuck::cast_slice(polar),
bytemuck::cast_slice(jl),
(nw & 0xFFFF) as u16,
(nw >> 16) as u16,
);
let v_polar = &words[vp + h * k_polar + t * pw..][..pw];
values.append(
h,
bytemuck::cast_slice(v_polar),
(words[vn + h * seq_len + t] & 0xFFFF) as u16,
);
}
}
(
encode_compressed_keys(&keys),
encode_compressed_values(&values),
)
}
pub fn restore_layer(
&self,
ctx: &GpuContext,
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 {
ctx.queue.write_buffer(
&l.keys,
TqLayout::words_to_bytes(h * cap * pw),
&keys.polar_data[h],
);
ctx.queue.write_buffer(
&l.keys,
TqLayout::words_to_bytes(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),
);
}
ctx.queue.write_buffer(
&l.keys,
TqLayout::words_to_bytes(norm_off + h * cap),
bytemuck::cast_slice(&norm_words),
);
ctx.queue.write_buffer(
&l.values,
TqLayout::words_to_bytes(h * cap * pw),
&values.polar_data[h],
);
norm_words.clear();
norm_words.extend(values.norms[h].iter().map(|&b| u32::from(b)));
ctx.queue.write_buffer(
&l.values,
TqLayout::words_to_bytes(v_norm_off + h * cap),
bytemuck::cast_slice(&norm_words),
);
}
Some(seq_len)
}
}