use ruda_autodiff::{Autodiff,checkpoint::strategy::CheckpointStrategy};
use ruda_model::{module::Module,tensor::{Tensor,backend::Backend}};
use crate::{Dropout,attention::{DenseAttentionMask,DenseAttentionOptions},
transformer::{DenseTransformerBlock,DenseTransformerNorm}};
use super::{AttentionParallelGroups,TensorParallelGroupedQueryAttention,TensorParallelFeedForward,BroadcastTensorCollective};
mod cached;
mod packed;
mod stack;
pub use stack::*;
#[derive(Clone,Copy,Debug,PartialEq,Eq)]
pub enum TensorParallelResidualStage {
Attention,
FeedForward,
}
#[derive(Module,Debug)]
pub struct TensorParallelTransformerBlock<B: Backend> {
pub attention: TensorParallelGroupedQueryAttention<B>,
pub feed_forward: TensorParallelFeedForward<B>,
pub attention_norm: DenseTransformerNorm<B>,
pub feed_forward_norm: DenseTransformerNorm<B>,
pub residual_dropout: Dropout,
pub norm_first: bool,
}
pub(super) fn residual<B: Backend,E,F,R,const D: usize>(input: Tensor<B,D>,norm: &DenseTransformerNorm<B>,norm_first: bool,
branch: F,dropout: R) -> Result<Tensor<B,D>,E>
where F: FnOnce(Tensor<B,D>)->Result<Tensor<B,D>,E>,R: FnOnce(Tensor<B,D>)->Tensor<B,D> {
let source = if norm_first {norm.forward(input.clone())} else {input.clone()};
let output = input+dropout(branch(source)?);
Ok(if norm_first {output} else {norm.forward(output)})
}
impl<B: Backend> TensorParallelTransformerBlock<B> {
pub fn from_sharded_block(block: DenseTransformerBlock<B>) -> Self {
let width = block.attention.query.weight.val().dims()[0];
assert_eq!(block.attention.key.weight.val().dims()[0],width,"self-attention memory/residual width differs");
assert_eq!(block.feed_forward.up.weight.val().dims()[0],width,"parallel FFN/residual width differs");
assert_eq!(block.attention_norm.width(),width,"parallel attention norm/residual width differs");
assert_eq!(block.feed_forward_norm.width(),width,"parallel FFN norm/residual width differs");
Self {attention:TensorParallelGroupedQueryAttention::from_shard(block.attention),
feed_forward:TensorParallelFeedForward::from_shard(block.feed_forward),attention_norm:block.attention_norm,
feed_forward_norm:block.feed_forward_norm,residual_dropout:block.residual_dropout,norm_first:block.norm_first}
}
pub fn into_local_block(self) -> DenseTransformerBlock<B> {
DenseTransformerBlock {attention:self.attention.local,feed_forward:self.feed_forward.local,
attention_norm:self.attention_norm,feed_forward_norm:self.feed_forward_norm,residual_dropout:self.residual_dropout,norm_first:self.norm_first}
}
pub fn forward_attention_inference<C,F>(&self,input: 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>) {
residual(input,&self.attention_norm,self.norm_first,|source| {
let (query,key,value) = self.attention.local.project(source.clone(),source.clone(),source);
let geometry = (query.dims(),key.dims());let (query,key) = positions(query,key);
assert_eq!((query.dims(),key.dims()),geometry,"native parallel attention positions changed local geometry");
self.attention.forward_projected_inference(query,key,value,masks,options,communicator)
},|branch|self.residual_dropout.forward(branch))
}
pub fn forward_feed_forward_inference<C: BroadcastTensorCollective<B>>(&self,input: Tensor<B,3>,communicator: C) -> Result<Tensor<B,3>,C::Error> {
residual(input,&self.feed_forward_norm,self.norm_first,|source|self.feed_forward.forward_inference(source,communicator),
|branch|self.residual_dropout.forward(branch))
}
pub fn forward_inference<C,F>(&self,input: 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 hidden = residual(input,&self.attention_norm,self.norm_first,|source| {
let (query,key,value) = self.attention.local.project(source.clone(),source.clone(),source);
let geometry = (query.dims(),key.dims());
let (query,key) = positions(query,key);
assert_eq!((query.dims(),key.dims()),geometry,"parallel inference positions changed actual head geometry");
self.attention.forward_projected_inference(query,key,value,masks,options,communicator.clone())
},|branch|self.residual_dropout.forward(branch))?;
residual(hidden,&self.feed_forward_norm,self.norm_first,|source|self.feed_forward.forward_inference(source,communicator),
|branch|self.residual_dropout.forward(branch))
}
}
impl<B: Backend,S: CheckpointStrategy> TensorParallelTransformerBlock<Autodiff<B,S>> {
pub fn forward<C,K>(&self,input: Tensor<Autodiff<B,S>,3>,masks: DenseAttentionMask<Autodiff<B,S>>,options: DenseAttentionOptions,
groups: &AttentionParallelGroups<C,K>) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error> {
self.forward_with_positions(input,masks,options,groups,|query,key|(query,key))
}
pub fn forward_with_positions<C,K,F>(&self,input: 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>) {
self.forward_with_residual(input,masks,options,groups,positions,|_,branch|self.residual_dropout.forward(branch))
}
pub fn forward_with_residual<C,K,F,R>(&self,input: Tensor<Autodiff<B,S>,3>,masks: DenseAttentionMask<Autodiff<B,S>>,options: DenseAttentionOptions,
groups: &AttentionParallelGroups<C,K>,positions: F,mut branch_output: R) -> 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>),
R: FnMut(TensorParallelResidualStage,Tensor<Autodiff<B,S>,3>)->Tensor<Autodiff<B,S>,3> {
let hidden = residual(input,&self.attention_norm,self.norm_first,|source| {
let (query,key,value) = self.attention.project_self(source,groups)?;
let geometry = (query.dims(),key.dims());
let (query,key) = positions(query,key);
assert_eq!((query.dims(),key.dims()),geometry,"parallel transformer positions changed local head geometry");
self.attention.forward_projected(query,key,value,masks,options,groups.heads.clone())
},|branch|branch_output(TensorParallelResidualStage::Attention,branch))?;
residual(hidden,&self.feed_forward_norm,self.norm_first,|source|self.feed_forward.forward(source,groups.heads.clone()),
|branch|branch_output(TensorParallelResidualStage::FeedForward,branch))
}
pub fn forward_attention_with_positions<C,K,F>(&self,input: 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>) {
residual(input,&self.attention_norm,self.norm_first,|source| {
let (query,key,value) = self.attention.project_self(source,groups)?;
let geometry = (query.dims(),key.dims());let (query,key) = positions(query,key);
assert_eq!((query.dims(),key.dims()),geometry,"parallel self-attention positions changed local geometry");
self.attention.forward_projected(query,key,value,masks,options,groups.heads.clone())
},|branch|self.residual_dropout.forward(branch))
}
pub fn forward_feed_forward<C: BroadcastTensorCollective<B>>(&self,input: Tensor<Autodiff<B,S>,3>,communicator: C)
-> Result<Tensor<Autodiff<B,S>,3>,C::Error> {
residual(input,&self.feed_forward_norm,self.norm_first,|source|self.feed_forward.forward(source,communicator),
|branch|self.residual_dropout.forward(branch))
}
}