use super::*;
pub trait ChunkedShardedCommunicator<B: Backend>:
ChunkedAllReduceCommunicator<B> + zero2::ShardedCommunicator<B>
{
fn reduce_scatter_float_chunked(&self, value: B::FloatTensorPrimitive, max_chunk_bytes: usize)
-> Result<B::FloatTensorPrimitive, TensorDeviceError>;
fn all_gather_float_chunked(&self, value: B::FloatTensorPrimitive, max_chunk_bytes: usize)
-> Result<B::FloatTensorPrimitive, TensorDeviceError>;
}
impl<B: Backend> ChunkedShardedCommunicator<B> for RankCommunicator<TensorDevice<B>> {
fn reduce_scatter_float_chunked(&self, value: B::FloatTensorPrimitive, max_chunk_bytes: usize)
-> Result<B::FloatTensorPrimitive, TensorDeviceError> {
RankCommunicator::reduce_scatter_float_chunked(self, value, ReduceOperation::Sum, max_chunk_bytes)
}
fn all_gather_float_chunked(&self, value: B::FloatTensorPrimitive, max_chunk_bytes: usize)
-> Result<B::FloatTensorPrimitive, TensorDeviceError> {
RankCommunicator::all_gather_float_chunked(self, value, max_chunk_bytes)
}
}
#[derive(Clone, Debug)]
pub struct BoundedShardedCommunicator<B: Backend, C: ChunkedShardedCommunicator<B>> {
reductions: BoundedAllReduceCommunicator<B, C>,
}
impl<B: Backend, C: ChunkedShardedCommunicator<B>> BoundedShardedCommunicator<B, C> {
pub fn new(inner: C, max_chunk_bytes: usize) -> Self {
Self { reductions: BoundedAllReduceCommunicator::new(inner, max_chunk_bytes) }
}
pub fn max_chunk_bytes(&self) -> usize { self.reductions.max_chunk_bytes() }
pub fn inner(&self) -> &C { self.reductions.inner() }
pub fn into_inner(self) -> C { self.reductions.into_inner() }
}
impl<B: Backend, C: ChunkedShardedCommunicator<B>> DataParallelCommunicator<B>
for BoundedShardedCommunicator<B, C>
{
fn rank(&self) -> u32 { self.reductions.rank() }
fn world_size(&self) -> u32 { self.reductions.world_size() }
fn device(&self) -> &B::Device { self.reductions.device() }
fn all_gather_bytes(&self, payload: Vec<u8>) -> Result<Vec<Vec<u8>>, DataParallelError> {
self.reductions.all_gather_bytes(payload)
}
fn broadcast_float(&self, value: B::FloatTensorPrimitive, root: u32)
-> Result<B::FloatTensorPrimitive, TensorDeviceError> {
self.reductions.broadcast_float(value, root)
}
fn broadcast_int(&self, value: B::IntTensorPrimitive, root: u32)
-> Result<B::IntTensorPrimitive, TensorDeviceError> {
self.reductions.broadcast_int(value, root)
}
fn all_reduce_float(&self, value: B::FloatTensorPrimitive, operation: ReduceOperation)
-> Result<B::FloatTensorPrimitive, TensorDeviceError> {
self.reductions.all_reduce_float(value, operation)
}
}
impl<B: Backend, C: ChunkedShardedCommunicator<B>> zero2::ShardedCommunicator<B>
for BoundedShardedCommunicator<B, C>
{
fn reduce_scatter_float(&self, value: B::FloatTensorPrimitive)
-> Result<B::FloatTensorPrimitive, TensorDeviceError> {
self.inner().reduce_scatter_float_chunked(value, self.max_chunk_bytes())
}
fn all_gather_float(&self, value: B::FloatTensorPrimitive)
-> Result<B::FloatTensorPrimitive, TensorDeviceError> {
self.inner().all_gather_float_chunked(value, self.max_chunk_bytes())
}
}