use super::{Backend,BroadcastTensorCollective,Tensor,Bool,DenseAttentionMask,DenseAttentionOptions,ProjectedKvCache,EncoderDecoderKvCache,Vec};
use super::{TensorParallelCrossAttentionBlock,TensorParallelEncoderDecoderLayer,TensorParallelEncoderDecoderStack};
impl<B: Backend> TensorParallelCrossAttentionBlock<B> {
pub fn prepare_memory_inference<F>(&self,memory: Tensor<B,3>,visible: Option<Tensor<B,2,Bool>>,positions: F) -> ProjectedKvCache<B>
where F: FnOnce(Tensor<B,4>)->Tensor<B,4> {
let (key,value) = self.attention.project_memory_inference(self.memory(memory));
let geometry = key.dims();let key = positions(key);
assert_eq!(key.dims(),geometry,"native parallel prepared memory positions changed actual keys");
ProjectedKvCache::from_projected(key,value,visible,0)
}
pub fn forward_cached_inference<C,F>(&self,input: Tensor<B,3>,memory: &ProjectedKvCache<B>,masks: DenseAttentionMask<B>,
options: DenseAttentionOptions,communicator: C,position: usize,positions: F) -> Result<Tensor<B,3>,C::Error>
where C: BroadcastTensorCollective<B>,F: FnOnce(Tensor<B,4>,usize)->Tensor<B,4> {
let query = self.attention.project_query_inference(self.source(input.clone()));
let geometry = query.dims();let query = positions(query,position);
assert_eq!(query.dims(),geometry,"native parallel cached cross positions changed actual query heads");
Ok(self.finish(input,self.attention.forward_cached_memory_inference(query,memory,masks,options,communicator)?))
}
}
impl<B: Backend> TensorParallelEncoderDecoderLayer<B> {
pub fn forward_inference<C,D,F,G>(&self,input: Tensor<B,3>,memory: Tensor<B,3>,self_masks: DenseAttentionMask<B>,self_options: DenseAttentionOptions,
memory_masks: DenseAttentionMask<B>,memory_options: DenseAttentionOptions,self_communicator: C,cross_communicator: D,
self_positions: F,cross_positions: G) -> Result<Tensor<B,3>,C::Error>
where C: BroadcastTensorCollective<B>,D: BroadcastTensorCollective<B,Error=C::Error>,
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_inference(input,self_masks,self_options,self_communicator.clone(),self_positions)?;
let hidden = self.cross_attention.forward_inference(hidden,memory,memory_masks,memory_options,cross_communicator,cross_positions)?;
self.backbone.forward_feed_forward_inference(hidden,self_communicator)
}
pub fn forward_cached_inference<C,D,F,G>(&self,input: Tensor<B,3>,visible: Option<Tensor<B,2,Bool>>,
decoder: &mut ProjectedKvCache<B>,memory: &ProjectedKvCache<B>,self_masks: DenseAttentionMask<B>,self_options: DenseAttentionOptions,
memory_masks: DenseAttentionMask<B>,memory_options: DenseAttentionOptions,self_communicator: C,cross_communicator: D,
self_positions: F,cross_positions: G) -> Result<Tensor<B,3>,C::Error>
where C: BroadcastTensorCollective<B>,D: BroadcastTensorCollective<B,Error=C::Error>,
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 = decoder.position();
let hidden = self.backbone.forward_cached_attention_inference(input,visible,decoder,self_masks,self_options,self_communicator.clone(),self_positions)?;
let hidden = self.cross_attention.forward_cached_inference(hidden,memory,memory_masks,memory_options,cross_communicator,position,cross_positions)?;
self.backbone.forward_feed_forward_inference(hidden,self_communicator)
}
}
impl<B: Backend> TensorParallelEncoderDecoderStack<B> {
pub fn forward_inference_with<E,F>(&self,mut input: Tensor<B,3>,memory: Tensor<B,3>,mut layer: F) -> Result<Tensor<B,3>,E>
where F: FnMut(usize,&TensorParallelEncoderDecoderLayer<B>,Tensor<B,3>,Tensor<B,3>)->Result<Tensor<B,3>,E> {
for (index,block) in self.layers.iter().enumerate() {input = layer(index,block,input,memory.clone())?;}
Ok(input)
}
pub fn prepare_kv_cache_inference_with<E,F>(&self,memory: Tensor<B,3>,initial_capacity: usize,mut layer: F) -> Result<EncoderDecoderKvCache<B>,E>
where F: FnMut(usize,&TensorParallelCrossAttentionBlock<B>,Tensor<B,3>)->Result<ProjectedKvCache<B>,E> {
let mut prepared = Vec::with_capacity(self.layers.len());
for (index,block) in self.layers.iter().enumerate() {prepared.push(layer(index,&block.cross_attention,memory.clone())?);}
Ok(EncoderDecoderKvCache::new(self.new_decoder_cache(initial_capacity),prepared))
}
pub fn forward_cached_inference_with<E,F>(&self,mut input: Tensor<B,3>,cache: &mut EncoderDecoderKvCache<B>,mut layer: F)
-> Result<Tensor<B,3>,E>
where F: FnMut(usize,&TensorParallelEncoderDecoderLayer<B>,Tensor<B,3>,&mut ProjectedKvCache<B>,&ProjectedKvCache<B>)->Result<Tensor<B,3>,E> {
assert!(cache.is_consistent(),"native parallel source/decoder cache pairing is inconsistent");
cache.decoder().validate_layers(self.layers.len());
let rows = (input.dims()[0],input.dims()[1]);let next = cache.position().checked_add(rows.1).expect("native parallel decoder position overflow");
let (decoder,memory) = cache.parts_mut();
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,"native parallel decoder changed actual chunk rows");
}
decoder.finish_chunk(next);
Ok(input)
}
}