use core::fmt;
use ruda_model::{module::{Module,Param},tensor::{Tensor,Int,DType,FloatDType,TensorPrimitive,MoeOps,MoeOptions,MoeGradientSelection,backend::Backend}};
use crate::transformer::{TransformerProjectionShape,TransformerProjection};
#[derive(Module,Debug)]
pub struct NativeSwiGluExperts<B:Backend> {
pub gate:Param<Tensor<B,3>>,
pub up:Param<Tensor<B,3>>,
pub down:Param<Tensor<B,3>>,
}
impl<B:Backend> NativeSwiGluExperts<B> {
pub fn from_parameters(gate:Param<Tensor<B,3>>,up:Param<Tensor<B,3>>,down:Param<Tensor<B,3>>) -> Self {
let experts=Self {gate,up,down};experts.validate();experts
}
pub fn dimensions(&self) -> [usize;3] {let [experts,inner,hidden]=self.gate.val().dims();[experts,hidden,inner]}
pub fn validate(&self) {
let gate=self.gate.val();let up=self.up.val();let down=self.down.val();let [experts,inner,hidden]=gate.dims();
assert!(experts>0 && inner>0 && hidden>0,"native MoE expert axes must be positive");
assert_eq!(up.dims(),gate.dims(),"native gate/up expert geometry differs");assert_eq!(down.dims(),[experts,hidden,inner],"native down expert geometry differs");
assert!(matches!(gate.dtype(),DType::F16|DType::BF16|DType::F32),"unsupported native expert storage");
for value in [&up,&down] {assert_eq!(value.dtype(),gate.dtype(),"native expert storage differs");assert_eq!(value.device(),gate.device(),"native expert devices differ");}
}
}
impl<B:MoeOps> NativeSwiGluExperts<B> {
pub fn forward(&self,input:Tensor<B,2>,logits:Tensor<B,2>,correction_bias:Option<Tensor<B,1>>,options:MoeOptions) -> Result<Tensor<B,2>,B::MoeError> {
self.validate();B::moe_inference(input.into_primitive().tensor(),logits.into_primitive().tensor(),correction_bias.map(|bias|bias.into_primitive().tensor()),
self.gate.val().into_primitive().tensor(),self.up.val().into_primitive().tensor(),self.down.val().into_primitive().tensor(),options)
.map(|output|Tensor::from_primitive(TensorPrimitive::Float(output)))
}
pub fn forward_with_state(&self,input:Tensor<B,2>,logits:Tensor<B,2>,correction_bias:Option<Tensor<B,1>>,options:MoeOptions)
-> Result<(Tensor<B,2>,B::MoeState),B::MoeError> {
self.validate();let (output,state)=B::moe_forward(input.into_primitive().tensor(),logits.into_primitive().tensor(),correction_bias.map(|bias|bias.into_primitive().tensor()),
self.gate.val().into_primitive().tensor(),self.up.val().into_primitive().tensor(),self.down.val().into_primitive().tensor(),options)?;
Ok((Tensor::from_primitive(TensorPrimitive::Float(output)),state))
}
pub fn forward_selected_with_state(&self,input:Tensor<B,2>,logits:Tensor<B,2>,correction_bias:Option<Tensor<B,1>>,options:MoeOptions,selection:MoeGradientSelection)
-> Result<(Tensor<B,2>,B::MoeState),B::MoeError> {
self.validate();let (output,state)=B::moe_forward_selected(input.into_primitive().tensor(),logits.into_primitive().tensor(),correction_bias.map(|bias|bias.into_primitive().tensor()),
self.gate.val().into_primitive().tensor(),self.up.val().into_primitive().tensor(),self.down.val().into_primitive().tensor(),options,selection)?;
Ok((Tensor::from_primitive(TensorPrimitive::Float(output)),state))
}
pub fn routing_indices(&self,state:&B::MoeState) -> Tensor<B,2,Int> {Tensor::from_primitive(B::moe_route_indices(state))}
pub fn backward_with_state(&self,state:B::MoeState,gradient:Tensor<B,2>) -> Result<NativeMoeBackward<B>,B::MoeError> {
let result=B::moe_backward(state,gradient.into_primitive().tensor())?;
Ok(NativeMoeBackward {input:Tensor::from_primitive(TensorPrimitive::Float(result.input)),logits:Tensor::from_primitive(TensorPrimitive::Float(result.logits)),
gate:Tensor::from_primitive(TensorPrimitive::Float(result.gate)),up:Tensor::from_primitive(TensorPrimitive::Float(result.up)),down:Tensor::from_primitive(TensorPrimitive::Float(result.down))})
}
pub fn backward_selected_with_state(&self,state:B::MoeState,gradient:Tensor<B,2>,selection:MoeGradientSelection)
-> Result<NativeMoeBackwardSelected<B>,B::MoeError> {
let result=B::moe_backward_selected(state,gradient.into_primitive().tensor(),selection)?;
Ok(NativeMoeBackwardSelected {input:result.input.map(|value|Tensor::from_primitive(TensorPrimitive::Float(value))),
logits:result.logits.map(|value|Tensor::from_primitive(TensorPrimitive::Float(value))),gate:result.gate.map(|value|Tensor::from_primitive(TensorPrimitive::Float(value))),
up:result.up.map(|value|Tensor::from_primitive(TensorPrimitive::Float(value))),down:result.down.map(|value|Tensor::from_primitive(TensorPrimitive::Float(value)))})
}
}
#[derive(Debug)]
pub struct NativeMoeBackward<B:Backend> {
pub input:Tensor<B,2>,
pub logits:Tensor<B,2>,
pub gate:Tensor<B,3>,
pub up:Tensor<B,3>,
pub down:Tensor<B,3>,
}
#[derive(Debug)]
pub struct NativeMoeBackwardSelected<B:Backend> {
pub input:Option<Tensor<B,2>>,
pub logits:Option<Tensor<B,2>>,
pub gate:Option<Tensor<B,3>>,
pub up:Option<Tensor<B,3>>,
pub down:Option<Tensor<B,3>>,
}
#[derive(Debug)]
pub enum NativeMoeLayerError<P:fmt::Debug,M:fmt::Debug> {
Router(P),
Experts(M),
}
impl<P:fmt::Debug,M:fmt::Debug> fmt::Display for NativeMoeLayerError<P,M> {
fn fmt(&self,f:&mut fmt::Formatter<'_>) -> fmt::Result {
match self {Self::Router(error)=>write!(f,"native router projection: {error:?}"),Self::Experts(error)=>write!(f,"native routed experts: {error:?}")}
}
}
impl<P:fmt::Debug,M:fmt::Debug> core::error::Error for NativeMoeLayerError<P,M> {}
#[derive(Module,Debug)]
pub struct NativeMoeLayer<B:Backend,P:Module<B>> {
pub router:P,
pub experts:NativeSwiGluExperts<B>,
pub correction_bias:Option<Param<Tensor<B,1>>>,
#[module(skip)]
pub options:MoeOptions,
#[module(skip)]
pub router_input_dtype:Option<FloatDType>,
}
impl<B:Backend,P:TransformerProjectionShape<B>> NativeMoeLayer<B,P> {
pub fn from_parts(router:P,experts:NativeSwiGluExperts<B>,correction_bias:Option<Param<Tensor<B,1>>>,options:MoeOptions,router_input_dtype:Option<FloatDType>) -> Self {
let layer=Self {router,experts,correction_bias,options,router_input_dtype};layer.validate();layer
}
pub fn width(&self) -> usize {self.experts.dimensions()[1]}
pub fn validate(&self) {
self.experts.validate();let [experts,hidden,_]=self.experts.dimensions();assert_eq!(self.router.dimensions(),[hidden,experts],"native router/expert geometry differs");
if let Some(bias)=&self.correction_bias {let bias=bias.val();assert_eq!(bias.dims(),[experts],"native correction bias width differs");
assert_eq!(bias.dtype(),DType::F32,"native correction bias must retain FP32");assert_eq!(bias.device(),self.experts.gate.val().device(),"native correction bias device differs");}
if let Some(dtype)=self.router_input_dtype {assert!(matches!(DType::from(dtype),DType::F16|DType::BF16|DType::F32),"unsupported native router compute storage");}
}
}
#[derive(Debug)]
pub struct NativeMoeLayerOutput<B:MoeOps,const D:usize> {
pub output:Tensor<B,D>,
pub router_logits:Tensor<B,2>,
pub selected_experts:Tensor<B,2,Int>,
pub state:B::MoeState,
}
impl<B:MoeOps,P:TransformerProjection<B>> NativeMoeLayer<B,P> {
fn project<const D:usize>(&self,input:Tensor<B,D>) -> Result<(Tensor<B,2>,Tensor<B,2>,[usize;D]),P::Error> {
self.validate();assert!(D>0,"native MoE input must have a feature axis");let shape=input.dims();assert_eq!(shape[D-1],self.width(),"native branch input width differs");
let rows=shape[..D-1].iter().try_fold(1usize,|count,&axis|count.checked_mul(axis)).expect("native token count overflows");
let input=input.reshape([rows,self.width()]);let routed=if let Some(dtype)=self.router_input_dtype {input.clone().cast(dtype)} else {input.clone()};
let logits=self.router.forward(routed)?;Ok((input,logits,shape))
}
pub fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,NativeMoeLayerError<P::Error,B::MoeError>> {
let (input,logits,shape)=self.project(input).map_err(NativeMoeLayerError::Router)?;
self.experts.forward(input,logits,self.correction_bias.as_ref().map(Param::val),self.options).map(|output|output.reshape(shape)).map_err(NativeMoeLayerError::Experts)
}
pub fn forward_detailed<const D:usize>(&self,input:Tensor<B,D>) -> Result<NativeMoeLayerOutput<B,D>,NativeMoeLayerError<P::Error,B::MoeError>> {
let (input,logits,shape)=self.project(input).map_err(NativeMoeLayerError::Router)?;
let (output,state)=self.experts.forward_with_state(input,logits.clone(),self.correction_bias.as_ref().map(Param::val),self.options).map_err(NativeMoeLayerError::Experts)?;
Ok(NativeMoeLayerOutput {output:output.reshape(shape),router_logits:logits,selected_experts:self.experts.routing_indices(&state),state})
}
}