use super::*;
use crate::{NativeSwiGluExperts,NativeMoeLayer,NativeMoeLayerOutput,NativeMoeLayerError};
use crate::transformer::TransformerProjection;
use ruda_model::{module::ModuleDisplay,tensor::{MoeOps,MoeOptions,IntegerTensorCollective}};
#[derive(Module,Debug)]
pub struct FullyShardedNativeSwiGluExperts<B:Backend> {
pub gate:ShardedParameter<B>,
pub up:ShardedParameter<B>,
pub down:ShardedParameter<B>,
}
impl<B:Backend> FullyShardedNativeSwiGluExperts<B> {
pub fn from_full(experts:NativeSwiGluExperts<B>,rank:usize,world:usize) -> Self {ShardingContext::new(rank,world).moe_experts(experts)}
pub fn dimensions(&self) -> [usize;3] {
assert_eq!(self.gate.logical_shape.len(),3,"native expert logical rank differs");let shape=&self.gate.logical_shape;[shape[0],shape[2],shape[1]]
}
pub fn validate(&self) {
let [experts,hidden,inner]=self.dimensions();assert!(experts>0 && hidden>0 && inner>0,"native expert axes must be positive");
assert_eq!(self.up.logical_shape,self.gate.logical_shape,"native sharded gate/up geometry differs");
assert_eq!(self.down.logical_shape,[experts,hidden,inner],"native sharded down geometry differs");
let source=self.gate.local.val();assert!(matches!(source.dtype(),DType::F16|DType::BF16|DType::F32),"unsupported native expert storage");
for parameter in [&self.gate,&self.up,&self.down] {
let value=parameter.local.val();assert_eq!(value.dtype(),source.dtype(),"native expert storage differs");assert_eq!(value.device(),source.device(),"native expert devices differ");
assert_eq!((parameter.rank,parameter.world_size),(self.gate.rank,self.gate.world_size),"native expert topology differs");
let _=ShardedParameter::from_local(parameter.local.clone(),parameter.logical_shape.clone(),parameter.rank,parameter.world_size);
}
}
}
#[derive(Module,Debug)]
pub struct FullyShardedNativeMoeLayer<B:Backend,P:Module<B>> {
pub router:P,
pub experts:FullyShardedNativeSwiGluExperts<B>,
pub correction_bias:Option<ShardedParameter<B>>,
#[module(skip)]
pub options:MoeOptions,
#[module(skip)]
pub router_input_dtype:Option<FloatDType>,
}
impl<B:Backend> ShardingContext<B> {
pub fn moe_experts(&mut self,experts:NativeSwiGluExperts<B>) -> FullyShardedNativeSwiGluExperts<B> {
experts.validate();let experts=FullyShardedNativeSwiGluExperts {gate:self.parameter(experts.gate),up:self.parameter(experts.up),down:self.parameter(experts.down)};
experts.validate();experts
}
pub fn moe_layer<P:ShardTransformerProjection<B>>(&mut self,layer:NativeMoeLayer<B,P>) -> FullyShardedNativeMoeLayer<B,P::Sharded> {
layer.validate();FullyShardedNativeMoeLayer {router:self.awq_projection(layer.router),experts:self.moe_experts(layer.experts),
correction_bias:layer.correction_bias.map(|bias|self.parameter(bias)),options:layer.options,router_input_dtype:layer.router_input_dtype}
}
}
impl<B:Backend,P:Module<B>> FullyShardedNativeMoeLayer<B,P> {
pub fn from_full<Q:ShardTransformerProjection<B,Sharded=P>>(layer:NativeMoeLayer<B,Q>,rank:usize,world:usize) -> Self {ShardingContext::new(rank,world).moe_layer(layer)}
}
impl<B:Backend,P:FullyShardedModule<B>+ModuleDisplay> FullyShardedNativeMoeLayer<B,P> {
pub fn width(&self) -> usize {self.experts.dimensions()[1]}
pub fn validate(&self) {
self.experts.validate();let device=self.experts.gate.local.val().device();let topology=(self.experts.gate.rank,self.experts.gate.world_size);
self.router.visit_shards(&mut |parameter| {assert_eq!(parameter.local.val().device(),device,"router/expert devices differ");
assert_eq!((parameter.rank,parameter.world_size),topology,"router/expert topologies differ");});
self.router.visit_packed_shards(&mut |parameter| {assert_eq!(parameter.local.val().device(),device,"packed router/expert devices differ");
assert_eq!((parameter.rank,parameter.world_size),topology,"packed router/expert topologies differ");});
if let Some(bias)=&self.correction_bias {assert_eq!(bias.logical_shape,[self.experts.dimensions()[0]],"native correction bias logical width differs");
assert_eq!(bias.local.val().dtype(),DType::F32,"native correction bias must retain FP32");assert_eq!(bias.local.val().device(),device,"native correction bias device differs");
assert_eq!((bias.rank,bias.world_size),topology,"native correction bias topology differs");}
}
}
#[derive(Debug)]
pub enum FullyShardedNativeMoeError<C:core::fmt::Debug,P:core::fmt::Debug,M:core::fmt::Debug> {
Collective(C),
Layer(NativeMoeLayerError<P,M>),
}
impl<C:core::fmt::Debug,P:core::fmt::Debug,M:core::fmt::Debug> core::fmt::Display for FullyShardedNativeMoeError<C,P,M> {
fn fmt(&self,f:&mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {Self::Collective(error)=>write!(f,"native MoE data transport: {error:?}"),Self::Layer(error)=>write!(f,"native MoE branch: {error}")}
}
}
impl<C:core::fmt::Debug,P:core::fmt::Debug,M:core::fmt::Debug> core::error::Error for FullyShardedNativeMoeError<C,P,M> {}
macro_rules! native_moe_gather {
($backend:ty,[$($generics:tt)*],$gather:ident) => {
impl<$($generics)*> FullyShardedNativeSwiGluExperts<$backend> {
pub fn $gather<C:IntegerTensorCollective<B>>(&self,communicator:C) -> Result<NativeSwiGluExperts<$backend>,C::Error> {
self.validate();Ok(NativeSwiGluExperts::from_parameters(Param::initialized(self.gate.local.id,self.gate.$gather::<C,3>(communicator.clone())?),
Param::initialized(self.up.local.id,self.up.$gather::<C,3>(communicator.clone())?),Param::initialized(self.down.local.id,self.down.$gather::<C,3>(communicator)?)))
}
}
impl<$($generics)*,P:GatherTransformerProjection<$backend,B>> FullyShardedNativeMoeLayer<$backend,P> {
pub fn $gather<C:IntegerTensorCollective<B>>(&self,communicator:C) -> Result<NativeMoeLayer<$backend,P::Gathered>,C::Error> {
self.validate();Ok(NativeMoeLayer::from_parts(self.router.gather_projection(communicator.clone())?,self.experts.$gather(communicator.clone())?,
self.correction_bias.as_ref().map(|bias|bias.$gather::<C,1>(communicator).map(|value|Param::initialized(bias.local.id,value))).transpose()?,self.options,self.router_input_dtype))
}
}
};
}
native_moe_gather!(B,[B:Backend],gather_inference);
native_moe_gather!(Autodiff<B,S>,[B:Backend,S:CheckpointStrategy],gather);
macro_rules! native_moe_execution {
($backend:ty,[$($generics:tt)*],$gather:ident,$forward:ident,$detailed:ident) => {
impl<$($generics)*,P:GatherTransformerProjection<$backend,B>> FullyShardedNativeMoeLayer<$backend,P>
where P::Gathered:TransformerProjection<$backend> {
pub fn $forward<C:IntegerTensorCollective<B>,const D:usize>(&self,input:Tensor<$backend,D>,communicator:C)
-> Result<Tensor<$backend,D>,FullyShardedNativeMoeError<C::Error,<P::Gathered as TransformerProjection<$backend>>::Error,<$backend as MoeOps>::MoeError>> {
self.$gather(communicator).map_err(FullyShardedNativeMoeError::Collective)?.forward(input).map_err(FullyShardedNativeMoeError::Layer)
}
pub fn $detailed<C:IntegerTensorCollective<B>,const D:usize>(&self,input:Tensor<$backend,D>,communicator:C)
-> Result<NativeMoeLayerOutput<$backend,D>,FullyShardedNativeMoeError<C::Error,<P::Gathered as TransformerProjection<$backend>>::Error,<$backend as MoeOps>::MoeError>> {
self.$gather(communicator).map_err(FullyShardedNativeMoeError::Collective)?.forward_detailed(input).map_err(FullyShardedNativeMoeError::Layer)
}
}
};
}
native_moe_execution!(B,[B:MoeOps],gather_inference,forward_inference,forward_detailed_inference);
native_moe_execution!(Autodiff<B,S>,[B:MoeOps,S:CheckpointStrategy],gather,forward,forward_detailed);