From 41032be30967bd2eecf30999aa51ff4c218087af Mon Sep 17 00:00:00 2001 From: Dan Yeh Date: Thu, 30 Jul 2026 22:10:53 +0800 Subject: [PATCH] state_kv_projection fusion gemm append_kv and cleanups --- .../backends/cpu/kernel/attention/qkv_norm.rs | 5 +- .../metal/kernel/attention/qkv_norm.metal | 10 +- .../backend-uzu/src/encodable_block/dflash.rs | 52 ++++-- .../encodable_block/mixer/attention/mod.rs | 19 +- .../encodable_block/mixer/attention/mode.rs | 168 ++++-------------- .../mixer/attention/qkv_norm.rs | 72 ++++---- .../backends/common/kernel/qkv_norm_test.rs | 3 +- 7 files changed, 118 insertions(+), 211 deletions(-) diff --git a/crates/backend-uzu/src/backends/cpu/kernel/attention/qkv_norm.rs b/crates/backend-uzu/src/backends/cpu/kernel/attention/qkv_norm.rs index 0fae3b4a1..99c9dca96 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/attention/qkv_norm.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/attention/qkv_norm.rs @@ -19,8 +19,7 @@ pub fn qkv_norm< #[optional(!scale_free)] scales: Option<*const ScaleT>, qkv_output: *mut OutputT, batch_size: u32, - num_q_heads: u32, - num_kv_heads: u32, + total_heads: u32, head_dim: u32, epsilon: f32, scale_offset: f32, @@ -38,7 +37,7 @@ pub fn qkv_norm< let head_dim = head_dim as usize; let head_offset = head_offset as usize; let head_count = head_count as usize; - let qkv_stride = (num_q_heads + 2 * num_kv_heads) as usize * head_dim; + let qkv_stride = total_heads as usize * head_dim; let head_dim_accum = AccumT::from(head_dim).unwrap(); let epsilon = AccumT::from(epsilon).unwrap(); let scale_offset = AccumT::from(scale_offset).unwrap(); diff --git a/crates/backend-uzu/src/backends/metal/kernel/attention/qkv_norm.metal b/crates/backend-uzu/src/backends/metal/kernel/attention/qkv_norm.metal index 2cd29ea4c..770be0f72 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/attention/qkv_norm.metal +++ b/crates/backend-uzu/src/backends/metal/kernel/attention/qkv_norm.metal @@ -20,8 +20,7 @@ PUBLIC KERNEL(QKVNorm)( const device ScaleT* scales OPTIONAL(!scale_free), device OutputT* qkv_output, constant uint& batch_size, - constant uint& num_q_heads, - constant uint& num_kv_heads, + constant uint& total_heads, constant uint& head_dim, constant float& epsilon, constant float& scale_offset, @@ -41,13 +40,8 @@ PUBLIC KERNEL(QKVNorm)( if (head_count == 0u || head_dim == 0u) return; - const uint total_heads_in_buffer = num_q_heads + 2u * num_kv_heads; - const uint logical_head_idx = head_offset + head_idx; - if (logical_head_idx >= total_heads_in_buffer) - return; - const ulong slice_offset = - (ulong)batch_idx * (ulong)total_heads_in_buffer * (ulong)head_dim + (ulong)logical_head_idx * (ulong)head_dim; + (ulong)batch_idx * (ulong)total_heads * (ulong)head_dim + (ulong)(head_offset + head_idx) * (ulong)head_dim; const device InputT* input_data = qkv_input + slice_offset; const device ScaleT* scales_data = scales; diff --git a/crates/backend-uzu/src/encodable_block/dflash.rs b/crates/backend-uzu/src/encodable_block/dflash.rs index d125f0f4a..0461c7328 100644 --- a/crates/backend-uzu/src/encodable_block/dflash.rs +++ b/crates/backend-uzu/src/encodable_block/dflash.rs @@ -45,6 +45,8 @@ impl DFlashState { pub struct DFlashDraft { target_feature_projection: Box>, projected_feature_norm: Normalization, + state_kv_projection: Box>, + layer_kv_dim: usize, layers: Box<[DFlashDraftLayer]>, output_norm: Normalization, top_k: ::RadixTopKSmall, @@ -119,7 +121,6 @@ impl DFlashDraft { "DFlash block_size must be in 1..=attention suffix capacity", )); } - let target_feature_projection = >::new( config.model_dim * config.target_layer_ids.len(), [config.model_dim], @@ -139,6 +140,16 @@ impl DFlashDraft { context, )?; let layers_tree = parameter_tree.subtree("layers")?; + let layer_kv_dim = + 2 * config.layer_configs[0].attention_config.num_groups * config.layer_configs[0].attention_config.head_dim; + let state_kv_projection = >::new( + config.model_dim, + [config.layer_configs.len() * layer_kv_dim], + false, + context, + data_type, + ¶meter_tree.subtree("state_kv_projection")?, + )?; let layers = config .layer_configs .iter() @@ -172,6 +183,8 @@ impl DFlashDraft { Ok(Self { target_feature_projection, projected_feature_norm, + state_kv_projection, + layer_kv_dim, layers, output_norm, top_k, @@ -244,20 +257,27 @@ impl DFlashDraft { let token_positions = (state.context_length..state.context_length + num_tokens).collect::>(); let rope = PrecalculatedRoPE::precalculate(&self.rope_config, &token_positions, encoder)?; - let dflash_layer_count = self.layers.len(); - let mut normalized_features = Some(normalized_features); - for (layer_index, (layer, attention_state)) in self.layers.iter().zip(state.layer_states.iter_mut()).enumerate() + let projected_kv = self.state_kv_projection.encode(normalized_features, num_tokens, encoder)?; + let layer_kv_bytes = self.layer_kv_dim * self.data_type.size_in_bytes(); + let kv_chunk = |chunk_index: usize| chunk_index * layer_kv_bytes..(chunk_index + 1) * layer_kv_bytes; + let mut layer_key_values = (0..self.layers.len()) + .map(|_| encoder.allocate_scratch(num_tokens * layer_kv_bytes)) + .collect::, _>>()?; + for (layer_index, key_value) in layer_key_values.iter_mut().enumerate() { + for token_index in 0..num_tokens { + encoder.encode_copy( + &projected_kv, + kv_chunk(token_index * self.layers.len() + layer_index), + key_value, + kv_chunk(token_index), + ); + } + } + for ((layer, attention_state), key_value) in + self.layers.iter().zip(state.layer_states.iter_mut()).zip(layer_key_values) { attention_state.prepare(state.context_length, num_tokens, encoder.context())?; - let kv_input = if layer_index + 1 == dflash_layer_count { - normalized_features.take().expect("normalized features available for last layer") - } else { - let source = normalized_features.as_ref().expect("normalized features available"); - let mut kv_input = encoder.allocate_scratch(source.size())?; - encoder.encode_copy(source, .., &mut kv_input, ..); - kv_input - }; - layer.attention.append_kv(kv_input, Some(&rope), num_tokens, attention_state, encoder)?; + layer.attention.append_projected_kv(key_value, &rope, num_tokens, attention_state, encoder)?; } state.context_length += num_tokens; @@ -376,12 +396,6 @@ impl DFlashDraftLayer { parameter_tree: &ParameterTree, data_type: DataType, ) -> Result> { - if config.attention_config.is_causal { - return Err(DFlashDraftNewError::InvalidAttentionConfig("DFlash attention must be non-causal")); - } - if config.attention_config.is_kv_sharing { - return Err(DFlashDraftNewError::InvalidAttentionConfig("DFlash attention must not use KV sharing")); - } let (attention, _) = Attention::new( model_dim, data_type, diff --git a/crates/backend-uzu/src/encodable_block/mixer/attention/mod.rs b/crates/backend-uzu/src/encodable_block/mixer/attention/mod.rs index 9c0c60cd8..700c870ac 100644 --- a/crates/backend-uzu/src/encodable_block/mixer/attention/mod.rs +++ b/crates/backend-uzu/src/encodable_block/mixer/attention/mod.rs @@ -14,7 +14,7 @@ use crate::{ Mixer, MixerState, attention::{ core::{AttentionCoreNewArguments, AttentionCores}, - mode::{LinearProjection, QkvProjection}, + mode::LinearProjection, qkv_norm::{QKVNorm, QKVNormError}, rope::PrecalculatedRoPE, }, @@ -41,7 +41,8 @@ pub struct Attention { sliding_window_size: Option, max_rope_length: Option, data_type: DataType, - projection: QkvProjection, + qkv: LinearProjection, + prepare: ::AttentionPrepareKernel, gate_projection: Option>>, sinks: Option>, flat_core: AttentionCores, @@ -158,14 +159,6 @@ impl Attention { rope_config.is_some(), ) .map_err(AttentionNewError::Backend)?; - let projection = QkvProjection::Packed { - qkv: LinearProjection { - lin: qkv_projection, - norm: qkv_norm, - }, - prepare, - }; - let sinks = config .has_sinks .then(|| parameter_tree.leaf("sinks")?.validate(&[num_q_heads], data_type)?.read_allocation()) @@ -230,7 +223,11 @@ impl Attention { sliding_window_size, max_rope_length, data_type, - projection, + qkv: LinearProjection { + lin: qkv_projection, + norm: qkv_norm, + }, + prepare, gate_projection, sinks, flat_core, diff --git a/crates/backend-uzu/src/encodable_block/mixer/attention/mode.rs b/crates/backend-uzu/src/encodable_block/mixer/attention/mode.rs index cc77db861..e2b205bc5 100644 --- a/crates/backend-uzu/src/encodable_block/mixer/attention/mode.rs +++ b/crates/backend-uzu/src/encodable_block/mixer/attention/mode.rs @@ -1,7 +1,7 @@ use crate::{ array::size_for_shape, backends::common::{ - Allocation, Backend, BufferArgMut, Encoder, Kernels, + Allocation, Backend, BufferArgMut, Encoder, gpu_types::trie::TrieNode, kernel::{AttentionPrepareKernel, SigmoidGateKernel}, }, @@ -19,8 +19,6 @@ use crate::{ utils::maybe_mut::MaybeMut, }; -type PrepareKernel = <::Kernels as Kernels>::AttentionPrepareKernel; - pub(super) struct LinearProjection { pub(super) lin: Box>, pub(super) norm: Option>, @@ -41,22 +39,6 @@ impl LinearProjection { } } -pub(super) enum QkvProjection { - /// Fused QKV — or Q-only under KV sharing (`num_kv_heads == None`). - Packed { - qkv: LinearProjection, - prepare: PrepareKernel, - }, - /// Separate Q and KV projections. - #[allow(dead_code)] // TODO: remove when wiring with DFlash. - Split { - q: LinearProjection, - kv: LinearProjection, - q_prepare: PrepareKernel, - kv_prepare: PrepareKernel, - }, -} - impl Attention { pub(super) fn attend( &self, @@ -76,17 +58,10 @@ impl Attention { (hidden, None) }; - let mut attention_output = match (&self.projection, state) { - ( - QkvProjection::Packed { - qkv, - prepare, - }, - Some(MaybeMut::Mut(state)), - ) => { - let qkv = qkv.project(hidden, batch_dim.size(), encoder)?; + let mut attention_output = match state { + Some(MaybeMut::Mut(state)) => { + let qkv = self.qkv.project(hidden, batch_dim.size(), encoder)?; let queries = self.prepare_kv_and_queries( - prepare, &qkv, state.keys.as_mut(), state.values.as_mut(), @@ -98,31 +73,19 @@ impl Attention { )?; self.run_core(&queries, batch_dim, state, encoder)? }, - ( - QkvProjection::Packed { - qkv, - prepare, - }, - Some(MaybeMut::Const(state)), - ) => { + Some(MaybeMut::Const(state)) => { // KV sharing: the packed projection produces queries only. - let query = qkv.project(hidden, batch_dim.size(), encoder)?; - let queries = self.prepare_queries(prepare, &query, precalculated_rope, batch_dim.size(), encoder)?; + let query = self.qkv.project(hidden, batch_dim.size(), encoder)?; + let queries = self.prepare_queries(&query, precalculated_rope, batch_dim.size(), encoder)?; self.run_core(&queries, batch_dim, state, encoder)? }, - ( - QkvProjection::Packed { - qkv, - prepare, - }, - None, - ) => { + None => { let Some(num_kv_heads) = self.num_kv_heads else { panic!("stateless attention doesn't support query-only projection"); }; assert!(batch_dim.is_flat(), "stateless attention doesn't support trie"); - let qkv = qkv.project(hidden, batch_dim.size(), encoder)?; + let qkv = self.qkv.project(hidden, batch_dim.size(), encoder)?; let mut keys = encoder.allocate_scratch(size_for_shape( &[batch_dim.size(), num_kv_heads, self.head_dim], self.data_type, @@ -133,7 +96,6 @@ impl Attention { ))?; let queries = self.prepare_kv_and_queries( - prepare, &qkv, &mut keys, &mut values, @@ -170,37 +132,6 @@ impl Attention { encoder, )? }, - ( - QkvProjection::Split { - q, - kv, - q_prepare, - kv_prepare, - }, - Some(MaybeMut::Mut(state)), - ) => { - // Linear::encode may consume/mutate its input; split Q/KV attention needs the same hidden for both projections. - let mut hidden_for_key_value = encoder.allocate_scratch(hidden.size())?; - encoder.encode_copy(&hidden, .., &mut hidden_for_key_value, ..); - let query = q.project(hidden, batch_dim.size(), encoder)?; - let key_value = kv.project(hidden_for_key_value, batch_dim.size(), encoder)?; - let precalculated_rope = precalculated_rope.expect("split attention requires RoPE"); - let queries = - self.prepare_queries(q_prepare, &query, Some(precalculated_rope), batch_dim.size(), encoder)?; - self.prepare_kv_and_queries( - kv_prepare, - &key_value, - state.keys.as_mut(), - state.values.as_mut(), - state.state_type.physical_prefix_length(), - 0, - Some(precalculated_rope), - batch_dim.size(), - encoder, - )?; - self.run_core(&queries, batch_dim, state, encoder)? - }, - _ => panic!("attention projection/state combination is invalid"), }; if let Some(gate_kernel) = &self.gate_kernel { @@ -214,52 +145,27 @@ impl Attention { self.out_projection.encode(attention_output, batch_dim.size(), encoder) } - pub(crate) fn append_kv( + pub fn append_projected_kv( &self, - hidden: Allocation, - precalculated_rope: Option<&PrecalculatedRoPE>, + mut key_value: Allocation, + precalculated_rope: &PrecalculatedRoPE, batch_dim: usize, state: &mut AttentionState, encoder: &mut Encoder, ) -> Result<(), B::Error> { - match &self.projection { - QkvProjection::Split { - kv, - kv_prepare, - .. - } => { - let precalculated_rope = precalculated_rope.expect("split attention requires RoPE"); - let key_value = kv.project(hidden, batch_dim, encoder)?; - self.prepare_kv_and_queries( - kv_prepare, - &key_value, - state.keys.as_mut(), - state.values.as_mut(), - state.state_type.physical_prefix_length(), - 0, - Some(precalculated_rope), - batch_dim, - encoder, - )?; - }, - QkvProjection::Packed { - qkv, - prepare, - } => { - let projected = qkv.project(hidden, batch_dim, encoder)?; - self.prepare_kv_and_queries( - prepare, - &projected, - state.keys.as_mut(), - state.values.as_mut(), - state.state_type.physical_prefix_length(), - self.num_q_heads as u32, - precalculated_rope, - batch_dim, - encoder, - )?; - }, + if let Some(norm) = &self.qkv.norm { + norm.encode_key_value(&mut key_value, batch_dim, encoder)?; } + self.prepare_kv_and_queries( + &key_value, + state.keys.as_mut(), + state.values.as_mut(), + state.state_type.physical_prefix_length(), + 0, + Some(precalculated_rope), + batch_dim, + encoder, + )?; state.append_full(batch_dim); Ok(()) } @@ -293,10 +199,8 @@ impl Attention { ) } - /// With `num_q_heads == 0`, only keys/values are scattered into the cache (KV append). fn prepare_kv_and_queries<'keys, 'values>( &self, - prepare: &PrepareKernel, input: &Allocation, keys: impl BufferArgMut<'keys, B>, values: impl BufferArgMut<'values, B>, @@ -309,9 +213,9 @@ impl Attention { let mut queries = if num_q_heads == 0 { encoder.allocate_scratch(self.data_type.size_in_bytes())? } else { - self.allocate_queries(batch_dim, encoder)? + encoder.allocate_scratch(size_for_shape(&[self.num_q_heads, batch_dim, self.head_dim], self.data_type))? }; - prepare.encode( + self.prepare.encode( input, &mut queries, Some(keys), @@ -331,36 +235,28 @@ impl Attention { fn prepare_queries( &self, - prepare: &PrepareKernel, query: &Allocation, precalculated_rope: Option<&PrecalculatedRoPE>, batch_dim: usize, encoder: &mut Encoder, ) -> Result, B::Error> { - let mut queries = self.allocate_queries(batch_dim, encoder)?; - prepare.encode( + let mut queries = + encoder.allocate_scratch(size_for_shape(&[self.num_q_heads, batch_dim, self.head_dim], self.data_type))?; + self.prepare.encode( query, &mut queries, None::<&mut Allocation>, None::<&mut Allocation>, - precalculated_rope.map(|precalculated_rope| &precalculated_rope.cosines), - precalculated_rope.map(|precalculated_rope| &precalculated_rope.sines), + precalculated_rope.map(|rope| &rope.cosines), + precalculated_rope.map(|rope| &rope.sines), self.num_q_heads as u32, None, self.head_dim as u32, - precalculated_rope.map(|precalculated_rope| precalculated_rope.dim as u32), + precalculated_rope.map(|rope| rope.dim as u32), None, batch_dim as u32, encoder, ); Ok(queries) } - - fn allocate_queries( - &self, - batch_dim: usize, - encoder: &mut Encoder, - ) -> Result, B::Error> { - encoder.allocate_scratch(size_for_shape(&[self.num_q_heads, batch_dim, self.head_dim], self.data_type)) - } } diff --git a/crates/backend-uzu/src/encodable_block/mixer/attention/qkv_norm.rs b/crates/backend-uzu/src/encodable_block/mixer/attention/qkv_norm.rs index a052bcc67..73fa3748a 100644 --- a/crates/backend-uzu/src/encodable_block/mixer/attention/qkv_norm.rs +++ b/crates/backend-uzu/src/encodable_block/mixer/attention/qkv_norm.rs @@ -119,42 +119,50 @@ impl QKVNorm { batch_dim: usize, encoder: &mut Encoder, ) -> Result<(), B::Error> { - if let Some(query) = &self.query { - self.encode_head(query, qkv, batch_dim, 0, self.num_q_heads as u32, encoder); - } - if let Some(key) = &self.key { - self.encode_head(key, qkv, batch_dim, self.num_q_heads as u32, self.num_kv_heads as u32, encoder); - } - if let Some(value) = &self.value { - let value_offset = (self.num_q_heads + self.num_kv_heads) as u32; - self.encode_head(value, qkv, batch_dim, value_offset, self.num_kv_heads as u32, encoder); - } - Ok(()) + self.encode_packed(qkv, batch_dim, self.num_q_heads, encoder) } - fn encode_head( + pub fn encode_key_value( &self, - head: &Head, - qkv: &mut Allocation, + key_value: &mut Allocation, + batch_dim: usize, + encoder: &mut Encoder, + ) -> Result<(), B::Error> { + self.encode_packed(key_value, batch_dim, 0, encoder) + } + + fn encode_packed( + &self, + buffer: &mut Allocation, batch_dim: usize, - range_start: u32, - range_end: u32, + q_heads: usize, encoder: &mut Encoder, - ) { - head.kernel.encode( - None::<&Allocation>, - head.scales.as_ref(), - &mut *qkv, - batch_dim as u32, - self.num_q_heads as u32, - self.num_kv_heads as u32, - self.head_dim as u32, - head.config.epsilon, - head.config.scale_offset.unwrap_or(0.0), - range_start, - range_end, - head.config.upcast_mode == UpcastMode::FullLayer, - encoder, - ); + ) -> Result<(), B::Error> { + let kv = self.num_kv_heads; + let total_heads = q_heads + 2 * kv; + let heads = [(&self.query, 0, q_heads), (&self.key, q_heads, kv), (&self.value, q_heads + kv, kv)]; + for (head, head_offset, head_count) in heads { + let Some(head) = head else { + continue; + }; + if head_count == 0 { + continue; + } + head.kernel.encode( + None::<&Allocation>, + head.scales.as_ref(), + &mut *buffer, + batch_dim as u32, + total_heads as u32, + self.head_dim as u32, + head.config.epsilon, + head.config.scale_offset.unwrap_or(0.0), + head_offset as u32, + head_count as u32, + head.config.upcast_mode == UpcastMode::FullLayer, + encoder, + ); + } + Ok(()) } } diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/qkv_norm_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/qkv_norm_test.rs index f95ca75ef..3b6f09a46 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/qkv_norm_test.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/qkv_norm_test.rs @@ -140,8 +140,7 @@ fn get_output< scales.as_ref(), &mut qkv, input.batch_size, - input.num_q_heads, - input.num_kv_heads, + input.num_q_heads + 2 * input.num_kv_heads, input.head_dim, input.epsilon, input.scale_offset,