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}