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#[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}