use alloc::{vec::Vec,collections::BTreeSet};
use ruda_model::{module::{Module,ModuleVisitor,Param,ParamId},tensor::{Tensor,Int,Bool,MoeDispatchOps,MoeReceivedOps,VariableTensorCollective,backend::Backend}};
use ruda_autodiff::{Autodiff,checkpoint::strategy::CheckpointStrategy,collective::{CollectiveScope,ScopedTensorCollective,ScopedCollectiveError}};
use crate::{Dropout,expert_parallel::{ExpertParallelMoeLayer,ExpertParallelMoeError,ExpertParallelSwiGluExperts,ExpertParallelGeometry,ExpertParallelReceived},attention::{DenseAttentionMask,DenseAttentionOptions,
PackedSequenceLayout,PackedAttentionOptions,PackedDocumentAttentionMask},cache::{ProjectedKvCache,TransformerKvCache},
loss::CausalCrossEntropyConfig,fully_sharded::{FullyShardedLoss,complete_fully_sharded_loss}};
use super::{ProjectedGroupedQueryAttention,ProjectedFeedForward,DenseTransformerNorm,NativeMoeTransformerLayer,NativeMoeTransformerError,
TransformerProjectionShape,TransformerProjection,TransformerEmbeddings,ProjectedTransformerHead,ProjectedTransformerInput};
use super::{dense::try_residual_branch,native_attention::{attention_branch,packed_attention_branch,cached_attention_branch},
projected_paired_model::{embed_projected,embed_packed_projected,check_block}};
#[derive(Module,Debug)]
pub struct ExpertParallelTransformerBlock<B:Backend,P:Module<B>,E:Module<B> =ExpertParallelSwiGluExperts<B>> {
pub attention:ProjectedGroupedQueryAttention<B,P>,
pub routed:ExpertParallelMoeLayer<B,P,E>,
pub shared:Option<ProjectedFeedForward<B,P>>,
pub attention_norm:DenseTransformerNorm<B>,
pub feed_forward_norm:DenseTransformerNorm<B>,
pub residual_dropout:Dropout,
pub norm_first:bool,
}
impl<B:Backend,P:TransformerProjectionShape<B>,E:ExpertParallelGeometry<B>> ExpertParallelTransformerBlock<B,P,E> {
pub fn validate(&self) {
self.routed.validate();let width=self.routed.width();
for projection in [&self.attention.query,&self.attention.key,&self.attention.value] {assert_eq!(projection.dimensions()[0],width,"owned-expert attention input width differs");}
assert_eq!(self.attention.output.dimensions()[1],width,"owned-expert attention output width differs");
assert_eq!((self.attention_norm.width(),self.feed_forward_norm.width()),(width,width),"owned-expert original norm widths differ");
if let Some(shared)=&self.shared {assert_eq!((shared.up.dimensions()[0],shared.down.dimensions()[1]),(width,width),"owned-expert shared residual width differs");}
}
}
#[derive(Debug)]
pub enum ExpertParallelTransformerError<C:core::fmt::Debug,P:core::fmt::Debug,M:core::fmt::Debug> {
Local(NativeMoeTransformerError<P,M>),
Expert(ExpertParallelMoeError<C,P,M>),
}
impl<C:core::fmt::Debug,P:core::fmt::Debug,M:core::fmt::Debug> core::fmt::Display for ExpertParallelTransformerError<C,P,M> {
fn fmt(&self,f:&mut core::fmt::Formatter<'_>) -> core::fmt::Result {match self {Self::Local(error)=>write!(f,"native model stage: {error}"),Self::Expert(error)=>write!(f,"owned expert stage: {error}")}}
}
impl<C:core::fmt::Debug,P:core::fmt::Debug,M:core::fmt::Debug> core::error::Error for ExpertParallelTransformerError<C,P,M> {}
#[derive(Module,Debug)]
pub enum ExpertParallelTransformerLayer<B:Backend,P:Module<B>,E:Module<B> =ExpertParallelSwiGluExperts<B>> {
Local(NativeMoeTransformerLayer<B,P>),
Parallel(ExpertParallelTransformerBlock<B,P,E>),
}
#[derive(Module,Debug)]
pub struct ExpertParallelTransformerModel<B:Backend,P:Module<B>,E:Module<B> =ExpertParallelSwiGluExperts<B>> {
pub embeddings:TransformerEmbeddings<B>,
pub layers:Vec<ExpertParallelTransformerLayer<B,P,E>>,
pub normalization:Option<DenseTransformerNorm<B>>,
pub head:ProjectedTransformerHead<B,P>,
}
struct FloatIds(BTreeSet<ParamId>);
impl<B:Backend> ModuleVisitor<B> for FloatIds {
fn visit_float<const D:usize>(&mut self,parameter:&Param<Tensor<B,D>>) {self.0.insert(parameter.id);}
}
impl<B:Backend,P:TransformerProjectionShape<B>,E:ExpertParallelGeometry<B>> ExpertParallelTransformerModel<B,P,E> {
pub fn from_expert_parts(embeddings:TransformerEmbeddings<B>,layers:Vec<ExpertParallelTransformerLayer<B,P,E>>,normalization:Option<DenseTransformerNorm<B>>,head:ProjectedTransformerHead<B,P>) -> Self {
let width=embeddings.token.weight.val().dims()[1];assert_eq!(head.projection.dimensions()[0],width,"expert model original head/input width differs");
if let Some(norm)=&normalization {assert_eq!(norm.width(),width,"expert model final norm width differs");}
for layer in &layers {match layer {
ExpertParallelTransformerLayer::Local(layer)=>{assert_eq!(layer.width(),width,"local model residual width differs");match layer {
NativeMoeTransformerLayer::Dense(block)=>check_block(block,width),NativeMoeTransformerLayer::Routed(block)=>block.validate()}},
ExpertParallelTransformerLayer::Parallel(block)=>{block.validate();assert_eq!(block.routed.width(),width,"owned-expert residual width differs");},
}}Self {embeddings,layers,normalization,head}
}
pub fn new_kv_cache(&self,capacity:usize) -> TransformerKvCache<B> {TransformerKvCache::new(self.layers.len(),capacity)}
pub fn owned_expert_parameter_ids(&self) -> Vec<ParamId> {
let mut ids=BTreeSet::new();for layer in &self.layers {if let ExpertParallelTransformerLayer::Parallel(block)=layer {
ids.extend(block.routed.experts.parameter_ids());}}ids.into_iter().collect()
}
pub fn non_expert_parameter_ids(&self) -> Vec<ParamId> {
let owned=self.owned_expert_parameter_ids().into_iter().collect::<BTreeSet<_>>();let mut ids=FloatIds(BTreeSet::new());self.visit(&mut ids);
ids.0.difference(&owned).copied().collect()
}
fn normalize<const D:usize>(&self,hidden:Tensor<B,D>) -> Tensor<B,D> {if let Some(norm)=&self.normalization {norm.forward(hidden)} else {hidden}}
}
impl<B:Backend,P:TransformerProjectionShape<B>> ExpertParallelTransformerModel<B,P> {
pub fn from_parts(embeddings:TransformerEmbeddings<B>,layers:Vec<ExpertParallelTransformerLayer<B,P>>,normalization:Option<DenseTransformerNorm<B>>,head:ProjectedTransformerHead<B,P>) -> Self {
Self::from_expert_parts(embeddings,layers,normalization,head)
}
}
macro_rules! expert_model_execution {
($backend:ty,[$($generics:tt)*],$routed:ident,$feed:ident,$forward:ident,$packed:ident,$cached:ident,$hidden_with:ident,$packed_hidden_with:ident) => {
impl<$($generics)*,P:TransformerProjection<$backend>,E:ExpertParallelReceived<$backend>> ExpertParallelTransformerBlock<$backend,P,E> {
fn $feed<C:VariableTensorCollective<B>,const D:usize>(&self,hidden:Tensor<$backend,D>,communicator:C)
-> Result<Tensor<$backend,D>,ExpertParallelTransformerError<C::Error,P::Error,<$backend as ruda_model::tensor::MoeOps>::MoeError>> {
try_residual_branch(hidden,&self.feed_forward_norm,&self.residual_dropout,self.norm_first,|source| {
let output=self.routed.$routed(source.clone(),communicator).map_err(ExpertParallelTransformerError::Expert)?;
if let Some(shared)=&self.shared {Ok(output+shared.forward(source).map_err(|error|ExpertParallelTransformerError::Local(NativeMoeTransformerError::Projection(error)))?)} else {Ok(output)}
})
}
pub fn $forward<C:VariableTensorCollective<B>,F>(&self,input:Tensor<$backend,3>,masks:DenseAttentionMask<$backend>,options:DenseAttentionOptions,communicator:C,positions:F)
-> Result<Tensor<$backend,3>,ExpertParallelTransformerError<C::Error,P::Error,<$backend as ruda_model::tensor::MoeOps>::MoeError>>
where F:FnOnce(Tensor<$backend,4>,Tensor<$backend,4>)->(Tensor<$backend,4>,Tensor<$backend,4>) {
self.validate();let hidden=attention_branch(&self.attention,&self.attention_norm,&self.residual_dropout,self.norm_first,input,masks,options,positions)
.map_err(|error|ExpertParallelTransformerError::Local(NativeMoeTransformerError::Projection(error)))?;self.$feed(hidden,communicator)
}
pub fn $packed<C:VariableTensorCollective<B>,F>(&self,input:Tensor<$backend,2>,layout:&PackedSequenceLayout,masks:Option<&[PackedDocumentAttentionMask<$backend>]>,
options:PackedAttentionOptions,communicator:C,positions:F) -> Result<Tensor<$backend,2>,ExpertParallelTransformerError<C::Error,P::Error,<$backend as ruda_model::tensor::MoeOps>::MoeError>>
where F:FnOnce(Tensor<$backend,3>,Tensor<$backend,3>)->(Tensor<$backend,3>,Tensor<$backend,3>) {
self.validate();let hidden=packed_attention_branch(&self.attention,&self.attention_norm,&self.residual_dropout,self.norm_first,input,layout,masks,options,positions)
.map_err(|error|ExpertParallelTransformerError::Local(NativeMoeTransformerError::Projection(error)))?;self.$feed(hidden,communicator)
}
pub fn $cached<C:VariableTensorCollective<B>,F>(&self,input:Tensor<$backend,3>,visible:Option<Tensor<$backend,2,Bool>>,cache:&mut ProjectedKvCache<$backend>,
masks:DenseAttentionMask<$backend>,options:DenseAttentionOptions,communicator:C,positions:F)
-> Result<Tensor<$backend,3>,ExpertParallelTransformerError<C::Error,P::Error,<$backend as ruda_model::tensor::MoeOps>::MoeError>>
where F:FnOnce(Tensor<$backend,4>,Tensor<$backend,4>,usize)->(Tensor<$backend,4>,Tensor<$backend,4>) {
self.validate();let hidden=cached_attention_branch(&self.attention,&self.attention_norm,&self.residual_dropout,self.norm_first,input,visible,cache,masks,options,positions)
.map_err(|error|ExpertParallelTransformerError::Local(NativeMoeTransformerError::Projection(error)))?;self.$feed(hidden,communicator)
}
}
impl<$($generics)*,P:TransformerProjection<$backend>,E:ExpertParallelReceived<$backend>> ExpertParallelTransformerLayer<$backend,P,E> {
pub fn $cached<C:VariableTensorCollective<B>,F>(&self,input:Tensor<$backend,3>,visible:Option<Tensor<$backend,2,Bool>>,cache:&mut ProjectedKvCache<$backend>,
masks:DenseAttentionMask<$backend>,options:DenseAttentionOptions,communicator:C,positions:F)
-> Result<Tensor<$backend,3>,ExpertParallelTransformerError<C::Error,P::Error,<$backend as ruda_model::tensor::MoeOps>::MoeError>>
where F:FnOnce(Tensor<$backend,4>,Tensor<$backend,4>,usize)->(Tensor<$backend,4>,Tensor<$backend,4>) {
match self {Self::Local(layer)=>layer.forward_cached_with_positions(input,visible,cache,masks,options,positions).map_err(ExpertParallelTransformerError::Local),
Self::Parallel(block)=>block.$cached(input,visible,cache,masks,options,communicator,positions)}
}
pub fn $forward<C:VariableTensorCollective<B>,F>(&self,input:Tensor<$backend,3>,masks:DenseAttentionMask<$backend>,options:DenseAttentionOptions,communicator:C,positions:F)
-> Result<Tensor<$backend,3>,ExpertParallelTransformerError<C::Error,P::Error,<$backend as ruda_model::tensor::MoeOps>::MoeError>>
where F:FnOnce(Tensor<$backend,4>,Tensor<$backend,4>)->(Tensor<$backend,4>,Tensor<$backend,4>) {
match self {Self::Local(layer)=>layer.forward_with_positions(input,masks,options,positions).map_err(ExpertParallelTransformerError::Local),
Self::Parallel(block)=>block.$forward(input,masks,options,communicator,positions)}
}
pub fn $packed<C:VariableTensorCollective<B>,F>(&self,input:Tensor<$backend,2>,layout:&PackedSequenceLayout,masks:Option<&[PackedDocumentAttentionMask<$backend>]>,
options:PackedAttentionOptions,communicator:C,positions:F) -> Result<Tensor<$backend,2>,ExpertParallelTransformerError<C::Error,P::Error,<$backend as ruda_model::tensor::MoeOps>::MoeError>>
where F:FnOnce(Tensor<$backend,3>,Tensor<$backend,3>)->(Tensor<$backend,3>,Tensor<$backend,3>) {
match self {Self::Local(layer)=>layer.forward_packed_with_positions(input,layout,masks,options,positions).map_err(ExpertParallelTransformerError::Local),
Self::Parallel(block)=>block.$packed(input,layout,masks,options,communicator,positions)}
}
}
impl<$($generics)*,P:TransformerProjection<$backend>,E:ExpertParallelReceived<$backend>> ExpertParallelTransformerModel<$backend,P,E> {
pub fn $cached<C:VariableTensorCollective<B>,F>(&self,input:ProjectedTransformerInput<$backend>,visible:Option<Tensor<$backend,2,Bool>>,cache:&mut TransformerKvCache<$backend>,
masks:DenseAttentionMask<$backend>,options:DenseAttentionOptions,communicator:C,mut positions:F)
-> Result<Tensor<$backend,3>,ExpertParallelTransformerError<C::Error,P::Error,<$backend as ruda_model::tensor::MoeOps>::MoeError>>
where F:FnMut(usize,Tensor<$backend,4>,Tensor<$backend,4>,usize)->(Tensor<$backend,4>,Tensor<$backend,4>) {
cache.validate_layers(self.layers.len());let rows=input.tokens.dims();let next=cache.position().checked_add(rows[1]).expect("expert model cache position overflows");
let mut hidden=embed_projected(&self.embeddings,input);for (index,layer) in self.layers.iter().enumerate() {
hidden=layer.$cached(hidden,visible.clone(),&mut cache.layers_mut()[index],masks.clone(),options,communicator.clone(),|query,key,position|positions(index,query,key,position))?;
assert_eq!((hidden.dims()[0],hidden.dims()[1]),(rows[0],rows[1]),"expert model cached layer changed actual token rows");}
cache.finish_chunk(next);self.head.forward(self.normalize(hidden)).map_err(|error|ExpertParallelTransformerError::Local(NativeMoeTransformerError::Projection(error)))
}
pub fn $hidden_with<C:VariableTensorCollective<B>,F>(&self,input:ProjectedTransformerInput<$backend>,communicator:C,mut layer:F)
-> Result<Tensor<$backend,3>,ExpertParallelTransformerError<C::Error,P::Error,<$backend as ruda_model::tensor::MoeOps>::MoeError>>
where F:FnMut(usize,&ExpertParallelTransformerLayer<$backend,P,E>,Tensor<$backend,3>,C)
-> Result<Tensor<$backend,3>,ExpertParallelTransformerError<C::Error,P::Error,<$backend as ruda_model::tensor::MoeOps>::MoeError>> {
let rows=input.tokens.dims();let width=self.embeddings.token.weight.val().dims()[1];let mut hidden=embed_projected(&self.embeddings,input);
for (index,block) in self.layers.iter().enumerate() {hidden=layer(index,block,hidden,communicator.clone())?;
assert_eq!(hidden.dims(),[rows[0],rows[1],width],"expert model layer changed actual source token axes");}Ok(self.normalize(hidden))
}
pub fn $packed_hidden_with<C:VariableTensorCollective<B>,F>(&self,input:ProjectedTransformerInput<$backend,1>,layout:&PackedSequenceLayout,communicator:C,mut layer:F)
-> Result<Tensor<$backend,2>,ExpertParallelTransformerError<C::Error,P::Error,<$backend as ruda_model::tensor::MoeOps>::MoeError>>
where F:FnMut(usize,&ExpertParallelTransformerLayer<$backend,P,E>,Tensor<$backend,2>,C)
-> Result<Tensor<$backend,2>,ExpertParallelTransformerError<C::Error,P::Error,<$backend as ruda_model::tensor::MoeOps>::MoeError>> {
let shape=[layout.tokens(),self.embeddings.token.weight.val().dims()[1]];let mut hidden=embed_packed_projected(&self.embeddings,input,layout);
for (index,block) in self.layers.iter().enumerate() {hidden=layer(index,block,hidden,communicator.clone())?;
assert_eq!(hidden.dims(),shape,"expert model packed layer changed original document rows");}Ok(self.normalize(hidden))
}
pub fn $forward<C:VariableTensorCollective<B>,F>(&self,input:ProjectedTransformerInput<$backend>,masks:DenseAttentionMask<$backend>,options:DenseAttentionOptions,communicator:C,mut positions:F)
-> Result<Tensor<$backend,3>,ExpertParallelTransformerError<C::Error,P::Error,<$backend as ruda_model::tensor::MoeOps>::MoeError>>
where F:FnMut(usize,Tensor<$backend,4>,Tensor<$backend,4>)->(Tensor<$backend,4>,Tensor<$backend,4>) {
let hidden=self.$hidden_with(input,communicator,|index,layer,hidden,transport|layer.$forward(hidden,masks.clone(),options,transport,|query,key|positions(index,query,key)))?;
self.head.forward(hidden).map_err(|error|ExpertParallelTransformerError::Local(NativeMoeTransformerError::Projection(error)))
}
pub fn $packed<C:VariableTensorCollective<B>,F>(&self,input:ProjectedTransformerInput<$backend,1>,layout:&PackedSequenceLayout,masks:Option<&[PackedDocumentAttentionMask<$backend>]>,
options:PackedAttentionOptions,communicator:C,mut positions:F) -> Result<Tensor<$backend,2>,ExpertParallelTransformerError<C::Error,P::Error,<$backend as ruda_model::tensor::MoeOps>::MoeError>>
where F:FnMut(usize,Tensor<$backend,3>,Tensor<$backend,3>)->(Tensor<$backend,3>,Tensor<$backend,3>) {
let hidden=self.$packed_hidden_with(input,layout,communicator,|index,layer,hidden,transport|layer.$packed(hidden,layout,masks,options,transport,|query,key|positions(index,query,key)))?;
self.head.forward(hidden).map_err(|error|ExpertParallelTransformerError::Local(NativeMoeTransformerError::Projection(error)))
}
}
};
}
expert_model_execution!(B,[B:MoeDispatchOps+MoeReceivedOps],forward_inference,feed_forward_inference,forward_with_positions_inference,forward_packed_with_positions_inference,forward_cached_with_positions_inference,forward_hidden_with_inference,forward_packed_hidden_with_inference);
expert_model_execution!(Autodiff<B,S>,[B:MoeDispatchOps+MoeReceivedOps,S:CheckpointStrategy],forward,feed_forward,forward_with_positions,forward_packed_with_positions,forward_cached_with_positions,forward_hidden_with,forward_packed_hidden_with);
pub type ExpertParallelLoss<B,S> = FullyShardedLoss<B,S>;
#[derive(Debug)]
pub enum ExpertParallelTrainingError<C:core::fmt::Debug,P:core::fmt::Debug,M:core::fmt::Debug> {
Model(ExpertParallelTransformerError<C,P,M>),
Loss(ScopedCollectiveError<C>),
}
impl<C:core::fmt::Debug,P:core::fmt::Debug,M:core::fmt::Debug> core::fmt::Display for ExpertParallelTrainingError<C,P,M> {
fn fmt(&self,f:&mut core::fmt::Formatter<'_>) -> core::fmt::Result {match self {Self::Model(error)=>write!(f,"expert model: {error}"),Self::Loss(error)=>write!(f,"expert loss: {error}")}}
}
impl<C:core::fmt::Debug,P:core::fmt::Debug,M:core::fmt::Debug> core::error::Error for ExpertParallelTrainingError<C,P,M> {}
impl<B:MoeDispatchOps+MoeReceivedOps,S:CheckpointStrategy,P:TransformerProjection<Autodiff<B,S>>,E:ExpertParallelReceived<Autodiff<B,S>>>
ExpertParallelTransformerModel<Autodiff<B,S>,P,E> {
pub fn forward_causal_with<C:VariableTensorCollective<B>,F>(&self,input:ProjectedTransformerInput<Autodiff<B,S>>,labels:Tensor<Autodiff<B,S>,2,Int>,
criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C,layer:F)
-> Result<ExpertParallelLoss<B,S>,ExpertParallelTrainingError<C::Error,P::Error,<Autodiff<B,S> as ruda_model::tensor::MoeOps>::MoeError>>
where F:FnMut(usize,&ExpertParallelTransformerLayer<Autodiff<B,S>,P,E>,Tensor<Autodiff<B,S>,3>,ScopedTensorCollective<C,B,S>)
-> Result<Tensor<Autodiff<B,S>,3>,ExpertParallelTransformerError<C::Error,P::Error,<Autodiff<B,S> as ruda_model::tensor::MoeOps>::MoeError>> {
let scope=CollectiveScope::<B,S>::new();let transport=scope.bind(communicator.clone());
let hidden=self.forward_hidden_with(input,transport,layer).map_err(ExpertParallelTrainingError::Model)?;
let loss=criterion.try_forward_hidden_with_smoothing(hidden,labels,|rows|self.head.forward(rows),label_smoothing)
.map_err(|error|ExpertParallelTrainingError::Model(ExpertParallelTransformerError::Local(NativeMoeTransformerError::Projection(error))))?;
complete_fully_sharded_loss(&scope,loss.loss_sum,loss.valid_tokens,communicator).map_err(ExpertParallelTrainingError::Loss)
}
pub fn forward_packed_causal_with<C:VariableTensorCollective<B>,F>(&self,input:ProjectedTransformerInput<Autodiff<B,S>,1>,labels:Tensor<Autodiff<B,S>,1,Int>,
layout:&PackedSequenceLayout,criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C,layer:F)
-> Result<ExpertParallelLoss<B,S>,ExpertParallelTrainingError<C::Error,P::Error,<Autodiff<B,S> as ruda_model::tensor::MoeOps>::MoeError>>
where F:FnMut(usize,&ExpertParallelTransformerLayer<Autodiff<B,S>,P,E>,Tensor<Autodiff<B,S>,2>,ScopedTensorCollective<C,B,S>)
-> Result<Tensor<Autodiff<B,S>,2>,ExpertParallelTransformerError<C::Error,P::Error,<Autodiff<B,S> as ruda_model::tensor::MoeOps>::MoeError>> {
let scope=CollectiveScope::<B,S>::new();let transport=scope.bind(communicator.clone());
let hidden=self.forward_packed_hidden_with(input,layout,transport,layer).map_err(ExpertParallelTrainingError::Model)?;
let loss=criterion.try_forward_packed_hidden_with_smoothing(hidden,labels,layout,|rows|self.head.forward(rows),label_smoothing)
.map_err(|error|ExpertParallelTrainingError::Model(ExpertParallelTransformerError::Local(NativeMoeTransformerError::Projection(error))))?;
complete_fully_sharded_loss(&scope,loss.loss_sum,loss.valid_tokens,communicator).map_err(ExpertParallelTrainingError::Loss)
}
}