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}