Skip to main content

ruda_tensor/api/
exchange.rs

1use super::{Tensor,Int};
2use crate::{Backend,primitive::TensorPrimitive,collective::{VariableTensorCollective,VariableTensorExchange}};
3
4impl<B:Backend,const D:usize> Tensor<B,D> {
5    /// Native variable row exchange over the explicit underlying backend communicator.
6    /// Differentiable exchanges use ruda-autodiff's collective function on its original transport backend.
7    pub fn all_to_all_v<C:VariableTensorCollective<B>>(self,communicator:C,send_counts:&[usize])
8        -> Result<VariableTensorExchange<Self>,C::Error> {
9        let result=communicator.all_to_all_v_float(self.into_primitive().tensor(),send_counts)?;
10        Ok(VariableTensorExchange {value:Tensor::from_primitive(TensorPrimitive::Float(result.value)),receive_counts:result.receive_counts})
11    }
12}
13impl<B:Backend,const D:usize> Tensor<B,D,Int> {
14    /// Actual U8/U32/I32/I64 source-rank row exchange without numerical widening or floating conversion.
15    pub fn all_to_all_v_int<C:VariableTensorCollective<B>>(self,communicator:C,send_counts:&[usize])
16        -> Result<VariableTensorExchange<Self>,C::Error> {
17        let result=communicator.all_to_all_v_int(self.into_primitive(),send_counts)?;
18        Ok(VariableTensorExchange {value:Tensor::from_primitive(result.value),receive_counts:result.receive_counts})
19    }
20}