Skip to main content

ruda_tensor/api/
moe_exchange.rs

1use super::{Tensor,Int};
2use crate::{TensorPrimitive,moe::{MoeOptions,MoeCombineGradientStrategy},moe_exchange::{MoeDispatchOps,MoeReceivedOps,MoeReceivedOptions}};
3
4/// Actual typed native source-side dispatched rows and differentiable FP32 continuous weights.
5#[derive(Debug)]
6pub struct NativeMoeDispatched<B:MoeDispatchOps> {
7    /// Original expert-sorted activation COPY rows.
8    pub values:Tensor<B,2>,
9    /// Original FP32 continuous weights with the original source logits graph.
10    pub weights:Tensor<B,2>,
11    /// Original selected native U32 IDs per source token.
12    pub selected_experts:Tensor<B,2,Int>,
13    /// Original native U32 expert IDs aligned with sorted assignment rows.
14    pub row_experts:Tensor<B,1,Int>,
15    /// Original valid private native dispatch and router metadata.
16    pub state:B::MoeDispatchState,
17}
18/// Run actual original device routing and COPY dispatch before expert-parallel transport.
19pub fn dispatch_moe<B:MoeDispatchOps>(input:Tensor<B,2>,logits:Tensor<B,2>,bias:Option<Tensor<B,1>>,options:MoeOptions)
20    -> Result<NativeMoeDispatched<B>,B::MoeError> {
21    let result=B::moe_dispatch(input.into_primitive().tensor(),logits.into_primitive().tensor(),bias.map(|value|value.into_primitive().tensor()),options)?;
22    Ok(NativeMoeDispatched {values:Tensor::from_primitive(TensorPrimitive::Float(result.values)),weights:Tensor::from_primitive(TensorPrimitive::Float(result.weights)),
23        selected_experts:Tensor::from_primitive(result.selected_experts),row_experts:Tensor::from_primitive(result.row_experts),state:result.state})
24}
25/// Apply original expert-ID-ordered combine after actual rows have returned to their original source.
26pub fn combine_moe<B:MoeDispatchOps>(state:B::MoeDispatchState,expert_values:Tensor<B,2>,weights:Tensor<B,2>,strategy:MoeCombineGradientStrategy)
27    -> Result<Tensor<B,2>,B::MoeError> {
28    B::moe_combine(state,expert_values.into_primitive().tensor(),weights.into_primitive().tensor(),strategy)
29        .map(|value|Tensor::from_primitive(TensorPrimitive::Float(value)))
30}
31/// Execute only the actual local expert cubes and restore original source-rank receive order.
32/// Native inference retains no backward cache; actual AD parents retain their required VJPs.
33pub fn received_moe_experts<B:MoeReceivedOps>(input:Tensor<B,2>,global_ids:Tensor<B,1,Int>,gate:Tensor<B,3>,up:Tensor<B,3>,down:Tensor<B,3>,options:MoeReceivedOptions)
34    -> Result<Tensor<B,2>,B::MoeError> {
35    B::moe_received_inference(input.into_primitive().tensor(),global_ids.into_primitive(),gate.into_primitive().tensor(),up.into_primitive().tensor(),down.into_primitive().tensor(),options)
36        .map(|value|Tensor::from_primitive(TensorPrimitive::Float(value)))
37}