use crate::{Backend, tensor::{FloatTensor, IntTensor}};
use core::fmt::Debug;
use alloc::vec::Vec;
pub trait TensorCollective<B: Backend>: Clone + Debug + Send + 'static {
type Error: Debug;
fn world_size(&self) -> u32;
fn autodiff_context(&self) -> Option<&(dyn core::any::Any+Send+Sync)> {None}
fn all_gather_float(&self, value: FloatTensor<B>) -> Result<FloatTensor<B>, Self::Error>;
fn reduce_scatter_sum(&self, value: FloatTensor<B>) -> Result<FloatTensor<B>, Self::Error>;
}
pub trait ReplicatedTensorCollective<B: Backend>: TensorCollective<B> {
fn all_reduce_sum(&self, value: FloatTensor<B>) -> Result<FloatTensor<B>, Self::Error>;
}
pub trait BroadcastTensorCollective<B: Backend>: ReplicatedTensorCollective<B> {
fn rank(&self) -> u32;
fn broadcast_float(
&self,
value: FloatTensor<B>,
root: u32,
) -> Result<FloatTensor<B>, Self::Error>;
}
pub trait IntegerTensorCollective<B: Backend>: BroadcastTensorCollective<B> {
fn all_gather_int(&self, value: IntTensor<B>) -> Result<IntTensor<B>, Self::Error>;
}
#[derive(Debug)]
pub struct VariableTensorExchange<T> {
pub value:T,
pub receive_counts:Vec<usize>,
}
pub trait VariableTensorCollective<B:Backend>:IntegerTensorCollective<B> {
fn all_to_all_v_float(&self,value:FloatTensor<B>,send_counts:&[usize]) -> Result<VariableTensorExchange<FloatTensor<B>>,Self::Error>;
fn all_to_all_v_int(&self,value:IntTensor<B>,send_counts:&[usize]) -> Result<VariableTensorExchange<IntTensor<B>>,Self::Error>;
}