ruda_tensor_device/dispatch/
expert_projection.rs1use 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#[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}