use alloc::vec::Vec;
use ruda_model::tensor::{Bool,Tensor,backend::Backend};
use crate::{attention::{DenseAttentionMask,DenseAttentionOptions},cache::{ProjectedKvCache,TransformerKvCache,EncoderDecoderKvCache}};
use super::{DenseCrossAttentionBlock,AdaptedCrossAttentionBlock,DecoderCrossAttention,DenseEncoderDecoderLayer,
AdaptedEncoderDecoderLayer,DenseEncoderDecoderStack,AdaptedEncoderDecoderStack};
use super::dense::residual_branch;
impl<B: Backend> DenseCrossAttentionBlock<B> {
pub fn prepare_cached_memory<F>(&self,memory: Tensor<B,3>,visible: Option<Tensor<B,2,Bool>>,
start_position: usize,positions: F) -> ProjectedKvCache<B>
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 (key,value) = self.attention.project_key_value(memory.clone(),memory);
let shape = key.dims();
let key = positions(key,start_position);
assert_eq!(key.dims(),shape,"prepared cross key positions changed actual memory geometry");
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) -> Tensor<B,3>
where F: FnOnce(Tensor<B,4>,usize)->Tensor<B,4> {
residual_branch(input,&self.query_norm,&self.residual_dropout,self.norm_first,|source| {
let query = self.attention.project_query(source);
let shape = query.dims();
let query = positions(query,query_position);
assert_eq!(query.dims(),shape,"cached cross query positions changed actual new geometry");
self.attention.forward_cached_memory(query,memory,masks,options)
})
}
}
impl<B: Backend> AdaptedCrossAttentionBlock<B> {
pub fn prepare_cached_memory<F>(&self,memory: Tensor<B,3>,visible: Option<Tensor<B,2,Bool>>,
start_position: usize,positions: F) -> ProjectedKvCache<B>
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 (key,value) = self.attention.project_key_value(memory.clone(),memory);
let shape = key.dims();
let key = positions(key,start_position);
assert_eq!(key.dims(),shape,"prepared adapted cross key positions changed actual memory geometry");
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) -> Tensor<B,3>
where F: FnOnce(Tensor<B,4>,usize)->Tensor<B,4> {
residual_branch(input,&self.query_norm,&self.residual_dropout,self.norm_first,|source| {
let query = self.attention.project_query(source);
let shape = query.dims();
let query = positions(query,query_position);
assert_eq!(query.dims(),shape,"cached adapted cross positions changed actual query geometry");
self.attention.forward_cached_memory(query,memory,masks,options)
})
}
}
impl<B: Backend> DecoderCrossAttention<B> {
pub fn prepare_cached_memory<F>(&self,memory: Tensor<B,3>,visible: Option<Tensor<B,2,Bool>>,
start_position: usize,positions: F) -> ProjectedKvCache<B>
where F: FnOnce(Tensor<B,4>,usize)->Tensor<B,4> {
match self {
Self::Dense(block)=>block.prepare_cached_memory(memory,visible,start_position,positions),
Self::Adapted(block)=>block.prepare_cached_memory(memory,visible,start_position,positions),
}
}
pub fn forward_cached_memory<F>(&self,input: Tensor<B,3>,memory: &ProjectedKvCache<B>,
masks: DenseAttentionMask<B>,options: DenseAttentionOptions,query_position: usize,positions: F) -> Tensor<B,3>
where F: FnOnce(Tensor<B,4>,usize)->Tensor<B,4> {
match self {
Self::Dense(block)=>block.forward_cached_memory(input,memory,masks,options,query_position,positions),
Self::Adapted(block)=>block.forward_cached_memory(input,memory,masks,options,query_position,positions),
}
}
}
impl<B: Backend> DenseEncoderDecoderLayer<B> {
pub fn forward_cached(&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) -> Tensor<B,3> {
self.forward_cached_with_positions(input,new_visible,cache,memory,self_masks,self_options,cross_masks,cross_options,
|query,key,_|(query,key),|query,_|query)
}
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) -> Tensor<B,3>
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)
}
}
impl<B: Backend> AdaptedEncoderDecoderLayer<B> {
pub fn forward_cached(&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) -> Tensor<B,3> {
self.forward_cached_with_positions(input,new_visible,cache,memory,self_masks,self_options,cross_masks,cross_options,
|query,key,_|(query,key),|query,_|query)
}
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) -> Tensor<B,3>
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)
}
}
impl<B: Backend> DenseEncoderDecoderStack<B> {
pub fn prepare_kv_cache<F>(&self,memory: Tensor<B,3>,visible: Option<Tensor<B,2,Bool>>,start_position: usize,
initial_capacity: usize,positions: F) -> EncoderDecoderKvCache<B>
where F: FnMut(usize,Tensor<B,4>,usize)->Tensor<B,4> {
EncoderDecoderKvCache::new(self.new_kv_cache(initial_capacity),self.prepare_cached_memory(memory,visible,start_position,positions))
}
pub fn forward_cached_pair(&self,input: Tensor<B,3>,new_visible: Option<Tensor<B,2,Bool>>,cache: &mut EncoderDecoderKvCache<B>,
self_masks: DenseAttentionMask<B>,self_options: DenseAttentionOptions,cross_masks: DenseAttentionMask<B>,
cross_options: DenseAttentionOptions) -> Tensor<B,3> {
let (decoder,memory) = cache.parts_mut();
self.forward_cached(input,new_visible,decoder,memory,self_masks,self_options,cross_masks,cross_options)
}
pub fn prepare_cached_memory<F>(&self,memory: Tensor<B,3>,visible: Option<Tensor<B,2,Bool>>,
start_position: usize,mut positions: F) -> Vec<ProjectedKvCache<B>>
where F: FnMut(usize,Tensor<B,4>,usize)->Tensor<B,4> {
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()
}
pub fn new_kv_cache(&self,initial_capacity: usize) -> TransformerKvCache<B> {
TransformerKvCache::new(self.layers.len(),initial_capacity)
}
pub fn forward_cached(&self,input: Tensor<B,3>,new_visible: Option<Tensor<B,2,Bool>>,cache: &mut TransformerKvCache<B>,
memory: &[ProjectedKvCache<B>],self_masks: DenseAttentionMask<B>,self_options: DenseAttentionOptions,
cross_masks: DenseAttentionMask<B>,cross_options: DenseAttentionOptions) -> Tensor<B,3> {
self.forward_cached_with(input,cache,memory,|_,layer,input,cache,memory|
layer.forward_cached(input,new_visible.clone(),cache,memory,self_masks.clone(),self_options,cross_masks.clone(),cross_options))
}
pub fn forward_cached_with<F>(&self,mut input: Tensor<B,3>,cache: &mut TransformerKvCache<B>,
memory: &[ProjectedKvCache<B>],mut layer: F) -> Tensor<B,3>
where F: FnMut(usize,&DenseEncoderDecoderLayer<B>,Tensor<B,3>,&mut ProjectedKvCache<B>,&ProjectedKvCache<B>)->Tensor<B,3> {
cache.validate_layers(self.layers.len());
assert_eq!(memory.len(),self.layers.len(),"actual cached encoder memory and decoder layer counts differ");
let shape = (input.dims()[0],input.dims()[1]);
let next = cache.position().checked_add(shape.1).expect("cached decoder position overflow");
for (index,block) in self.layers.iter().enumerate() {
input = layer(index,block,input,&mut cache.layers_mut()[index],&memory[index]);
assert_eq!((input.dims()[0],input.dims()[1]),shape,"cached decoder layer changed actual new rows");
}
cache.finish_chunk(next);
input
}
}
impl<B: Backend> AdaptedEncoderDecoderStack<B> {
pub fn prepare_kv_cache<F>(&self,memory: Tensor<B,3>,visible: Option<Tensor<B,2,Bool>>,start_position: usize,
initial_capacity: usize,positions: F) -> EncoderDecoderKvCache<B>
where F: FnMut(usize,Tensor<B,4>,usize)->Tensor<B,4> {
EncoderDecoderKvCache::new(self.new_kv_cache(initial_capacity),self.prepare_cached_memory(memory,visible,start_position,positions))
}
pub fn forward_cached_pair(&self,input: Tensor<B,3>,new_visible: Option<Tensor<B,2,Bool>>,cache: &mut EncoderDecoderKvCache<B>,
self_masks: DenseAttentionMask<B>,self_options: DenseAttentionOptions,cross_masks: DenseAttentionMask<B>,
cross_options: DenseAttentionOptions) -> Tensor<B,3> {
let (decoder,memory) = cache.parts_mut();
self.forward_cached(input,new_visible,decoder,memory,self_masks,self_options,cross_masks,cross_options)
}
pub fn prepare_cached_memory<F>(&self,memory: Tensor<B,3>,visible: Option<Tensor<B,2,Bool>>,
start_position: usize,mut positions: F) -> Vec<ProjectedKvCache<B>>
where F: FnMut(usize,Tensor<B,4>,usize)->Tensor<B,4> {
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()
}
pub fn new_kv_cache(&self,initial_capacity: usize) -> TransformerKvCache<B> {
TransformerKvCache::new(self.layers.len(),initial_capacity)
}
pub fn forward_cached(&self,input: Tensor<B,3>,new_visible: Option<Tensor<B,2,Bool>>,cache: &mut TransformerKvCache<B>,
memory: &[ProjectedKvCache<B>],self_masks: DenseAttentionMask<B>,self_options: DenseAttentionOptions,
cross_masks: DenseAttentionMask<B>,cross_options: DenseAttentionOptions) -> Tensor<B,3> {
self.forward_cached_with(input,cache,memory,|_,layer,input,cache,memory|
layer.forward_cached(input,new_visible.clone(),cache,memory,self_masks.clone(),self_options,cross_masks.clone(),cross_options))
}
pub fn forward_cached_with<F>(&self,mut input: Tensor<B,3>,cache: &mut TransformerKvCache<B>,
memory: &[ProjectedKvCache<B>],mut layer: F) -> Tensor<B,3>
where F: FnMut(usize,&AdaptedEncoderDecoderLayer<B>,Tensor<B,3>,&mut ProjectedKvCache<B>,&ProjectedKvCache<B>)->Tensor<B,3> {
cache.validate_layers(self.layers.len());
assert_eq!(memory.len(),self.layers.len(),"actual adapted memory cache and decoder layer counts differ");
let shape = (input.dims()[0],input.dims()[1]);
let next = cache.position().checked_add(shape.1).expect("cached adapted decoder position overflow");
for (index,block) in self.layers.iter().enumerate() {
input = layer(index,block,input,&mut cache.layers_mut()[index],&memory[index]);
assert_eq!((input.dims()[0],input.dims()[1]),shape,"cached adapted decoder changed actual new rows");
}
cache.finish_chunk(next);
input
}
}