use super::{TensorDevice, TensorDeviceError};
use crate::{ReduceOperation, rank::communicator::RankCommunicator};
use ruda_tensor::{Backend, collective::{TensorCollective, ReplicatedTensorCollective,
BroadcastTensorCollective, IntegerTensorCollective, VariableTensorCollective, VariableTensorExchange}};
#[derive(Clone, Debug)]
pub struct BoundedTensorCommunicator<B: Backend> {
inner: RankCommunicator<TensorDevice<B>>,
max_chunk_bytes: usize,
}
impl<B: Backend> BoundedTensorCommunicator<B> {
pub fn new(inner: RankCommunicator<TensorDevice<B>>, max_chunk_bytes: usize) -> Self {
Self { inner, max_chunk_bytes }
}
pub fn max_chunk_bytes(&self) -> usize { self.max_chunk_bytes }
pub fn inner(&self) -> &RankCommunicator<TensorDevice<B>> { &self.inner }
pub fn into_inner(self) -> RankCommunicator<TensorDevice<B>> { self.inner }
pub fn broadcast_int(&self, value: B::IntTensorPrimitive, root: u32)
-> Result<B::IntTensorPrimitive, TensorDeviceError> {
self.inner.broadcast_int_chunked(value, root, self.max_chunk_bytes)
}
}
impl<B: Backend> TensorCollective<B> for BoundedTensorCommunicator<B> {
type Error = TensorDeviceError;
fn world_size(&self) -> u32 { self.inner.world_size() }
fn all_gather_float(&self, value: B::FloatTensorPrimitive) -> Result<B::FloatTensorPrimitive, Self::Error> {
self.inner.all_gather_float_chunked(value, self.max_chunk_bytes)
}
fn reduce_scatter_sum(&self, value: B::FloatTensorPrimitive) -> Result<B::FloatTensorPrimitive, Self::Error> {
self.inner.reduce_scatter_float_chunked(value, ReduceOperation::Sum, self.max_chunk_bytes)
}
}
impl<B: Backend> ReplicatedTensorCollective<B> for BoundedTensorCommunicator<B> {
fn all_reduce_sum(&self, value: B::FloatTensorPrimitive) -> Result<B::FloatTensorPrimitive, Self::Error> {
self.inner.all_reduce_float_chunked(value, ReduceOperation::Sum, self.max_chunk_bytes)
}
}
impl<B: Backend> BroadcastTensorCollective<B> for BoundedTensorCommunicator<B> {
fn rank(&self) -> u32 { self.inner.rank() }
fn broadcast_float(&self, value: B::FloatTensorPrimitive, root: u32) -> Result<B::FloatTensorPrimitive, Self::Error> {
self.inner.broadcast_float_chunked(value, root, self.max_chunk_bytes)
}
}
impl<B: Backend> IntegerTensorCollective<B> for BoundedTensorCommunicator<B> {
fn all_gather_int(&self, value: B::IntTensorPrimitive) -> Result<B::IntTensorPrimitive, Self::Error> {
self.inner.all_gather_int_chunked(value, self.max_chunk_bytes)
}
}
impl<B: Backend> VariableTensorCollective<B> for BoundedTensorCommunicator<B> {
fn all_to_all_v_float(&self, value: B::FloatTensorPrimitive, send_counts: &[usize])
-> Result<VariableTensorExchange<B::FloatTensorPrimitive>, Self::Error> {
self.inner.all_to_all_v_float_chunked(value, send_counts, self.max_chunk_bytes)
}
fn all_to_all_v_int(&self, value: B::IntTensorPrimitive, send_counts: &[usize])
-> Result<VariableTensorExchange<B::IntTensorPrimitive>, Self::Error> {
self.inner.all_to_all_v_int_chunked(value, send_counts, self.max_chunk_bytes)
}
}