use alloc::vec::Vec;
use ruda_model::{module::Module,tensor::{Bool,Tensor,backend::Backend}};
use crate::{Dropout,attention::{DenseAttentionMask,DenseAttentionOptions,PackedSequenceLayout,PackedAttentionOptions,PackedDocumentAttentionMask},
cache::{ProjectedKvCache,TransformerKvCache,EncoderDecoderKvCache}};
use super::{TransformerProjectionShape,TransformerProjection,ProjectedGroupedQueryAttention,ProjectedTransformerBlock,DenseTransformerNorm};
use super::dense::try_residual_branch;
#[derive(Module,Debug)]
pub struct ProjectedCrossAttentionBlock<B:Backend,P:Module<B>> {
pub attention:ProjectedGroupedQueryAttention<B,P>,
pub query_norm:DenseTransformerNorm<B>,
pub memory_norm:Option<DenseTransformerNorm<B>>,
pub residual_dropout:Dropout,
pub norm_first:bool,
}
impl<B:Backend,P:TransformerProjectionShape<B>> ProjectedCrossAttentionBlock<B,P> {
pub fn from_parts(attention:ProjectedGroupedQueryAttention<B,P>,query_norm:DenseTransformerNorm<B>,memory_norm:Option<DenseTransformerNorm<B>>,
residual_dropout:Dropout,norm_first:bool) -> Self {
assert_eq!(query_norm.width(),attention.query.dimensions()[0],"cross query norm width differs");
if let Some(norm)=&memory_norm {assert_eq!(norm.width(),attention.key.dimensions()[0],"cross memory norm width differs");}
Self {attention,query_norm,memory_norm,residual_dropout,norm_first}
}
}
impl<B:Backend,P:TransformerProjection<B>> ProjectedCrossAttentionBlock<B,P> {
pub fn forward_with_positions<F>(&self,input:Tensor<B,3>,memory:Tensor<B,3>,masks:DenseAttentionMask<B>,options:DenseAttentionOptions,positions:F)
-> Result<Tensor<B,3>,P::Error> where F:FnOnce(Tensor<B,4>,Tensor<B,4>)->(Tensor<B,4>,Tensor<B,4>) {
let memory=if let Some(norm)=&self.memory_norm {norm.forward(memory)} else {memory};
try_residual_branch(input,&self.query_norm,&self.residual_dropout,self.norm_first,|source| {
let (query,key,value)=self.attention.project(source,memory.clone(),memory)?;
let geometry=(query.dims(),key.dims());let (query,key)=positions(query,key);
assert_eq!((query.dims(),key.dims()),geometry,"cross positions changed actual head geometry");
self.attention.forward_projected(query,key,value,masks,options)
})
}
pub fn forward_packed_with_positions<F>(&self,input:Tensor<B,2>,memory:Tensor<B,2>,query_layout:&PackedSequenceLayout,memory_layout:&PackedSequenceLayout,
masks:Option<&[PackedDocumentAttentionMask<B>]>,options:PackedAttentionOptions,positions:F) -> Result<Tensor<B,2>,P::Error>
where F:FnOnce(Tensor<B,3>,Tensor<B,3>)->(Tensor<B,3>,Tensor<B,3>) {
assert_eq!(query_layout.documents(),memory_layout.documents(),"paired source/target document counts differ");
assert_eq!(input.dims()[0],query_layout.tokens(),"paired target boundaries differ from actual rows");
assert_eq!(memory.dims()[0],memory_layout.tokens(),"paired memory boundaries differ from actual rows");
let memory=if let Some(norm)=&self.memory_norm {norm.forward(memory)} else {memory};
try_residual_branch(input,&self.query_norm,&self.residual_dropout,self.norm_first,|source| {
let (query,key,value)=self.attention.project_packed(source,memory.clone(),memory)?;
let geometry=(query.dims(),key.dims());let (query,key)=positions(query,key);
assert_eq!((query.dims(),key.dims()),geometry,"packed cross positions changed actual geometry");
if let Some(masks)=masks {self.attention.forward_packed_masked_projected(query,key,value,query_layout,memory_layout,masks,options)}
else {self.attention.forward_packed_projected(query,key,value,query_layout,memory_layout,options)}
})
}
pub fn prepare_cached_memory<F>(&self,memory:Tensor<B,3>,visible:Option<Tensor<B,2,Bool>>,start_position:usize,positions:F)
-> Result<ProjectedKvCache<B>,P::Error> where F:FnOnce(Tensor<B,4>,usize)->Tensor<B,4> {
let memory=if let Some(norm)=&self.memory_norm {norm.forward(memory)} else {memory};let [batch,tokens,_]=memory.dims();
let key=self.attention.key.forward(memory.clone())?.reshape([batch,tokens,self.attention.kv_heads,self.attention.head_dimension]).swap_dims(1,2);
let value=self.attention.value.forward(memory)?.reshape([batch,tokens,self.attention.kv_heads,self.attention.head_dimension]).swap_dims(1,2);
let geometry=key.dims();let key=positions(key,start_position);assert_eq!(key.dims(),geometry,"prepared memory positions changed geometry");
Ok(ProjectedKvCache::from_projected(key,value,visible,start_position))
}
pub fn forward_cached_memory<F>(&self,input:Tensor<B,3>,memory:&ProjectedKvCache<B>,masks:DenseAttentionMask<B>,options:DenseAttentionOptions,
query_position:usize,positions:F) -> Result<Tensor<B,3>,P::Error> where F:FnOnce(Tensor<B,4>,usize)->Tensor<B,4> {
try_residual_branch(input,&self.query_norm,&self.residual_dropout,self.norm_first,|source| {
let [batch,tokens,_]=source.dims();
let query=self.attention.query.forward(source)?.reshape([batch,tokens,self.attention.query_heads,self.attention.head_dimension]).swap_dims(1,2);
let geometry=query.dims();let query=positions(query,query_position);assert_eq!(query.dims(),geometry,"cached cross query positions changed geometry");
let (key,value,visible)=memory.prefix().expect("prepare actual encoder memory before cached cross attention");
self.attention.forward_projected(query,key,value,crate::attention::cached_projected_masks(masks,visible),options)
})
}
}
#[derive(Module,Debug)]
pub struct ProjectedEncoderDecoderLayer<B:Backend,P:Module<B>> {
pub backbone:ProjectedTransformerBlock<B,P>,
pub cross_attention:ProjectedCrossAttentionBlock<B,P>,
}
impl<B:Backend,P:TransformerProjectionShape<B>> ProjectedEncoderDecoderLayer<B,P> {
pub fn from_parts(backbone:ProjectedTransformerBlock<B,P>,cross_attention:ProjectedCrossAttentionBlock<B,P>) -> Self {
assert_eq!(backbone.attention.query.dimensions()[0],cross_attention.attention.query.dimensions()[0],"decoder residual widths differ");
Self {backbone,cross_attention}
}
}
impl<B:Backend,P:TransformerProjection<B>> ProjectedEncoderDecoderLayer<B,P> {
pub fn forward_with_positions<F,G>(&self,input:Tensor<B,3>,memory:Tensor<B,3>,self_masks:DenseAttentionMask<B>,self_options:DenseAttentionOptions,
cross_masks:DenseAttentionMask<B>,cross_options:DenseAttentionOptions,self_positions:F,cross_positions:G) -> Result<Tensor<B,3>,P::Error>
where F:FnOnce(Tensor<B,4>,Tensor<B,4>)->(Tensor<B,4>,Tensor<B,4>),G:FnOnce(Tensor<B,4>,Tensor<B,4>)->(Tensor<B,4>,Tensor<B,4>) {
let hidden=self.backbone.forward_attention_with_positions(input,self_masks,self_options,self_positions)?;
let hidden=self.cross_attention.forward_with_positions(hidden,memory,cross_masks,cross_options,cross_positions)?;
self.backbone.forward_feed_forward(hidden)
}
pub fn forward_packed_with_positions<F,G>(&self,input:Tensor<B,2>,memory:Tensor<B,2>,query_layout:&PackedSequenceLayout,memory_layout:&PackedSequenceLayout,
self_masks:Option<&[PackedDocumentAttentionMask<B>]>,self_options:PackedAttentionOptions,cross_masks:Option<&[PackedDocumentAttentionMask<B>]>,
cross_options:PackedAttentionOptions,self_positions:F,cross_positions:G) -> Result<Tensor<B,2>,P::Error>
where F:FnOnce(Tensor<B,3>,Tensor<B,3>)->(Tensor<B,3>,Tensor<B,3>),G:FnOnce(Tensor<B,3>,Tensor<B,3>)->(Tensor<B,3>,Tensor<B,3>) {
let hidden=self.backbone.forward_packed_attention_with_positions(input,query_layout,self_masks,self_options,self_positions)?;
let hidden=self.cross_attention.forward_packed_with_positions(hidden,memory,query_layout,memory_layout,cross_masks,cross_options,cross_positions)?;
self.backbone.forward_feed_forward(hidden)
}
pub fn forward_cached_with_positions<F,G>(&self,input:Tensor<B,3>,new_visible:Option<Tensor<B,2,Bool>>,cache:&mut ProjectedKvCache<B>,memory:&ProjectedKvCache<B>,
self_masks:DenseAttentionMask<B>,self_options:DenseAttentionOptions,cross_masks:DenseAttentionMask<B>,cross_options:DenseAttentionOptions,
self_positions:F,cross_positions:G) -> Result<Tensor<B,3>,P::Error>
where F:FnOnce(Tensor<B,4>,Tensor<B,4>,usize)->(Tensor<B,4>,Tensor<B,4>),G:FnOnce(Tensor<B,4>,usize)->Tensor<B,4> {
let position=cache.position();
let hidden=self.backbone.forward_cached_attention_with_positions(input,new_visible,cache,self_masks,self_options,self_positions)?;
let hidden=self.cross_attention.forward_cached_memory(hidden,memory,cross_masks,cross_options,position,cross_positions)?;
self.backbone.forward_feed_forward(hidden)
}
}
#[derive(Module,Debug)]
pub struct ProjectedEncoderDecoderStack<B:Backend,P:Module<B>> {
pub layers:Vec<ProjectedEncoderDecoderLayer<B,P>>,
}
impl<B:Backend,P:TransformerProjection<B>> ProjectedEncoderDecoderStack<B,P> {
pub fn forward_with<F>(&self,mut input:Tensor<B,3>,memory:Tensor<B,3>,mut layer:F) -> Result<Tensor<B,3>,P::Error>
where F:FnMut(usize,&ProjectedEncoderDecoderLayer<B,P>,Tensor<B,3>,Tensor<B,3>)->Result<Tensor<B,3>,P::Error> {
for (index,block) in self.layers.iter().enumerate() {input=layer(index,block,input,memory.clone())?;}Ok(input)
}
pub fn forward_packed_with<F>(&self,mut input:Tensor<B,2>,memory:Tensor<B,2>,query_layout:&PackedSequenceLayout,memory_layout:&PackedSequenceLayout,mut layer:F)
-> Result<Tensor<B,2>,P::Error>
where F:FnMut(usize,&ProjectedEncoderDecoderLayer<B,P>,Tensor<B,2>,Tensor<B,2>)->Result<Tensor<B,2>,P::Error> {
assert_eq!(query_layout.documents(),memory_layout.documents(),"paired document counts differ");
assert_eq!(input.dims()[0],query_layout.tokens(),"target boundaries differ");assert_eq!(memory.dims()[0],memory_layout.tokens(),"source boundaries differ");
for (index,block) in self.layers.iter().enumerate() {input=layer(index,block,input,memory.clone())?;}Ok(input)
}
pub fn prepare_kv_cache<F>(&self,memory:Tensor<B,3>,visible:Option<Tensor<B,2,Bool>>,start_position:usize,initial_capacity:usize,mut positions:F)
-> Result<EncoderDecoderKvCache<B>,P::Error> where F:FnMut(usize,Tensor<B,4>,usize)->Tensor<B,4> {
let memory=self.layers.iter().enumerate().map(|(index,layer)|layer.cross_attention.prepare_cached_memory(memory.clone(),visible.clone(),start_position,
|key,position|positions(index,key,position))).collect::<Result<Vec<_>,_>>()?;
Ok(EncoderDecoderKvCache::new(TransformerKvCache::new(self.layers.len(),initial_capacity),memory))
}
pub fn forward_cached_with<F>(&self,mut input:Tensor<B,3>,cache:&mut EncoderDecoderKvCache<B>,mut layer:F) -> Result<Tensor<B,3>,P::Error>
where F:FnMut(usize,&ProjectedEncoderDecoderLayer<B,P>,Tensor<B,3>,&mut ProjectedKvCache<B>,&ProjectedKvCache<B>)->Result<Tensor<B,3>,P::Error> {
let (decoder,memory)=cache.parts_mut();decoder.validate_layers(self.layers.len());assert_eq!(memory.len(),self.layers.len(),"cached memory/layer counts differ");
let rows=(input.dims()[0],input.dims()[1]);let next=decoder.position().checked_add(rows.1).expect("cached decoder position overflows");
for (index,block) in self.layers.iter().enumerate() {
input=layer(index,block,input,&mut decoder.layers_mut()[index],&memory[index])?;
assert_eq!((input.dims()[0],input.dims()[1]),rows,"cached decoder changed actual new rows");
}
decoder.finish_chunk(next);Ok(input)
}
}