use super::{Primitive, TensorDevice, TensorDeviceError};
use super::bounded_reduction::reduction_progress;
use crate::{ReduceOperation, rank::{ElementType, ReductionOperation, communicator::RankCommunicator,
device_collective::NativeChunkProgress}};
use ruda_core::tensor::collective::CollectiveShape;
use ruda_tensor::{Backend, DType, Shape, TensorMetadata};
#[derive(Clone, Copy)]
enum ScatterOperation { Float(ReduceOperation), Integer(ReductionOperation) }
impl<B: Backend> RankCommunicator<TensorDevice<B>> {
pub fn all_gather_float_chunked(&self, value: B::FloatTensorPrimitive, max_chunk_bytes: usize)
-> Result<B::FloatTensorPrimitive, TensorDeviceError> {
self.all_gather_float_chunked_with_progress(value, max_chunk_bytes, |_| {})
}
pub fn all_gather_float_chunked_with_progress<F: FnMut(NativeChunkProgress)>(
&self, value: B::FloatTensorPrimitive, max_chunk_bytes: usize, mut progress: F,
) -> Result<B::FloatTensorPrimitive, TensorDeviceError> {
match self.native_gather_chunked(Primitive::Float(value), max_chunk_bytes, &mut progress)? {
Primitive::Float(value) => Ok(value), Primitive::Int(_) => unreachable!("original floating gather kind"),
}
}
pub fn all_gather_int_chunked(&self, value: B::IntTensorPrimitive, max_chunk_bytes: usize)
-> Result<B::IntTensorPrimitive, TensorDeviceError> {
self.all_gather_int_chunked_with_progress(value, max_chunk_bytes, |_| {})
}
pub fn all_gather_int_chunked_with_progress<F: FnMut(NativeChunkProgress)>(
&self, value: B::IntTensorPrimitive, max_chunk_bytes: usize, mut progress: F,
) -> Result<B::IntTensorPrimitive, TensorDeviceError> {
match self.native_gather_chunked(Primitive::Int(value), max_chunk_bytes, &mut progress)? {
Primitive::Int(value) => Ok(value), Primitive::Float(_) => unreachable!("original integer gather kind"),
}
}
pub fn reduce_scatter_float_chunked(
&self, value: B::FloatTensorPrimitive, operation: ReduceOperation, max_chunk_bytes: usize,
) -> Result<B::FloatTensorPrimitive, TensorDeviceError> {
self.reduce_scatter_float_chunked_with_progress(value, operation, max_chunk_bytes, |_| {})
}
pub fn reduce_scatter_float_chunked_with_progress<F: FnMut(NativeChunkProgress)>(
&self, value: B::FloatTensorPrimitive, operation: ReduceOperation,
max_chunk_bytes: usize, mut progress: F,
) -> Result<B::FloatTensorPrimitive, TensorDeviceError> {
match self.native_scatter_chunked(Primitive::Float(value), ScatterOperation::Float(operation), max_chunk_bytes, &mut progress)? {
Primitive::Float(value) => Ok(value), Primitive::Int(_) => unreachable!("original floating scatter kind"),
}
}
pub fn reduce_scatter_int_chunked(
&self, value: B::IntTensorPrimitive, operation: ReductionOperation, max_chunk_bytes: usize,
) -> Result<B::IntTensorPrimitive, TensorDeviceError> {
self.reduce_scatter_int_chunked_with_progress(value, operation, max_chunk_bytes, |_| {})
}
pub fn reduce_scatter_int_chunked_with_progress<F: FnMut(NativeChunkProgress)>(
&self, value: B::IntTensorPrimitive, operation: ReductionOperation,
max_chunk_bytes: usize, mut progress: F,
) -> Result<B::IntTensorPrimitive, TensorDeviceError> {
match self.native_scatter_chunked(Primitive::Int(value), ScatterOperation::Integer(operation), max_chunk_bytes, &mut progress)? {
Primitive::Int(value) => Ok(value), Primitive::Float(_) => unreachable!("original integer scatter kind"),
}
}
fn native_gather_chunked<F: FnMut(NativeChunkProgress)>(
&self, value: Primitive<B>, max_chunk_bytes: usize, progress: &mut F,
) -> Result<Primitive<B>, TensorDeviceError> {
let (shape, dtype) = self.native_sharded_metadata(&value)?;
let element_type = native_element_type(dtype)?;
self.native_chunk_agreement(&shape, element_type, 0x100, max_chunk_bytes)?;
let layout = CollectiveShape::all_gather(shape, self.world_size() as usize)
.map_err(|error| TensorDeviceError::Data(format!("{error:?}")))?;
let plan = self.native_message_plan(layout.input_elements, element_type, max_chunk_bytes)?;
let value = native_reshape::<B>(value, Shape::new([layout.input_elements]));
let mut output = self.native_empty_like(&value, layout.output_elements, dtype);
let mut state = reduction_progress(plan)?;
progress(state);
while state.completed_elements < plan.elements() {
let length = plan.chunk_elements().min(plan.elements() - state.completed_elements);
let gathered = match value.clone().slice(state.completed_elements..state.completed_elements + length) {
Primitive::Float(value) => Primitive::Float(self.all_gather_float(value)?),
Primitive::Int(value) => Primitive::Int(self.all_gather_int(value)?),
};
for owner in 0..self.world_size() as usize {
let start = owner * layout.input_elements + state.completed_elements;
let piece = gathered.clone().slice(owner * length..(owner + 1) * length);
output = output.assign(start..start + length, piece);
}
B::sync(self.execution().device())?;
state.completed_elements += length;
state.completed_chunks += 1;
progress(state);
}
Ok(native_reshape::<B>(output, layout.output))
}
fn native_scatter_chunked<F: FnMut(NativeChunkProgress)>(
&self, value: Primitive<B>, operation: ScatterOperation, max_chunk_bytes: usize, progress: &mut F,
) -> Result<Primitive<B>, TensorDeviceError> {
let (shape, dtype) = self.native_sharded_metadata(&value)?;
let element_type = native_element_type(dtype)?;
let tag = match operation {
ScatterOperation::Float(operation) => 0x200 + operation as u64,
ScatterOperation::Integer(operation) => 0x200 + operation as u64,
};
self.native_chunk_agreement(&shape, element_type, tag, max_chunk_bytes)?;
let layout = CollectiveShape::reduce_scatter(shape, self.world_size() as usize)
.map_err(|error| TensorDeviceError::Data(format!("{error:?}")))?;
if matches!(dtype, DType::U8 | DType::U32) { return Err(TensorDeviceError::UnsupportedDType(dtype)); }
let plan = self.native_message_plan(layout.output_elements, element_type, max_chunk_bytes)?;
let value = native_reshape::<B>(value, Shape::new([layout.input_elements]));
let mut output = self.native_empty_like(&value, layout.output_elements, dtype);
let mut state = reduction_progress(plan)?;
progress(state);
while state.completed_elements < plan.elements() {
let length = plan.chunk_elements().min(plan.elements() - state.completed_elements);
let chunks = (0..self.world_size() as usize).map(|owner| {
let start = owner * layout.output_elements + state.completed_elements;
value.clone().slice(start..start + length)
}).collect::<Vec<_>>();
let reduced = match operation {
ScatterOperation::Float(operation) => {
let chunks = chunks.into_iter().map(|piece| match piece {
Primitive::Float(value) => value, Primitive::Int(_) => unreachable!("original floating scatter input"),
}).collect();
Primitive::Float(self.reduce_scatter_float(B::float_cat(chunks, 0), operation)?)
}
ScatterOperation::Integer(operation) => {
let chunks = chunks.into_iter().map(|piece| match piece {
Primitive::Int(value) => value, Primitive::Float(_) => unreachable!("original integer scatter input"),
}).collect();
Primitive::Int(self.reduce_scatter_int(B::int_cat(chunks, 0), operation)?)
}
};
output = output.assign(state.completed_elements..state.completed_elements + length, reduced);
B::sync(self.execution().device())?;
state.completed_elements += length;
state.completed_chunks += 1;
progress(state);
}
Ok(native_reshape::<B>(output, layout.output))
}
pub(super) fn native_sharded_metadata(&self, value: &Primitive<B>) -> Result<(Shape, DType), TensorDeviceError> {
let (shape, dtype, device) = match value {
Primitive::Float(value) => (value.shape(), value.dtype(), B::float_device(value)),
Primitive::Int(value) => (value.shape(), value.dtype(), B::int_device(value)),
};
if device != *self.execution().device() { return Err(TensorDeviceError::DeviceMismatch); }
Ok((shape, dtype))
}
pub(super) fn native_empty_like(&self, source: &Primitive<B>, elements: usize, dtype: DType) -> Primitive<B> {
let shape = Shape::new([elements]);
match source {
Primitive::Float(_) => Primitive::Float(B::float_empty(shape, self.execution().device(), dtype.into())),
Primitive::Int(_) => Primitive::Int(B::int_empty(shape, self.execution().device(), dtype.into())),
}
}
}
pub(super) fn native_reshape<B: Backend>(value: Primitive<B>, shape: Shape) -> Primitive<B> {
match value { Primitive::Float(value) => Primitive::Float(B::float_reshape(value, shape)),
Primitive::Int(value) => Primitive::Int(B::int_reshape(value, shape)) }
}
pub(super) fn native_element_type(dtype: DType) -> Result<ElementType, TensorDeviceError> {
Ok(match dtype {
DType::F32 => ElementType::F32, DType::F16 => ElementType::F16, DType::BF16 => ElementType::BF16,
DType::U8 => ElementType::U8, DType::U32 => ElementType::U32, DType::I32 => ElementType::I32, DType::I64 => ElementType::I64,
dtype => return Err(TensorDeviceError::UnsupportedDType(dtype)),
})
}