use crate::{Backend,tensor::{FloatTensor,IntTensor}};
use core::fmt;
#[derive(Clone,Copy,Debug,PartialEq,Eq)]
pub enum MoeRouterScoring {
Softmax,
Sigmoid,
}
#[derive(Clone,Copy,Debug,PartialEq)]
pub struct MoeRouterWeightOptions {
pub scoring:MoeRouterScoring,
pub renormalize:bool,
pub scale:f32,
}
#[derive(Clone,Copy,Debug,PartialEq)]
pub enum MoeSelectionOptions {
Softmax {
top_k:usize,
renormalize:bool,
},
SigmoidGrouped {
top_k:usize,
groups:usize,
selected_groups:usize,
group_top_two:bool,
renormalize:bool,
scale:f32,
},
}
#[derive(Clone,Copy,Debug,PartialEq,Eq)]
pub enum MoeExpertStrategy {
Scalar,
Auto,
TensorCore,
}
#[derive(Clone,Copy,Debug,PartialEq,Eq)]
pub enum MoeCombineGradientStrategy {
Serial,
Plane,
}
#[derive(Clone,Copy,Debug,PartialEq)]
pub struct MoeOptions {
pub selection:MoeSelectionOptions,
pub weights:MoeRouterWeightOptions,
pub forward:MoeExpertStrategy,
pub backward:MoeExpertStrategy,
pub combine_backward:MoeCombineGradientStrategy,
}
#[derive(Debug)]
pub struct MoeBackward<B:Backend> {
pub input:FloatTensor<B>,
pub logits:FloatTensor<B>,
pub gate:FloatTensor<B>,
pub up:FloatTensor<B>,
pub down:FloatTensor<B>,
}
#[derive(Clone,Copy,Debug,PartialEq,Eq)]
pub struct MoeGradientSelection {
pub input:bool,
pub logits:bool,
pub gate:bool,
pub up:bool,
pub down:bool,
}
#[derive(Debug)]
pub struct MoeBackwardSelected<B:Backend> {
pub input:Option<FloatTensor<B>>,
pub logits:Option<FloatTensor<B>>,
pub gate:Option<FloatTensor<B>>,
pub up:Option<FloatTensor<B>>,
pub down:Option<FloatTensor<B>>,
}
#[derive(Debug)]
pub enum MoeAutodiffError<E:fmt::Debug> {
Native(E),
HigherDerivativeUnsupported,
}
impl<E:fmt::Debug> fmt::Display for MoeAutodiffError<E> {
fn fmt(&self,f:&mut fmt::Formatter<'_>) -> fmt::Result {
match self {Self::Native(error)=>write!(f,"native MoE: {error:?}"),Self::HigherDerivativeUnsupported=>f.write_str("native MoE provides first-order derivatives only")}
}
}
impl<E:fmt::Debug> core::error::Error for MoeAutodiffError<E> {}
pub trait MoeOps:Backend {
type MoeError:fmt::Debug;
type MoeState:Clone+Send+fmt::Debug+'static;
fn moe_selected_weights(logits:FloatTensor<Self>,indices:IntTensor<Self>,options:MoeRouterWeightOptions) -> Result<FloatTensor<Self>,Self::MoeError>;
fn moe_selected_weights_backward(logits:FloatTensor<Self>,indices:IntTensor<Self>,gradient:FloatTensor<Self>,options:MoeRouterWeightOptions)
-> Result<FloatTensor<Self>,Self::MoeError>;
fn moe_forward(input:FloatTensor<Self>,logits:FloatTensor<Self>,correction_bias:Option<FloatTensor<Self>>,
gate:FloatTensor<Self>,up:FloatTensor<Self>,down:FloatTensor<Self>,options:MoeOptions) -> Result<(FloatTensor<Self>,Self::MoeState),Self::MoeError>;
fn moe_forward_selected(input:FloatTensor<Self>,logits:FloatTensor<Self>,correction_bias:Option<FloatTensor<Self>>,
gate:FloatTensor<Self>,up:FloatTensor<Self>,down:FloatTensor<Self>,options:MoeOptions,selection:MoeGradientSelection)
-> Result<(FloatTensor<Self>,Self::MoeState),Self::MoeError> {
let _=selection;Self::moe_forward(input,logits,correction_bias,gate,up,down,options)
}
fn moe_inference(input:FloatTensor<Self>,logits:FloatTensor<Self>,correction_bias:Option<FloatTensor<Self>>,
gate:FloatTensor<Self>,up:FloatTensor<Self>,down:FloatTensor<Self>,options:MoeOptions) -> Result<FloatTensor<Self>,Self::MoeError>;
fn moe_route_indices(state:&Self::MoeState) -> IntTensor<Self>;
fn moe_backward(state:Self::MoeState,gradient:FloatTensor<Self>) -> Result<MoeBackward<Self>,Self::MoeError>;
fn moe_backward_selected(state:Self::MoeState,gradient:FloatTensor<Self>,selection:MoeGradientSelection)
-> Result<MoeBackwardSelected<Self>,Self::MoeError> {
let result=Self::moe_backward(state,gradient)?;
Ok(MoeBackwardSelected {input:selection.input.then_some(result.input),logits:selection.logits.then_some(result.logits),
gate:selection.gate.then_some(result.gate),up:selection.up.then_some(result.up),down:selection.down.then_some(result.down)})
}
}