Skip to main content

ruda_tensor/
moe.rs

1//! Native local MoE training contracts over existing device routing and expert kernels.
2use crate::{Backend,tensor::{FloatTensor,IntTensor}};
3use core::fmt;
4
5/// Original continuous weight scoring, independent of discrete expert selection.
6#[derive(Clone,Copy,Debug,PartialEq,Eq)]
7pub enum MoeRouterScoring {
8    /// FP32 softmax over every expert before selecting continuous weights.
9    Softmax,
10    /// Original pointwise stable sigmoid scores.
11    Sigmoid,
12}
13/// Explicit original continuous routing-weight policy; no model family or default is inferred.
14#[derive(Clone,Copy,Debug,PartialEq)]
15pub struct MoeRouterWeightOptions {
16    /// Original full-softmax or pointwise-sigmoid scoring.
17    pub scoring:MoeRouterScoring,
18    /// Normalize the actual selected slots, retaining duplicate-slot gather semantics.
19    pub renormalize:bool,
20    /// Finite positive FP32 multiplier, applied after selected normalization.
21    pub scale:f32,
22}
23/// Original discrete expert selection. Correction bias affects selection only.
24#[derive(Clone,Copy,Debug,PartialEq)]
25pub enum MoeSelectionOptions {
26    /// Original full softmax/top-k, with lower expert IDs resolving exact ties.
27    Softmax {
28        /// Actual selected experts per token.
29        top_k:usize,
30        /// Original inference selection-weight normalization; training weights are explicit separately.
31        renormalize:bool,
32    },
33    /// Original group-limited sigmoid selection with optional FP32 correction bias.
34    SigmoidGrouped {
35        /// Actual selected experts per token.
36        top_k:usize,
37        /// Actual equal expert-group count.
38        groups:usize,
39        /// Actual selected group count.
40        selected_groups:usize,
41        /// Use the sum of the two highest corrected scores rather than the group maximum.
42        group_top_two:bool,
43        /// Original inference selected-weight normalization.
44        renormalize:bool,
45        /// Original finite positive inference weight multiplier.
46        scale:f32,
47    },
48}
49/// Original segmented expert GEMM strategy, selected independently for forward and backward.
50#[derive(Clone,Copy,Debug,PartialEq,Eq)]
51pub enum MoeExpertStrategy {
52    /// Original FP32-accumulating device scalar kernel, not host computation.
53    Scalar,
54    /// Original capability-only choice; failed compilation/launch never selects recovery fallback.
55    Auto,
56    /// Require original supported half/BF16 cooperative matrix kernels.
57    TensorCore,
58}
59/// Original routing-weight VJP reduction order.
60#[derive(Clone,Copy,Debug,PartialEq,Eq)]
61pub enum MoeCombineGradientStrategy {
62    /// Original serial per-slot reduction order.
63    Serial,
64    /// Explicit original full-plane reduction, requiring supported hardware/grid.
65    Plane,
66}
67/// Actual complete local expert execution choices, with no implicit tuning or retry policy.
68#[derive(Clone,Copy,Debug,PartialEq)]
69pub struct MoeOptions {
70    /// Original actual discrete selection policy.
71    pub selection:MoeSelectionOptions,
72    /// Original actual continuous selected-weight policy, independent of correction bias.
73    pub weights:MoeRouterWeightOptions,
74    /// Original actual expert forward strategy.
75    pub forward:MoeExpertStrategy,
76    /// Original actual expert backward strategy.
77    pub backward:MoeExpertStrategy,
78    /// Original actual combine weight-gradient reduction strategy.
79    pub combine_backward:MoeCombineGradientStrategy,
80}
81/// Actual original first-order derivatives of the native routed expert branch.
82#[derive(Debug)]
83pub struct MoeBackward<B:Backend> {
84    /// Input gradient after dispatch-copy backward, without a second routing-weight multiplication.
85    pub input:FloatTensor<B>,
86    /// Original source-logits storage VJP, including nonselected softmax experts.
87    pub logits:FloatTensor<B>,
88    /// Original FP32 gate-weight gradients.
89    pub gate:FloatTensor<B>,
90    /// Original FP32 up-weight gradients.
91    pub up:FloatTensor<B>,
92    /// Original FP32 down-weight gradients.
93    pub down:FloatTensor<B>,
94}
95/// Actual first-order derivatives requested by the caller or tracked AD parents.
96#[derive(Clone,Copy,Debug,PartialEq,Eq)]
97pub struct MoeGradientSelection {
98    /// Preserve upstream input derivatives independently of expert-weight trainability.
99    pub input:bool,
100    /// Original fixed-selection logits derivative.
101    pub logits:bool,
102    /// Original FP32 gate cube derivative.
103    pub gate:bool,
104    /// Original FP32 up cube derivative.
105    pub up:bool,
106    /// Original FP32 down cube derivative.
107    pub down:bool,
108}
109/// Only requested original native derivatives; None is absence, never a synthetic zero tensor.
110#[derive(Debug)]
111pub struct MoeBackwardSelected<B:Backend> {
112    /// Requested source-token input derivative.
113    pub input:Option<FloatTensor<B>>,
114    /// Requested source-logit derivative.
115    pub logits:Option<FloatTensor<B>>,
116    /// Requested original FP32 gate cube derivative.
117    pub gate:Option<FloatTensor<B>>,
118    /// Requested original FP32 up cube derivative.
119    pub up:Option<FloatTensor<B>>,
120    /// Requested original FP32 down cube derivative.
121    pub down:Option<FloatTensor<B>>,
122}
123/// Native execution failure or an unsupported derivative of the first-order native kernels.
124#[derive(Debug)]
125pub enum MoeAutodiffError<E:fmt::Debug> {
126    /// Original native validation/launch error.
127    Native(E),
128    /// Native router/expert VJPs do not provide higher derivatives.
129    HigherDerivativeUnsupported,
130}
131impl<E:fmt::Debug> fmt::Display for MoeAutodiffError<E> {
132    fn fmt(&self,f:&mut fmt::Formatter<'_>) -> fmt::Result {
133        match self {Self::Native(error)=>write!(f,"native MoE: {error:?}"),Self::HigherDerivativeUnsupported=>f.write_str("native MoE provides first-order derivatives only")}
134    }
135}
136impl<E:fmt::Debug> core::error::Error for MoeAutodiffError<E> {}
137
138/// Optional actual device MoE extension. No dense host expert evaluation is provided.
139pub trait MoeOps:Backend {
140    /// Original native or first-order contract error.
141    type MoeError:fmt::Debug;
142    /// Actual opaque forward/dispatch state; contains original native handles, not host model copies.
143    type MoeState:Clone+Send+fmt::Debug+'static;
144    /// FP32 continuous weights for explicitly supplied integer selections. Invalid
145    /// indices retain the original bounds-safe whole-row NaN semantics.
146    fn moe_selected_weights(logits:FloatTensor<Self>,indices:IntTensor<Self>,options:MoeRouterWeightOptions) -> Result<FloatTensor<Self>,Self::MoeError>;
147    /// Original source-logits storage VJP for fixed selections; gradient weights are FP32.
148    fn moe_selected_weights_backward(logits:FloatTensor<Self>,indices:IntTensor<Self>,gradient:FloatTensor<Self>,options:MoeRouterWeightOptions)
149        -> Result<FloatTensor<Self>,Self::MoeError>;
150    /// Original top-k/group selection -> FP32 continuous weights -> native dispatch
151    /// -> original SwiGLU experts -> ordered combine. No tokens are capacity-dropped.
152    fn moe_forward(input:FloatTensor<Self>,logits:FloatTensor<Self>,correction_bias:Option<FloatTensor<Self>>,
153        gate:FloatTensor<Self>,up:FloatTensor<Self>,down:FloatTensor<Self>,options:MoeOptions) -> Result<(FloatTensor<Self>,Self::MoeState),Self::MoeError>;
154    /// Same original forward with an explicit backward-output requirement. A native
155    /// implementation may omit expert activation caches when no expert/input VJP is
156    /// required. Requesting an unavailable expert VJP from that state must error.
157    /// The compatibility default retains the original complete forward state.
158    fn moe_forward_selected(input:FloatTensor<Self>,logits:FloatTensor<Self>,correction_bias:Option<FloatTensor<Self>>,
159        gate:FloatTensor<Self>,up:FloatTensor<Self>,down:FloatTensor<Self>,options:MoeOptions,selection:MoeGradientSelection)
160        -> Result<(FloatTensor<Self>,Self::MoeState),Self::MoeError> {
161        let _=selection;Self::moe_forward(input,logits,correction_bias,gate,up,down,options)
162    }
163    /// Same original continuous weight policy and native expert output, without retained backward caches.
164    fn moe_inference(input:FloatTensor<Self>,logits:FloatTensor<Self>,correction_bias:Option<FloatTensor<Self>>,
165        gate:FloatTensor<Self>,up:FloatTensor<Self>,down:FloatTensor<Self>,options:MoeOptions) -> Result<FloatTensor<Self>,Self::MoeError>;
166    /// Actual original discrete U32 expert IDs, with no host readback or differentiable selection claim.
167    fn moe_route_indices(state:&Self::MoeState) -> IntTensor<Self>;
168    /// Original first-order input/logits/expert derivatives. Expert gradients retain
169    /// native FP32 accumulation/output; ordinary AD casts them at the parent-storage boundary.
170    fn moe_backward(state:Self::MoeState,gradient:FloatTensor<Self>) -> Result<MoeBackward<Self>,Self::MoeError>;
171    /// Explicit requested VJP outputs. Device implementations can omit unneeded native launches
172    /// and allocations; the compatibility default retains exact full-backward behavior.
173    fn moe_backward_selected(state:Self::MoeState,gradient:FloatTensor<Self>,selection:MoeGradientSelection)
174        -> Result<MoeBackwardSelected<Self>,Self::MoeError> {
175        let result=Self::moe_backward(state,gradient)?;
176        Ok(MoeBackwardSelected {input:selection.input.then_some(result.input),logits:selection.logits.then_some(result.logits),
177            gate:selection.gate.then_some(result.gate),up:selection.up.then_some(result.up),down:selection.down.then_some(result.down)})
178    }
179}