Skip to main content

ruda_tensor_device/dispatch/
moe_exchange.rs

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/// Source-local actual native discrete mappings and continuous router VJP.
9#[derive(Debug,Clone)]
10pub struct NativeMoeDispatchState<R:DeviceRuntime> {router:RouterTrainingPlan<R>,dispatched:DispatchedTokens<R>}
11/// Actual local received rows/cache, or the exact empty-owner original source metadata.
12#[derive(Debug,Clone)]
13pub enum NativeMoeReceivedState<R:DeviceRuntime> {
14    /// Validated native private row mapping and original expert cache when requested.
15    Grouped {rows:ReceivedExpertRows<R>,cache:Option<ExpertTrainingCache<R>>,options:MoeReceivedOptions},
16    /// True zero-expert owner: every original local cube and input has zero rows.
17    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}