use ruda_model::tensor::{Bool,Tensor,backend::Backend};
use crate::{cache::ProjectedKvCache,transformer::AdaptedGroupedQueryAttention};
use super::{GroupedQueryAttention,DenseAttentionMask,DenseAttentionOptions};
fn cached_masks<B: Backend>(mut masks: DenseAttentionMask<B>,visible: Tensor<B,2,Bool>) -> DenseAttentionMask<B> {
masks.key_valid = Some(if let Some(additional) = masks.key_valid {
assert_eq!(additional.dims(),visible.dims(),"cached additional key validity must describe the complete retained prefix");
assert_eq!(additional.device(),visible.device(),"cached additional visibility device differs");
visible.bool_and(additional)
} else { visible });
masks
}
fn append<B: Backend>(query: &Tensor<B,4>,key: Tensor<B,4>,value: Tensor<B,4>,
new_visible: Option<Tensor<B,2,Bool>>,cache: &mut ProjectedKvCache<B>,masks: DenseAttentionMask<B>,
geometry: (usize,usize,usize)) -> (Tensor<B,4>,Tensor<B,4>,DenseAttentionMask<B>) {
let [batch,heads,_,width] = query.dims();
assert_eq!((heads,width),(geometry.0,geometry.2),"cached projected query head geometry differs");
assert_eq!((key.dims()[0],key.dims()[1],key.dims()[3]),(batch,geometry.1,geometry.2),"cached projected key geometry differs");
assert_eq!((value.dims()[0],value.dims()[1],value.dims()[3]),(batch,geometry.1,geometry.2),"cached projected value geometry differs");
assert_eq!(query.dtype(),key.dtype(),"cached Q/K storage differs");
assert_eq!(query.device(),key.device(),"cached Q/K device differs");
cache.validate_append(&key,&value,new_visible.as_ref());
if let Some(additional) = &masks.key_valid {
let length = cache.len().checked_add(key.dims()[2]).expect("cached score length overflow");
assert_eq!(additional.dims(),[batch,length],"cached additional mask must cover old and new key slots");
assert_eq!(additional.device(),query.device(),"cached additional mask device differs");
}
let (key,value,visible) = cache.append(key,value,new_visible);
(key,value,cached_masks(masks,visible))
}
impl<B: Backend> GroupedQueryAttention<B> {
pub fn project_query(&self,input: Tensor<B,3>) -> Tensor<B,4> {
let [batch,tokens,_] = input.dims();
self.query.forward(input).reshape([batch,tokens,self.query_heads,self.head_dimension]).swap_dims(1,2)
}
pub fn project_key_value(&self,key: Tensor<B,3>,value: Tensor<B,3>) -> (Tensor<B,4>,Tensor<B,4>) {
let [batch,tokens,_] = key.dims();
assert_eq!((value.dims()[0],value.dims()[1]),(batch,tokens),"memory K/V input rows differ");
(self.key.forward(key).reshape([batch,tokens,self.kv_heads,self.head_dimension]).swap_dims(1,2),
self.value.forward(value).reshape([batch,tokens,self.kv_heads,self.head_dimension]).swap_dims(1,2))
}
pub fn forward_cached_projected(&self,query: Tensor<B,4>,key: Tensor<B,4>,value: Tensor<B,4>,
new_visible: Option<Tensor<B,2,Bool>>,cache: &mut ProjectedKvCache<B>,masks: DenseAttentionMask<B>,
options: DenseAttentionOptions) -> Tensor<B,3> {
let (key,value,masks) = append(&query,key,value,new_visible,cache,masks,(self.query_heads,self.kv_heads,self.head_dimension));
self.forward_projected(query,key,value,masks,options)
}
pub fn forward_cached(&self,input: Tensor<B,3>,new_visible: Option<Tensor<B,2,Bool>>,
cache: &mut ProjectedKvCache<B>,masks: DenseAttentionMask<B>,options: DenseAttentionOptions) -> Tensor<B,3> {
self.forward_cached_with_positions(input,new_visible,cache,masks,options,|query,key,_|(query,key))
}
pub fn forward_cached_with_positions<F>(&self,input: Tensor<B,3>,new_visible: Option<Tensor<B,2,Bool>>,
cache: &mut ProjectedKvCache<B>,masks: DenseAttentionMask<B>,options: DenseAttentionOptions,positions: F) -> Tensor<B,3>
where F: FnOnce(Tensor<B,4>,Tensor<B,4>,usize)->(Tensor<B,4>,Tensor<B,4>) {
let (query,key,value) = self.project(input.clone(),input.clone(),input);
let shape = (query.dims(),key.dims());
let (query,key) = positions(query,key,cache.position());
assert_eq!((query.dims(),key.dims()),shape,"cached position transform changed new head geometry");
self.forward_cached_projected(query,key,value,new_visible,cache,masks,options)
}
pub fn forward_cached_memory(&self,query: Tensor<B,4>,memory: &ProjectedKvCache<B>,
masks: DenseAttentionMask<B>,options: DenseAttentionOptions) -> Tensor<B,3> {
let (key,value,visible) = memory.prefix().expect("prepare actual projected memory before cached cross-attention");
self.forward_projected(query,key,value,cached_masks(masks,visible),options)
}
}
impl<B: Backend> AdaptedGroupedQueryAttention<B> {
pub fn project_query(&self,input: Tensor<B,3>) -> Tensor<B,4> {
let [batch,tokens,_] = input.dims();
self.query.forward(input).reshape([batch,tokens,self.query_heads,self.head_dimension]).swap_dims(1,2)
}
pub fn project_key_value(&self,key: Tensor<B,3>,value: Tensor<B,3>) -> (Tensor<B,4>,Tensor<B,4>) {
let [batch,tokens,_] = key.dims();
assert_eq!((value.dims()[0],value.dims()[1]),(batch,tokens),"adapted memory K/V rows differ");
(self.key.forward(key).reshape([batch,tokens,self.kv_heads,self.head_dimension]).swap_dims(1,2),
self.value.forward(value).reshape([batch,tokens,self.kv_heads,self.head_dimension]).swap_dims(1,2))
}
pub fn forward_cached_projected(&self,query: Tensor<B,4>,key: Tensor<B,4>,value: Tensor<B,4>,
new_visible: Option<Tensor<B,2,Bool>>,cache: &mut ProjectedKvCache<B>,masks: DenseAttentionMask<B>,
options: DenseAttentionOptions) -> Tensor<B,3> {
let (key,value,masks) = append(&query,key,value,new_visible,cache,masks,(self.query_heads,self.kv_heads,self.head_dimension));
self.forward_projected(query,key,value,masks,options)
}
pub fn forward_cached(&self,input: Tensor<B,3>,new_visible: Option<Tensor<B,2,Bool>>,
cache: &mut ProjectedKvCache<B>,masks: DenseAttentionMask<B>,options: DenseAttentionOptions) -> Tensor<B,3> {
self.forward_cached_with_positions(input,new_visible,cache,masks,options,|query,key,_|(query,key))
}
pub fn forward_cached_with_positions<F>(&self,input: Tensor<B,3>,new_visible: Option<Tensor<B,2,Bool>>,
cache: &mut ProjectedKvCache<B>,masks: DenseAttentionMask<B>,options: DenseAttentionOptions,positions: F) -> Tensor<B,3>
where F: FnOnce(Tensor<B,4>,Tensor<B,4>,usize)->(Tensor<B,4>,Tensor<B,4>) {
let (query,key,value) = self.project(input.clone(),input.clone(),input);
let shape = (query.dims(),key.dims());
let (query,key) = positions(query,key,cache.position());
assert_eq!((query.dims(),key.dims()),shape,"adapted cached positions changed new head geometry");
self.forward_cached_projected(query,key,value,new_visible,cache,masks,options)
}
pub fn forward_cached_memory(&self,query: Tensor<B,4>,memory: &ProjectedKvCache<B>,
masks: DenseAttentionMask<B>,options: DenseAttentionOptions) -> Tensor<B,3> {
let (key,value,visible) = memory.prefix().expect("prepare actual adapted projected memory before cached cross-attention");
self.forward_projected(query,key,value,cached_masks(masks,visible),options)
}
}