1use crate::{DeviceBackend,DeviceRuntime,FloatElement,IntElement,element::BoolElement};
2use ruda_tensor::{moe::{MoeOptions,MoeCombineGradientStrategy},moe_exchange::*,tensor::{FloatTensor,IntTensor}};
3use rudnn::moe::{self,MoeError,RouterTrainingPlan,DispatchedTokens,ReceivedExpertRows,ExpertTrainingCache,SwiGluExperts};
4use ruda_core::tensor::DType;
5use ruda_kernel::tensor::{RudaTensor,contiguous::into_contiguous,allocation::empty_device_contiguous_dtype};
6use super::moe::{weight_options,expert_strategy,select};
7
8#[derive(Debug,Clone)]
10pub struct NativeMoeDispatchState<R:DeviceRuntime> {router:RouterTrainingPlan<R>,dispatched:DispatchedTokens<R>}
11#[derive(Debug,Clone)]
13pub enum NativeMoeReceivedState<R:DeviceRuntime> {
14 Grouped {rows:ReceivedExpertRows<R>,cache:Option<ExpertTrainingCache<R>>,options:MoeReceivedOptions},
16 Empty {input:RudaTensor<R>,gate:RudaTensor<R>,up:RudaTensor<R>,down:RudaTensor<R>},
18}
19fn read_u32<R:DeviceRuntime>(tensor:RudaTensor<R>) -> Result<Vec<u32>,MoeError> {
20 ruda_core::future::block_on(ruda_kernel::tensor::readback::into_data(tensor)).map_err(|_|MoeError("native expert prefix metadata read failed"))?
21 .to_vec::<u32>().map_err(|_|MoeError("native expert prefix metadata storage mismatch"))
22}
23fn combine_strategy(strategy:MoeCombineGradientStrategy) -> moe::CombineGradientStrategy {
24 match strategy {MoeCombineGradientStrategy::Serial=>moe::CombineGradientStrategy::Serial,MoeCombineGradientStrategy::Plane=>moe::CombineGradientStrategy::Plane}
25}
26impl<R,F,I,BT> MoeDispatchOps for DeviceBackend<R,F,I,BT> where R:DeviceRuntime,F:FloatElement,I:IntElement,BT:BoolElement {
27 type MoeDispatchState=NativeMoeDispatchState<R>;
28 fn moe_dispatch(input:FloatTensor<Self>,logits:FloatTensor<Self>,bias:Option<FloatTensor<Self>>,options:MoeOptions) -> Result<MoeDispatched<Self>,Self::MoeError> {
29 let logits=into_contiguous(logits);
30 for value in [&input].into_iter().chain(bias.iter()) {if value.device!=logits.device || !value.client.same_execution_queue(&logits.client) {
31 return Err(MoeError("native dispatch operands must share the original device and queue"));}}
32 let router=select(logits.clone(),bias,options.selection)?.into_training(logits,weight_options(options.weights))?;
33 let dispatched=router.routing().clone().dispatch(input)?;
34 Ok(MoeDispatched {values:dispatched.values().clone(),weights:router.routing().weights().clone(),selected_experts:router.routing().expert_indices().clone(),
35 row_experts:dispatched.row_experts().clone(),state:NativeMoeDispatchState {router,dispatched}})
36 }
37 fn moe_dispatch_counts(state:&Self::MoeDispatchState,expert_prefix:&[usize]) -> Result<Vec<usize>,Self::MoeError> {
38 let experts=state.router.routing().experts();if expert_prefix.len()<2 || expert_prefix[0]!=0 || expert_prefix.last()!=Some(&experts)
39 || expert_prefix.windows(2).any(|pair|pair[0]>pair[1]) {return Err(MoeError("expert ownership must be a complete nondecreasing original global prefix"));}
40 let offsets=read_u32(state.dispatched.expert_offsets().clone())?;
41 if offsets.len()!=experts+1 || offsets.first()!=Some(&0) || offsets.windows(2).any(|pair|pair[0]>pair[1])
42 || offsets.last().map(|&value|value as usize)!=Some(state.dispatched.values().meta.shape()[0]) {return Err(MoeError("native expert assignment prefix metadata mismatch"));}
43 Ok(expert_prefix.windows(2).map(|pair|(offsets[pair[1]]-offsets[pair[0]]) as usize).collect())
44 }
45 fn moe_dispatch_backward(state:Self::MoeDispatchState,gradient:FloatTensor<Self>) -> Result<FloatTensor<Self>,Self::MoeError> {state.dispatched.dispatch_backward(gradient)}
46 fn moe_dispatch_weights_backward(state:Self::MoeDispatchState,gradient:FloatTensor<Self>) -> Result<FloatTensor<Self>,Self::MoeError> {state.router.backward(&gradient)}
47 fn moe_combine(state:Self::MoeDispatchState,expert_values:FloatTensor<Self>,weights:FloatTensor<Self>,_backward_strategy:MoeCombineGradientStrategy)
48 -> Result<FloatTensor<Self>,Self::MoeError> {
49 state.dispatched.with_weights(weights)?.combine(expert_values)
50 }
51 fn moe_combine_backward(state:Self::MoeDispatchState,expert_values:FloatTensor<Self>,weights:FloatTensor<Self>,gradient:FloatTensor<Self>,
52 strategy:MoeCombineGradientStrategy,selection:MoeCombineSelection) -> Result<MoeCombineBackward<Self>,Self::MoeError> {
53 let result=state.dispatched.with_weights(weights)?.combine_backward_selected(&expert_values,gradient,combine_strategy(strategy),
54 moe::CombineGradientSelection {experts:selection.experts,weights:selection.weights})?;
55 Ok(MoeCombineBackward {experts:result.dexpert,weights:result.dweights})
56 }
57}
58impl<R,F,I,BT> MoeReceivedOps for DeviceBackend<R,F,I,BT> where R:DeviceRuntime,F:FloatElement,I:IntElement,BT:BoolElement {
59 type MoeReceivedState=NativeMoeReceivedState<R>;
60 fn moe_received_forward(input:FloatTensor<Self>,global_ids:IntTensor<Self>,gate:FloatTensor<Self>,up:FloatTensor<Self>,down:FloatTensor<Self>,
61 options:MoeReceivedOptions,selection:MoeReceivedSelection) -> Result<(FloatTensor<Self>,Self::MoeReceivedState),Self::MoeError> {
62 if input.meta.num_dims()!=2 || gate.meta.num_dims()!=3 || up.meta.shape()!=gate.meta.shape() || down.meta.num_dims()!=3 {
63 return Err(MoeError("native received experts require original input and gate/up/down cube ranks"));}
64 let (experts,inner,hidden)=(gate.meta.shape()[0],gate.meta.shape()[1],gate.meta.shape()[2]);
65 if inner==0 || hidden==0 || down.meta.shape()[..]!=[experts,hidden,inner] || input.meta.shape()[1]!=hidden
66 || global_ids.meta.shape()[..]!=[input.meta.shape()[0]] || global_ids.dtype!=DType::U32 || global_ids.qparams.is_some()
67 || options.expert_start.checked_add(experts).is_none_or(|end|end>u32::MAX as usize) {return Err(MoeError("native received local cube geometry/global expert range mismatch"));}
68 for value in [&input,&gate,&up,&down] {
69 if !matches!(value.dtype,DType::F16|DType::BF16|DType::F32) || value.qparams.is_some() || value.dtype!=input.dtype || value.device!=input.device
70 || !value.client.same_execution_queue(&input.client) {return Err(MoeError("native received expert storage/device/queue mismatch"));}}
71 for value in [&input,&gate,&up,&down] {if value.meta.shape().iter().any(|&axis|axis>u32::MAX as usize)
72 || value.meta.shape().iter().try_fold(1usize,|size,&axis|size.checked_mul(axis)).is_none_or(|size|size>u32::MAX as usize) {
73 return Err(MoeError("native received expert tensor exceeds U32 indexing"));}}
74 if global_ids.device!=input.device || !global_ids.client.same_execution_queue(&input.client) {return Err(MoeError("received U32 assignments must share the actual input queue"));}
75 if experts==0 {
76 if input.meta.shape()[0]!=0 {return Err(MoeError("a zero-expert owner cannot receive nonempty assignment rows"));}
77 let output=empty_device_contiguous_dtype(input.client.clone(),input.device.clone(),input.meta.shape().clone(),input.dtype);
78 return Ok((output,NativeMoeReceivedState::Empty {input,gate,up,down}));
79 }
80 let rows=ReceivedExpertRows::new(input,global_ids,options.expert_start,experts)?;let experts=SwiGluExperts::new(gate,up,down)?;
81 let (expert_output,cache)=if selection.input || selection.gate || selection.up || selection.down {
82 let trained=experts.forward_grouped_training(rows.grouped(),expert_strategy(options.forward))?;(trained.output,Some(trained.cache))
83 } else {(experts.forward_grouped_rows(rows.grouped(),expert_strategy(options.forward))?,None)};
84 let output=rows.restore(expert_output)?;Ok((output,NativeMoeReceivedState::Grouped {rows,cache,options}))
85 }
86 fn moe_received_backward(state:Self::MoeReceivedState,gradient:FloatTensor<Self>,selection:MoeReceivedSelection)
87 -> Result<MoeReceivedBackward<Self>,Self::MoeError> {
88 match state {
89 NativeMoeReceivedState::Empty {input,gate,up,down}=>{
90 if gradient.meta.shape()!=input.meta.shape() || gradient.dtype!=input.dtype || gradient.device!=input.device
91 || !gradient.client.same_execution_queue(&input.client) || gradient.qparams.is_some() {return Err(MoeError("empty-owner received seed metadata mismatch"));}
92 let empty=|like:RudaTensor<R>,dtype|empty_device_contiguous_dtype(like.client.clone(),like.device.clone(),like.meta.shape().clone(),dtype);
93 Ok(MoeReceivedBackward {input:selection.input.then(||empty(input,gradient.dtype)),gate:selection.gate.then(||empty(gate,DType::F32)),
94 up:selection.up.then(||empty(up,DType::F32)),down:selection.down.then(||empty(down,DType::F32))})
95 },
96 NativeMoeReceivedState::Grouped {rows,cache,options}=>{
97 if !selection.input && !selection.gate && !selection.up && !selection.down {return Ok(MoeReceivedBackward {input:None,gate:None,up:None,down:None});}
98 let cache=cache.ok_or(MoeError("received forward did not retain the requested expert VJP cache"))?;
99 let result=cache.backward_selected(rows.sort_gradient(gradient)?,expert_strategy(options.backward),
100 moe::ExpertGradientSelection {input:selection.input,gate:selection.gate,up:selection.up,down:selection.down})?;
101 Ok(MoeReceivedBackward {input:result.dinput.map(|value|rows.restore(value)).transpose()?,gate:result.dgate,up:result.dup,down:result.ddown})
102 },
103 }
104 }
105}