Skip to main content

ruda_tensor/
collective.rs

1//! Explicit rank communicators for floating tensor collectives.
2use crate::{Backend, tensor::{FloatTensor, IntTensor}};
3use core::fmt::Debug;
4use alloc::vec::Vec;
5
6/// Matching rank-ordered collectives used by differentiable sharded tensors.
7/// All participating ranks must enter forward and backward in the same order.
8pub trait TensorCollective<B: Backend>: Clone + Debug + Send + 'static {
9    /// Transport or tensor-contract failure.
10    type Error: Debug;
11    /// Number of participating ranks, greater than zero.
12    fn world_size(&self) -> u32;
13    /// Optional typed AD execution context carried by an explicit communicator wrapper.
14    /// Native transports and ordinary AD calls retain no context by default.
15    fn autodiff_context(&self) -> Option<&(dyn core::any::Any+Send+Sync)> {None}
16    /// Gather equal leading-axis shards in rank order, retaining the input dtype.
17    fn all_gather_float(&self, value: FloatTensor<B>) -> Result<FloatTensor<B>, Self::Error>;
18    /// Sum tensors across ranks and return this rank's equal leading-axis shard.
19    fn reduce_scatter_sum(&self, value: FloatTensor<B>) -> Result<FloatTensor<B>, Self::Error>;
20}
21
22/// Replicated reductions in addition to leading-axis shard collectives.
23pub trait ReplicatedTensorCollective<B: Backend>: TensorCollective<B> {
24    /// Sum corresponding tensor elements across ranks, retaining shape and dtype.
25    fn all_reduce_sum(&self, value: FloatTensor<B>) -> Result<FloatTensor<B>, Self::Error>;
26}
27
28/// Root-owned broadcasts with replicated reductions for their backward pass.
29pub trait BroadcastTensorCollective<B: Backend>: ReplicatedTensorCollective<B> {
30    /// This communicator's rank.
31    fn rank(&self) -> u32;
32    /// Broadcast root's tensor while retaining each rank's input snapshot.
33    fn broadcast_float(
34        &self,
35        value: FloatTensor<B>,
36        root: u32,
37    ) -> Result<FloatTensor<B>, Self::Error>;
38}
39
40/// Exact packed-integer storage transport in addition to floating rank collectives.
41/// Immutable packed parameters do not acquire a floating surrogate or derivative.
42pub trait IntegerTensorCollective<B: Backend>: BroadcastTensorCollective<B> {
43    /// Rank-ordered leading-axis gather retaining the original integer dtype and
44    /// bit patterns. Implementations must not cast packed words through floating point.
45    fn all_gather_int(&self, value: IntTensor<B>) -> Result<IntTensor<B>, Self::Error>;
46}
47
48/// Actual variable leading-axis exchange output in source-rank order.
49#[derive(Debug)]
50pub struct VariableTensorExchange<T> {
51    /// Original storage tensor with unchanged trailing axes and dtype.
52    pub value:T,
53    /// Actual received ROW counts from each source rank, not byte or scalar-element counts.
54    pub receive_counts:Vec<usize>,
55}
56/// Variable row exchange for routed/sharded graphs through an explicit original communicator.
57/// Every rank enters even with zero rows. Input blocks are in destination-rank
58/// order; received blocks are in source-rank order. Trailing axes must be nonzero
59/// and match across peers; no implicit expert-to-rank assignment is defined here.
60pub trait VariableTensorCollective<B:Backend>:IntegerTensorCollective<B> {
61    /// Exchange actual floating row blocks, retaining original floating storage.
62    fn all_to_all_v_float(&self,value:FloatTensor<B>,send_counts:&[usize]) -> Result<VariableTensorExchange<FloatTensor<B>>,Self::Error>;
63    /// Exchange original U8/U32/I32/I64 row blocks without a floating surrogate.
64    fn all_to_all_v_int(&self,value:IntTensor<B>,send_counts:&[usize]) -> Result<VariableTensorExchange<IntTensor<B>>,Self::Error>;
65}