use super::{TensorDevice, TensorDeviceError};
use crate::{ReduceOperation, rank::{ElementType, ReductionOperation, communicator::RankCommunicator,
device_collective::{NativeChunkPlan, NativeChunkProgress}}};
use ruda_tensor::{Backend, DType, Shape, Slice, TensorMetadata};
impl<B: Backend> RankCommunicator<TensorDevice<B>> {
pub fn all_reduce_float_chunked(
&self, value: B::FloatTensorPrimitive, operation: ReduceOperation, max_chunk_bytes: usize,
) -> Result<B::FloatTensorPrimitive, TensorDeviceError> {
self.all_reduce_float_chunked_with_progress(value, operation, max_chunk_bytes, |_| {})
}
pub fn all_reduce_float_chunked_with_progress<F: FnMut(NativeChunkProgress)>(
&self, value: B::FloatTensorPrimitive, operation: ReduceOperation,
max_chunk_bytes: usize, mut progress: F,
) -> Result<B::FloatTensorPrimitive, TensorDeviceError> {
let element_type = match value.dtype() {
DType::F32 => ElementType::F32,
DType::F16 => ElementType::F16,
DType::BF16 => ElementType::BF16,
dtype => return Err(TensorDeviceError::UnsupportedDType(dtype)),
};
if B::float_device(&value) != *self.execution().device() { return Err(TensorDeviceError::DeviceMismatch); }
let shape = value.shape();
let plan = self.reduction_chunk_agreement(&shape, element_type, operation as u64, max_chunk_bytes)?;
let mut state = reduction_progress(plan)?;
progress(state);
let mut value = B::float_reshape(value, Shape::new([plan.elements()]));
while state.completed_elements < plan.elements() {
let end = state.completed_elements + plan.chunk_elements().min(plan.elements() - state.completed_elements);
let range = Slice::from(state.completed_elements..end);
let chunk = B::float_slice(value.clone(), core::slice::from_ref(&range));
let reduced = self.all_reduce_float(chunk, operation)?;
value = B::float_slice_assign(value, core::slice::from_ref(&range), reduced);
B::sync(self.execution().device())?;
state.completed_elements = end;
state.completed_chunks += 1;
progress(state);
}
Ok(B::float_reshape(value, shape))
}
pub fn all_reduce_int_chunked(
&self, value: B::IntTensorPrimitive, operation: ReductionOperation, max_chunk_bytes: usize,
) -> Result<B::IntTensorPrimitive, TensorDeviceError> {
self.all_reduce_int_chunked_with_progress(value, operation, max_chunk_bytes, |_| {})
}
pub fn all_reduce_int_chunked_with_progress<F: FnMut(NativeChunkProgress)>(
&self, value: B::IntTensorPrimitive, operation: ReductionOperation,
max_chunk_bytes: usize, mut progress: F,
) -> Result<B::IntTensorPrimitive, TensorDeviceError> {
let element_type = match value.dtype() {
DType::I32 => ElementType::I32,
DType::I64 => ElementType::I64,
dtype => return Err(TensorDeviceError::UnsupportedDType(dtype)),
};
if B::int_device(&value) != *self.execution().device() { return Err(TensorDeviceError::DeviceMismatch); }
let shape = value.shape();
let plan = self.reduction_chunk_agreement(&shape, element_type, operation as u64, max_chunk_bytes)?;
let mut state = reduction_progress(plan)?;
progress(state);
let mut value = B::int_reshape(value, Shape::new([plan.elements()]));
while state.completed_elements < plan.elements() {
let end = state.completed_elements + plan.chunk_elements().min(plan.elements() - state.completed_elements);
let range = Slice::from(state.completed_elements..end);
let chunk = B::int_slice(value.clone(), core::slice::from_ref(&range));
let reduced = self.all_reduce_int(chunk, operation)?;
value = B::int_slice_assign(value, core::slice::from_ref(&range), reduced);
B::sync(self.execution().device())?;
state.completed_elements = end;
state.completed_chunks += 1;
progress(state);
}
Ok(B::int_reshape(value, shape))
}
fn reduction_chunk_agreement(
&self, shape: &Shape, element_type: ElementType, operation: u64, max_chunk_bytes: usize,
) -> Result<NativeChunkPlan, TensorDeviceError> {
self.native_chunk_agreement(shape, element_type, operation, max_chunk_bytes)?;
let elements = if shape.contains(&0) { 0 } else { shape.iter().try_fold(1usize, |count, dimension| count.checked_mul(*dimension))
.ok_or(TensorDeviceError::InvalidBuffer("chunked reduction shape element count overflow"))?
};
self.native_message_plan(elements, element_type, max_chunk_bytes)
}
pub(super) fn native_chunk_agreement(
&self, shape: &Shape, element_type: ElementType, operation: u64, max_chunk_bytes: usize,
) -> Result<(), TensorDeviceError> {
let mut fields = vec![element_type as u64, operation, max_chunk_bytes as u64, self.world_size() as u64, shape.num_dims() as u64];
fields.extend(shape.iter().map(|dimension| *dimension as u64));
let payload = fields.iter().flat_map(|field| field.to_le_bytes()).collect::<Vec<_>>();
let (requests, _) = self.host().all_gather_host_staged(ElementType::U64, fields.len(), payload.clone())?;
if requests.chunks_exact(payload.len()).any(|other| other != payload.as_slice()) {
return Err(TensorDeviceError::InvalidOperation("chunked reduction geometry/type/operation/budget differs across ranks"));
}
Ok(())
}
pub(super) fn native_message_plan(
&self, elements: usize, element_type: ElementType, max_chunk_bytes: usize,
) -> Result<NativeChunkPlan, TensorDeviceError> {
let world = self.world_size() as usize;
if world == 0 { return Err(TensorDeviceError::InvalidOperation("native chunk communicator world must be positive")); }
let per_rank_bytes = if elements == 0 { element_type.byte_width() } else { max_chunk_bytes / world };
Ok(NativeChunkPlan::new(elements, element_type.byte_width(), per_rank_bytes)?)
}
}
pub(super) fn reduction_progress(plan: NativeChunkPlan) -> Result<NativeChunkProgress, TensorDeviceError> {
Ok(NativeChunkProgress { completed_elements: 0, total_elements: plan.elements(),
completed_chunks: 0, total_chunks: plan.chunk_count()?, restored_elements: 0 })
}