use alloc::vec::Vec;
use ruda_autodiff::{Autodiff,checkpoint::strategy::CheckpointStrategy};
use ruda_model::{module::Module,tensor::{Bool,Tensor,backend::Backend}};
use crate::{Dropout,attention::{DenseAttentionMask,DenseAttentionOptions},cache::{ProjectedKvCache,TransformerKvCache,EncoderDecoderKvCache},
transformer::{DenseCrossAttentionBlock,DenseTransformerNorm,DenseEncoderDecoderLayer,DenseEncoderDecoderStack}};
use super::{AttentionParallelGroups,TensorParallelGroupedQueryAttention,TensorParallelTransformerBlock,BroadcastTensorCollective};
#[derive(Module,Debug)]
pub struct TensorParallelCrossAttentionBlock<B: Backend> {
pub attention: TensorParallelGroupedQueryAttention<B>,
pub query_norm: DenseTransformerNorm<B>,
pub memory_norm: Option<DenseTransformerNorm<B>>,
pub residual_dropout: Dropout,
pub norm_first: bool,
}
impl<B: Backend> TensorParallelCrossAttentionBlock<B> {
pub fn from_sharded_block(block: DenseCrossAttentionBlock<B>) -> Self {
assert_eq!(block.query_norm.width(),block.attention.query.weight.val().dims()[0],"parallel cross query norm width differs");
if let Some(norm) = &block.memory_norm {assert_eq!(norm.width(),block.attention.key.weight.val().dims()[0],"parallel cross memory norm width differs");}
Self {attention:TensorParallelGroupedQueryAttention::from_shard(block.attention),query_norm:block.query_norm,
memory_norm:block.memory_norm,residual_dropout:block.residual_dropout,norm_first:block.norm_first}
}
fn source(&self,input: Tensor<B,3>) -> Tensor<B,3> {
if self.norm_first {self.query_norm.forward(input)} else {input}
}
fn memory(&self,input: Tensor<B,3>) -> Tensor<B,3> {
if let Some(norm) = &self.memory_norm {norm.forward(input)} else {input}
}
fn finish(&self,input: Tensor<B,3>,branch: Tensor<B,3>) -> Tensor<B,3> {
let output = input+self.residual_dropout.forward(branch);
if self.norm_first {output} else {self.query_norm.forward(output)}
}
pub fn forward_inference<C,F>(&self,input: Tensor<B,3>,memory: Tensor<B,3>,masks: DenseAttentionMask<B>,options: DenseAttentionOptions,
communicator: C,positions: F) -> Result<Tensor<B,3>,C::Error>
where C: BroadcastTensorCollective<B>,F: FnOnce(Tensor<B,4>,Tensor<B,4>)->(Tensor<B,4>,Tensor<B,4>) {
let source = self.source(input.clone());let memory = self.memory(memory);
let (query,key,value) = self.attention.local.project(source,memory.clone(),memory);
let geometry = (query.dims(),key.dims());let (query,key) = positions(query,key);
assert_eq!((query.dims(),key.dims()),geometry,"parallel cross positions changed actual query/source heads");
Ok(self.finish(input,self.attention.forward_projected_inference(query,key,value,masks,options,communicator)?))
}
}
impl<B: Backend,S: CheckpointStrategy> TensorParallelCrossAttentionBlock<Autodiff<B,S>> {
pub fn forward<C,K,F>(&self,input: Tensor<Autodiff<B,S>,3>,memory: Tensor<Autodiff<B,S>,3>,masks: DenseAttentionMask<Autodiff<B,S>>,
options: DenseAttentionOptions,groups: &AttentionParallelGroups<C,K>,positions: F) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error>,
F: FnOnce(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>)->(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>) {
let query = self.attention.project_query(self.source(input.clone()),groups.heads.clone())?;
let (key,value) = self.attention.project_memory(self.memory(memory),groups)?;
let geometry = (query.dims(),key.dims());let (query,key) = positions(query,key);
assert_eq!((query.dims(),key.dims()),geometry,"parallel cross positions changed actual query/memory head geometry");
Ok(self.finish(input,self.attention.forward_projected(query,key,value,masks,options,groups.heads.clone())?))
}
pub fn prepare_memory<C,K,F>(&self,memory: Tensor<Autodiff<B,S>,3>,visible: Option<Tensor<Autodiff<B,S>,2,Bool>>,
groups: &AttentionParallelGroups<C,K>,positions: F) -> Result<ProjectedKvCache<Autodiff<B,S>>,C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error>,
F: FnOnce(Tensor<Autodiff<B,S>,4>)->Tensor<Autodiff<B,S>,4> {
let (key,value) = self.attention.project_memory(self.memory(memory),groups)?;
let geometry = key.dims();let key = positions(key);
assert_eq!(key.dims(),geometry,"parallel prepared source positions changed actual key geometry");
Ok(ProjectedKvCache::from_projected(key,value,visible,0))
}
pub fn forward_cached<C,F>(&self,input: Tensor<Autodiff<B,S>,3>,memory: &ProjectedKvCache<Autodiff<B,S>>,
masks: DenseAttentionMask<Autodiff<B,S>>,options: DenseAttentionOptions,communicator: C,position: usize,positions: F)
-> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C: BroadcastTensorCollective<B>,F: FnOnce(Tensor<Autodiff<B,S>,4>,usize)->Tensor<Autodiff<B,S>,4> {
let query = self.attention.project_query(self.source(input.clone()),communicator.clone())?;
let geometry = query.dims();let query = positions(query,position);
assert_eq!(query.dims(),geometry,"parallel cached cross positions changed actual query geometry");
Ok(self.finish(input,self.attention.forward_cached_memory(query,memory,masks,options,communicator)?))
}
}
#[derive(Module,Debug)]
pub struct TensorParallelEncoderDecoderLayer<B: Backend> {
pub backbone: TensorParallelTransformerBlock<B>,
pub cross_attention: TensorParallelCrossAttentionBlock<B>,
}
impl<B: Backend> TensorParallelEncoderDecoderLayer<B> {
pub fn from_sharded_layer(layer: DenseEncoderDecoderLayer<B>) -> Self {
assert_eq!(layer.backbone.attention.query.weight.val().dims()[0],layer.cross_attention.attention.query.weight.val().dims()[0],
"parallel decoder/cross residual widths differ");
Self {backbone:TensorParallelTransformerBlock::from_sharded_block(layer.backbone),
cross_attention:TensorParallelCrossAttentionBlock::from_sharded_block(layer.cross_attention)}
}
}
impl<B: Backend,S: CheckpointStrategy> TensorParallelEncoderDecoderLayer<Autodiff<B,S>> {
pub fn forward<C,K,L,F,G>(&self,input: Tensor<Autodiff<B,S>,3>,memory: Tensor<Autodiff<B,S>,3>,
self_masks: DenseAttentionMask<Autodiff<B,S>>,self_options: DenseAttentionOptions,
memory_masks: DenseAttentionMask<Autodiff<B,S>>,memory_options: DenseAttentionOptions,
self_groups: &AttentionParallelGroups<C,K>,memory_groups: &AttentionParallelGroups<C,L>,self_positions: F,cross_positions: G)
-> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error>,L: BroadcastTensorCollective<B,Error=C::Error>,
F: FnOnce(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>)->(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>),
G: FnOnce(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>)->(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>) {
let hidden = self.backbone.forward_attention_with_positions(input,self_masks,self_options,self_groups,self_positions)?;
let hidden = self.cross_attention.forward(hidden,memory,memory_masks,memory_options,memory_groups,cross_positions)?;
self.backbone.forward_feed_forward(hidden,self_groups.heads.clone())
}
pub fn forward_cached<C,K,F,G>(&self,input: Tensor<Autodiff<B,S>,3>,visible: Option<Tensor<Autodiff<B,S>,2,Bool>>,
decoder: &mut ProjectedKvCache<Autodiff<B,S>>,memory: &ProjectedKvCache<Autodiff<B,S>>,
self_masks: DenseAttentionMask<Autodiff<B,S>>,self_options: DenseAttentionOptions,
memory_masks: DenseAttentionMask<Autodiff<B,S>>,memory_options: DenseAttentionOptions,
groups: &AttentionParallelGroups<C,K>,cross_communicator: C,self_positions: F,cross_positions: G) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error>,
F: FnOnce(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>,usize)->(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>),
G: FnOnce(Tensor<Autodiff<B,S>,4>,usize)->Tensor<Autodiff<B,S>,4> {
let position = decoder.position();
let hidden = self.backbone.forward_cached_attention(input,visible,decoder,self_masks,self_options,groups,self_positions)?;
let hidden = self.cross_attention.forward_cached(hidden,memory,memory_masks,memory_options,cross_communicator,position,cross_positions)?;
self.backbone.forward_feed_forward(hidden,groups.heads.clone())
}
}
#[derive(Module,Debug)]
pub struct TensorParallelEncoderDecoderStack<B: Backend> {
pub layers: Vec<TensorParallelEncoderDecoderLayer<B>>,
}
impl<B: Backend> TensorParallelEncoderDecoderStack<B> {
pub fn new(layers: Vec<TensorParallelEncoderDecoderLayer<B>>) -> Self {Self {layers}}
pub fn from_sharded_stack(stack: DenseEncoderDecoderStack<B>) -> Self {
Self::new(stack.layers.into_iter().map(TensorParallelEncoderDecoderLayer::from_sharded_layer).collect())
}
pub fn new_decoder_cache(&self,initial_capacity: usize) -> TransformerKvCache<B> {TransformerKvCache::new(self.layers.len(),initial_capacity)}
}
impl<B: Backend,S: CheckpointStrategy> TensorParallelEncoderDecoderStack<Autodiff<B,S>> {
pub fn forward_with<E,F>(&self,mut input: Tensor<Autodiff<B,S>,3>,memory: Tensor<Autodiff<B,S>,3>,mut layer: F)
-> Result<Tensor<Autodiff<B,S>,3>,E>
where F: FnMut(usize,&TensorParallelEncoderDecoderLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,3>,Tensor<Autodiff<B,S>,3>)
-> Result<Tensor<Autodiff<B,S>,3>,E> {
for (index,block) in self.layers.iter().enumerate() {input = layer(index,block,input,memory.clone())?;}
Ok(input)
}
pub fn prepare_kv_cache_with<E,F>(&self,memory: Tensor<Autodiff<B,S>,3>,initial_capacity: usize,mut layer: F)
-> Result<EncoderDecoderKvCache<Autodiff<B,S>>,E>
where F: FnMut(usize,&TensorParallelCrossAttentionBlock<Autodiff<B,S>>,Tensor<Autodiff<B,S>,3>)
-> Result<ProjectedKvCache<Autodiff<B,S>>,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_with<E,F>(&self,mut input: Tensor<Autodiff<B,S>,3>,cache: &mut EncoderDecoderKvCache<Autodiff<B,S>>,mut layer: F)
-> Result<Tensor<Autodiff<B,S>,3>,E>
where F: FnMut(usize,&TensorParallelEncoderDecoderLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,3>,
&mut ProjectedKvCache<Autodiff<B,S>>,&ProjectedKvCache<Autodiff<B,S>>)->Result<Tensor<Autodiff<B,S>,3>,E> {
assert!(cache.is_consistent(),"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("parallel encoder-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,"parallel decoder changed actual target chunk rows");
}
decoder.finish_chunk(next);
Ok(input)
}
}