Skip to main content

ruda_tensor_device/dispatch/
expert_projection.rs

1use crate::{DeviceBackend,DeviceRuntime,FloatElement,IntElement,element::BoolElement};
2use ruda_tensor::{FloatDType,expert_projection::*,ops::FloatTensorOps,tensor::{FloatTensor,IntTensor}};
3use rudnn::moe::{MoeError,ReceivedExpertRows,ExpertProjectionCache,expert_projection,swiglu_activation,swiglu_activation_backward,SwiGluActivationSelection};
4use rublas::tensor_grouped::GroupedGradientSelection;
5use super::moe::expert_strategy;
6
7/// Actual original native grouped operands, valid private row COPY mapping and independent VJP policy.
8#[derive(Clone,Debug)]
9pub struct NativeExpertProjectionState<R:DeviceRuntime> {rows:ReceivedExpertRows<R>,cache:ExpertProjectionCache<R>,options:ExpertProjectionOptions,dtype:FloatDType}
10impl<R,F,I,BT> ExpertProjectionOps for DeviceBackend<R,F,I,BT> where R:DeviceRuntime,F:FloatElement,I:IntElement,BT:BoolElement {
11    type ExpertProjectionError=MoeError;
12    type ExpertProjectionState=NativeExpertProjectionState<R>;
13    fn expert_projection_forward(input:FloatTensor<Self>,global_ids:IntTensor<Self>,weights:FloatTensor<Self>,options:ExpertProjectionOptions)
14        -> Result<(FloatTensor<Self>,Self::ExpertProjectionState),Self::ExpertProjectionError> {
15        if weights.meta.num_dims()!=3 {return Err(MoeError("native floating expert projection requires actual rank-three weights"));}
16        let experts=weights.meta.shape()[0];let dtype=input.dtype.into();let rows=ReceivedExpertRows::new(input,global_ids,options.expert_start,experts)?;
17        let (output,cache)=expert_projection(rows.grouped(),weights,expert_strategy(options.forward))?;
18        Ok((rows.restore(output)?,NativeExpertProjectionState {rows,cache,options,dtype}))
19    }
20    fn expert_projection_backward(state:Self::ExpertProjectionState,gradient:FloatTensor<Self>,selection:ExpertProjectionSelection)
21        -> Result<ExpertProjectionBackward<Self>,Self::ExpertProjectionError> {
22        let gradient=state.rows.sort_gradient(Self::float_cast(gradient,state.dtype))?;
23        let result=state.cache.backward(gradient,expert_strategy(state.options.backward),GroupedGradientSelection {input:selection.input,weights:selection.weights})?;
24        Ok(ExpertProjectionBackward {input:result.dinput.map(|value|state.rows.restore(value)).transpose()?,weights:result.dweights})
25    }
26}
27impl<R,F,I,BT> NativeSwiGluOps for DeviceBackend<R,F,I,BT> where R:DeviceRuntime,F:FloatElement,I:IntElement,BT:BoolElement {
28    type SwiGluError=MoeError;
29    fn native_swiglu(gate:FloatTensor<Self>,up:FloatTensor<Self>) -> Result<FloatTensor<Self>,Self::SwiGluError> {swiglu_activation(gate,up)}
30    fn native_swiglu_backward(gate:FloatTensor<Self>,up:FloatTensor<Self>,gradient:FloatTensor<Self>,selection:NativeSwiGluSelection)
31        -> Result<NativeSwiGluBackward<Self>,Self::SwiGluError> {
32        let gradient=Self::float_cast(gradient,gate.dtype.into());
33        let result=swiglu_activation_backward(gate,up,gradient,SwiGluActivationSelection {gate:selection.gate,up:selection.up})?;
34        Ok(NativeSwiGluBackward {gate:result.gate,up:result.up})
35    }
36}