Skip to main content

ruda_tensor/
moe_exchange.rs

1//! Native routing/dispatch, received expert rows and ordered combine for expert-parallel graphs.
2use alloc::vec::Vec;
3use core::fmt::Debug;
4use crate::{Backend,moe::{MoeOps,MoeOptions,MoeExpertStrategy,MoeCombineGradientStrategy},tensor::{FloatTensor,IntTensor}};
5
6/// Actual original discrete dispatch and continuous weights, before any cross-rank exchange.
7#[derive(Debug)]
8pub struct MoeDispatched<B:MoeDispatchOps> {
9    /// Original expert-sorted source row copies, one row per selected assignment.
10    pub values:FloatTensor<B>,
11    /// Original FP32 continuous `[tokens,top_k]` weights, with its actual logits derivative.
12    pub weights:FloatTensor<B>,
13    /// Original selected U32 `[tokens,top_k]` expert IDs.
14    pub selected_experts:IntTensor<B>,
15    /// Original sorted U32 `[assignments]` global expert IDs, aligned with values.
16    pub row_experts:IntTensor<B>,
17    /// Actual private original native mappings and router VJP state.
18    pub state:B::MoeDispatchState,
19}
20/// Requested original combine VJP outputs.
21#[derive(Clone,Copy,Debug,PartialEq,Eq)]
22pub struct MoeCombineSelection {
23    /// Original expert-row output derivative.
24    pub experts:bool,
25    /// Original FP32 continuous routing-weight derivative.
26    pub weights:bool,
27}
28/// Actual optional original combine derivatives, not placeholder zero tensors.
29#[derive(Debug)]
30pub struct MoeCombineBackward<B:Backend> {
31    /// Requested derivative of source expert-sorted row values.
32    pub experts:Option<FloatTensor<B>>,
33    /// Requested FP32 derivative of `[tokens,top_k]` continuous weights.
34    pub weights:Option<FloatTensor<B>>,
35}
36/// Independent native source-side dispatch and combine, retaining discrete selection semantics.
37pub trait MoeDispatchOps:MoeOps {
38    /// Original valid private native row mappings, with no unchecked imported offsets.
39    type MoeDispatchState:Clone+Send+Debug+'static;
40    /// Original selection and continuous weights followed by device COPY dispatch.
41    fn moe_dispatch(input:FloatTensor<Self>,logits:FloatTensor<Self>,bias:Option<FloatTensor<Self>>,options:MoeOptions)
42        -> Result<MoeDispatched<Self>,Self::MoeError>;
43    /// Actual row counts per caller-declared contiguous global expert range.
44    /// Reads only expert-prefix coordination metadata, not activation/weight values.
45    fn moe_dispatch_counts(state:&Self::MoeDispatchState,expert_prefix:&[usize]) -> Result<Vec<usize>,Self::MoeError>;
46    /// COPY VJP sums selected expert-row seeds without multiplying routing weights again.
47    fn moe_dispatch_backward(state:Self::MoeDispatchState,gradient:FloatTensor<Self>) -> Result<FloatTensor<Self>,Self::MoeError>;
48    /// Original fixed-selection source-logit VJP of FP32 continuous weights.
49    fn moe_dispatch_weights_backward(state:Self::MoeDispatchState,gradient:FloatTensor<Self>) -> Result<FloatTensor<Self>,Self::MoeError>;
50    /// Original ascending expert-ID combine after actual expert rows have returned to their source.
51    fn moe_combine(state:Self::MoeDispatchState,expert_values:FloatTensor<Self>,weights:FloatTensor<Self>,backward_strategy:MoeCombineGradientStrategy)
52        -> Result<FloatTensor<Self>,Self::MoeError>;
53    /// Original combine kernels with only requested actual expert/weight derivatives.
54    fn moe_combine_backward(state:Self::MoeDispatchState,expert_values:FloatTensor<Self>,weights:FloatTensor<Self>,gradient:FloatTensor<Self>,
55        strategy:MoeCombineGradientStrategy,selection:MoeCombineSelection) -> Result<MoeCombineBackward<Self>,Self::MoeError>;
56}
57/// Explicit original local expert range and independent forward/backward strategies.
58#[derive(Clone,Copy,Debug,PartialEq,Eq)]
59pub struct MoeReceivedOptions {
60    /// First global expert ID owned by these actual local cubes; count comes from their original shape.
61    pub expert_start:usize,
62    /// Original actual segmented forward policy.
63    pub forward:MoeExpertStrategy,
64    /// Original actual independent segmented backward policy.
65    pub backward:MoeExpertStrategy,
66}
67/// Actual original received-expert derivatives requested by native callers or tracked parents.
68#[derive(Clone,Copy,Debug,PartialEq,Eq)]
69pub struct MoeReceivedSelection {
70    /// Real upstream activation derivative even when every expert matrix is frozen.
71    pub input:bool,
72    /// Original FP32 local gate-cube derivative.
73    pub gate:bool,
74    /// Original FP32 local up-cube derivative.
75    pub up:bool,
76    /// Original FP32 local down-cube derivative.
77    pub down:bool,
78}
79/// Actual original optional received-expert VJP outputs in the original receive order.
80#[derive(Debug)]
81pub struct MoeReceivedBackward<B:Backend> {
82    /// Original unsorted receive-axis input derivative.
83    pub input:Option<FloatTensor<B>>,
84    /// Original FP32 local gate cube derivative.
85    pub gate:Option<FloatTensor<B>>,
86    /// Original FP32 local up cube derivative.
87    pub up:Option<FloatTensor<B>>,
88    /// Original FP32 local down cube derivative.
89    pub down:Option<FloatTensor<B>>,
90}
91/// Native local expert computation on actual transported assignments, without routing-weight application.
92pub trait MoeReceivedOps:MoeOps {
93    /// Private actual validated grouping, inverse permutation and original expert VJP cache.
94    type MoeReceivedState:Clone+Send+Debug+'static;
95    /// Sort received U32 assignments locally, execute original experts and restore receive order.
96    /// Empty expert owners participate with real empty local cubes and zero received rows.
97    fn moe_received_forward(input:FloatTensor<Self>,global_ids:IntTensor<Self>,gate:FloatTensor<Self>,up:FloatTensor<Self>,down:FloatTensor<Self>,
98        options:MoeReceivedOptions,selection:MoeReceivedSelection) -> Result<(FloatTensor<Self>,Self::MoeReceivedState),Self::MoeError>;
99    /// Same original values without retained caches unless actual AD parents require derivatives.
100    fn moe_received_inference(input:FloatTensor<Self>,global_ids:IntTensor<Self>,gate:FloatTensor<Self>,up:FloatTensor<Self>,down:FloatTensor<Self>,options:MoeReceivedOptions)
101        -> Result<FloatTensor<Self>,Self::MoeError> {
102        Self::moe_received_forward(input,global_ids,gate,up,down,options,MoeReceivedSelection {input:false,gate:false,up:false,down:false}).map(|(output,_)|output)
103    }
104    /// Original expert VJP and inverse COPY mapping; local cube gradients retain FP32.
105    fn moe_received_backward(state:Self::MoeReceivedState,gradient:FloatTensor<Self>,selection:MoeReceivedSelection)
106        -> Result<MoeReceivedBackward<Self>,Self::MoeError>;
107}