Skip to main content

ruda_tensor_device/dispatch/
moe.rs

1use crate::{DeviceBackend,DeviceRuntime,FloatElement,IntElement,element::BoolElement};
2use ruda_tensor::{moe::{MoeOps,MoeOptions,MoeSelectionOptions,MoeRouterScoring,MoeRouterWeightOptions,
3    MoeExpertStrategy,MoeCombineGradientStrategy,MoeBackward,MoeGradientSelection,MoeBackwardSelected},tensor::{FloatTensor,IntTensor}};
4use rudnn::moe::{self,RoutingPlan,RouterTrainingPlan,DispatchedTokens,ExpertTrainingCache,MoeError,SwiGluExperts};
5use ruda_kernel::tensor::{RudaTensor,contiguous::into_contiguous};
6
7/// Actual native forward state retaining unchanged original logits and dispatch metadata.
8#[derive(Debug,Clone)]
9pub struct NativeMoeState<R:DeviceRuntime> {
10    router:RouterTrainingPlan<R>,
11    dispatched:DispatchedTokens<R>,
12    experts:Option<ExpertTrainingCache<R>>,
13    expert_output:RudaTensor<R>,
14    options:MoeOptions,
15}
16pub(super) fn weight_options(options:MoeRouterWeightOptions) -> moe::RouterWeightOptions {
17    moe::RouterWeightOptions {scoring:match options.scoring {MoeRouterScoring::Softmax=>moe::RouterScoring::Softmax,MoeRouterScoring::Sigmoid=>moe::RouterScoring::Sigmoid},
18        renormalize:options.renormalize,scale:options.scale}
19}
20pub(super) fn expert_strategy(strategy:MoeExpertStrategy) -> moe::GroupedStrategy {
21    match strategy {MoeExpertStrategy::Scalar=>moe::GroupedStrategy::Scalar,MoeExpertStrategy::Auto=>moe::GroupedStrategy::Auto,MoeExpertStrategy::TensorCore=>moe::GroupedStrategy::TensorCore}
22}
23pub(super) fn select<R:DeviceRuntime>(logits:RudaTensor<R>,bias:Option<RudaTensor<R>>,options:MoeSelectionOptions) -> Result<RoutingPlan<R>,MoeError> {
24    match options {
25        MoeSelectionOptions::Softmax {top_k,renormalize}=>{
26            if bias.is_some() {return Err(MoeError("softmax selection does not accept a sigmoid correction bias"));}
27            moe::route(logits,moe::RoutingOptions {top_k,renormalize})
28        },
29        MoeSelectionOptions::SigmoidGrouped {top_k,groups,selected_groups,group_top_two,renormalize,scale}=>
30            moe::route_sigmoid_grouped(logits,bias,moe::GroupRoutingOptions {top_k,groups,selected_groups,group_top_two,renormalize,scale}),
31    }
32}
33fn prepare<R:DeviceRuntime>(input:RudaTensor<R>,logits:RudaTensor<R>,bias:Option<RudaTensor<R>>,gate:RudaTensor<R>,up:RudaTensor<R>,down:RudaTensor<R>,options:MoeOptions)
34    -> Result<(SwiGluExperts<R>,DispatchedTokens<R>,RouterTrainingPlan<R>),MoeError> {
35    let logits=into_contiguous(logits);
36    for value in [&input,&gate,&up,&down].into_iter().chain(bias.iter()) {
37        if value.device!=logits.device || !value.client.same_execution_queue(&logits.client) {return Err(MoeError("MoE operands must share an actual device and execution queue"));}
38    }
39    let experts=SwiGluExperts::new(gate,up,down)?;
40    let routing=select(logits.clone(),bias,options.selection)?.into_training(logits,weight_options(options.weights))?;
41    let dispatched=routing.routing().clone().dispatch(input)?;
42    Ok((experts,dispatched,routing))
43}
44impl<R,F,I,BT> MoeOps for DeviceBackend<R,F,I,BT>
45    where R:DeviceRuntime,F:FloatElement,I:IntElement,BT:BoolElement {
46    type MoeError=MoeError;
47    type MoeState=NativeMoeState<R>;
48    fn moe_selected_weights(logits:FloatTensor<Self>,indices:IntTensor<Self>,options:MoeRouterWeightOptions) -> Result<FloatTensor<Self>,Self::MoeError> {
49        moe::selected_router_weights(&logits,&indices,weight_options(options))
50    }
51    fn moe_selected_weights_backward(logits:FloatTensor<Self>,indices:IntTensor<Self>,gradient:FloatTensor<Self>,options:MoeRouterWeightOptions)
52        -> Result<FloatTensor<Self>,Self::MoeError> {
53        moe::selected_router_backward(&logits,&indices,&gradient,weight_options(options))
54    }
55    fn moe_forward(input:FloatTensor<Self>,logits:FloatTensor<Self>,correction_bias:Option<FloatTensor<Self>>,
56        gate:FloatTensor<Self>,up:FloatTensor<Self>,down:FloatTensor<Self>,options:MoeOptions) -> Result<(FloatTensor<Self>,Self::MoeState),Self::MoeError> {
57        let (experts,dispatched,routing)=prepare(input,logits,correction_bias,gate,up,down,options)?;
58        let trained=experts.forward_dispatched_training(&dispatched,expert_strategy(options.forward))?;
59        let output=dispatched.combine(trained.output.clone())?;
60        Ok((output,NativeMoeState {router:routing,dispatched,experts:Some(trained.cache),expert_output:trained.output,options}))
61    }
62    fn moe_forward_selected(input:FloatTensor<Self>,logits:FloatTensor<Self>,correction_bias:Option<FloatTensor<Self>>,
63        gate:FloatTensor<Self>,up:FloatTensor<Self>,down:FloatTensor<Self>,options:MoeOptions,selection:MoeGradientSelection)
64        -> Result<(FloatTensor<Self>,Self::MoeState),Self::MoeError> {
65        if selection.input || selection.gate || selection.up || selection.down {
66            return Self::moe_forward(input,logits,correction_bias,gate,up,down,options);
67        }
68        let (experts,dispatched,routing)=prepare(input,logits,correction_bias,gate,up,down,options)?;
69        let expert_output=experts.forward_dispatched_with_strategy(&dispatched,expert_strategy(options.forward))?;
70        let output=dispatched.combine(expert_output.clone())?;
71        Ok((output,NativeMoeState {router:routing,dispatched,experts:None,expert_output,options}))
72    }
73    fn moe_inference(input:FloatTensor<Self>,logits:FloatTensor<Self>,correction_bias:Option<FloatTensor<Self>>,
74        gate:FloatTensor<Self>,up:FloatTensor<Self>,down:FloatTensor<Self>,options:MoeOptions) -> Result<FloatTensor<Self>,Self::MoeError> {
75        let (experts,dispatched,_)=prepare(input,logits,correction_bias,gate,up,down,options)?;
76        let output=experts.forward_dispatched_with_strategy(&dispatched,expert_strategy(options.forward))?;
77        dispatched.combine(output)
78    }
79    fn moe_route_indices(state:&Self::MoeState) -> IntTensor<Self> {state.router.routing().expert_indices().clone()}
80    fn moe_backward(state:Self::MoeState,gradient:FloatTensor<Self>) -> Result<MoeBackward<Self>,Self::MoeError> {
81        let cache=state.experts.ok_or(MoeError("selected forward did not retain an expert-input/weight VJP cache"))?;
82        let combine=state.dispatched.combine_backward_with_strategy(&state.expert_output,gradient,
83            match state.options.combine_backward {MoeCombineGradientStrategy::Serial=>moe::CombineGradientStrategy::Serial,MoeCombineGradientStrategy::Plane=>moe::CombineGradientStrategy::Plane})?;
84        let logits=state.router.backward(&combine.dweights)?;
85        let experts=cache.backward_with_strategy(combine.dexpert,expert_strategy(state.options.backward))?;
86        let input=state.dispatched.dispatch_backward(experts.dinput)?;
87        Ok(MoeBackward {input,logits,gate:experts.dgate,up:experts.dup,down:experts.ddown})
88    }
89    fn moe_backward_selected(state:Self::MoeState,gradient:FloatTensor<Self>,selection:MoeGradientSelection)
90        -> Result<MoeBackwardSelected<Self>,Self::MoeError> {
91        if (selection.input || selection.gate || selection.up || selection.down) && state.experts.is_none() {
92            return Err(MoeError("selected forward did not retain the requested expert-input/weight VJP cache"));
93        }
94        let combine=state.dispatched.combine_backward_selected(&state.expert_output,gradient,
95            match state.options.combine_backward {MoeCombineGradientStrategy::Serial=>moe::CombineGradientStrategy::Serial,MoeCombineGradientStrategy::Plane=>moe::CombineGradientStrategy::Plane},
96            moe::CombineGradientSelection {experts:selection.input || selection.gate || selection.up || selection.down,weights:selection.logits})?;
97        let logits=combine.dweights.map(|gradient|state.router.backward(&gradient)).transpose()?;
98        let experts=combine.dexpert.map(|gradient|state.experts.expect("validated native expert cache").backward_selected(gradient,expert_strategy(state.options.backward),
99            moe::ExpertGradientSelection {input:selection.input,gate:selection.gate,up:selection.up,down:selection.down})).transpose()?;
100        let (input,gate,up,down)=if let Some(experts)=experts {
101            (experts.dinput.map(|gradient|state.dispatched.dispatch_backward(gradient)).transpose()?,experts.dgate,experts.dup,experts.ddown)
102        } else {(None,None,None,None)};
103        Ok(MoeBackwardSelected {input,logits,gate,up,down})
104    }
105}