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}