use super::{Primitive, TensorDevice, TensorDeviceError};
use super::bounded_sharded::{native_reshape, native_element_type};
use crate::rank::{ElementType, communicator::RankCommunicator, device_collective::NativeChunkPlan};
use ruda_tensor::{Backend, Shape, collective::{VariableTensorCollective, VariableTensorExchange}};
use serde::{Deserialize, Serialize};
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct RowExchangeProgress {
pub completed_rounds: usize,
pub total_rounds: usize,
pub sent_rows: usize,
pub total_send_rows: usize,
pub received_rows: usize,
pub total_receive_rows: usize,
}
impl<B: Backend> RankCommunicator<TensorDevice<B>> {
pub fn all_to_all_v_float_chunked(
&self, value: B::FloatTensorPrimitive, send_counts: &[usize], max_chunk_bytes: usize,
) -> Result<VariableTensorExchange<B::FloatTensorPrimitive>, TensorDeviceError> {
self.all_to_all_v_float_chunked_with_progress(value, send_counts, max_chunk_bytes, |_| {})
}
pub fn all_to_all_v_float_chunked_with_progress<F: FnMut(RowExchangeProgress)>(
&self, value: B::FloatTensorPrimitive, send_counts: &[usize], max_chunk_bytes: usize, mut progress: F,
) -> Result<VariableTensorExchange<B::FloatTensorPrimitive>, TensorDeviceError> {
let (value, receive_counts) = self.native_rows_chunked(Primitive::Float(value), send_counts, max_chunk_bytes, &mut progress)?;
match value { Primitive::Float(value) => Ok(VariableTensorExchange { value, receive_counts }),
Primitive::Int(_) => unreachable!("original floating variable-row kind") }
}
pub fn all_to_all_v_int_chunked(
&self, value: B::IntTensorPrimitive, send_counts: &[usize], max_chunk_bytes: usize,
) -> Result<VariableTensorExchange<B::IntTensorPrimitive>, TensorDeviceError> {
self.all_to_all_v_int_chunked_with_progress(value, send_counts, max_chunk_bytes, |_| {})
}
pub fn all_to_all_v_int_chunked_with_progress<F: FnMut(RowExchangeProgress)>(
&self, value: B::IntTensorPrimitive, send_counts: &[usize], max_chunk_bytes: usize, mut progress: F,
) -> Result<VariableTensorExchange<B::IntTensorPrimitive>, TensorDeviceError> {
let (value, receive_counts) = self.native_rows_chunked(Primitive::Int(value), send_counts, max_chunk_bytes, &mut progress)?;
match value { Primitive::Int(value) => Ok(VariableTensorExchange { value, receive_counts }),
Primitive::Float(_) => unreachable!("original integer variable-row kind") }
}
fn native_rows_chunked<F: FnMut(RowExchangeProgress)>(
&self, value: Primitive<B>, send_counts: &[usize], max_chunk_bytes: usize, progress: &mut F,
) -> Result<(Primitive<B>, Vec<usize>), TensorDeviceError> {
let (shape, dtype) = self.native_sharded_metadata(&value)?;
let world = self.world_size() as usize;
if world == 0 { return Err(TensorDeviceError::InvalidOperation("variable-row world must be positive")); }
let rows = shape.first().copied();
let count_sum = send_counts.iter().try_fold(0usize, |total, count| total.checked_add(*count));
let valid = !shape.is_empty() && send_counts.len() == world && count_sum == rows;
let status = [u64::from(valid), send_counts.len() as u64];
let (statuses, _) = self.host().all_gather_host_staged(ElementType::U64, status.len(),
status.into_iter().flat_map(u64::to_le_bytes).collect())?;
if statuses.chunks_exact(16).any(|state| u64::from_le_bytes(state[..8].try_into().unwrap()) == 0) {
return Err(TensorDeviceError::InvalidBuffer("at least one rank's row counts differ from its actual input/leading axis"));
}
let rows = rows.expect("collectively validated row axis");
let element_type = native_element_type(dtype)?;
let tail = Shape::from(shape[1..].to_vec());
self.native_chunk_agreement(&tail, element_type, 0x300, max_chunk_bytes)?;
let row_elements = tail.iter().try_fold(1usize, |count, dimension| count.checked_mul(*dimension))
.filter(|count| *count != 0).ok_or(TensorDeviceError::InvalidBuffer("variable-row trailing axes are zero or overflowed"))?;
let counts_payload = send_counts.iter().flat_map(|count| (*count as u64).to_le_bytes()).collect::<Vec<_>>();
let (matrix, _) = self.host().all_gather_host_staged(ElementType::U64, world, counts_payload)?;
let matrix = matrix.chunks_exact(8).map(|bytes| usize::try_from(u64::from_le_bytes(bytes.try_into().unwrap()))
.map_err(|_| TensorDeviceError::InvalidBuffer("variable-row wire count exceeds native index range")))
.collect::<Result<Vec<_>, _>>()?;
let receive_counts = (0..world).map(|source| matrix[source * world + self.rank() as usize]).collect::<Vec<_>>();
let send_prefix = row_prefix(send_counts)?;
let receive_prefix = row_prefix(&receive_counts)?;
let received_rows = *receive_prefix.last().expect("rank prefix has its zero origin");
let input_elements = rows.checked_mul(row_elements).ok_or(TensorDeviceError::InvalidBuffer("variable-row source size overflow"))?;
let output_elements = received_rows.checked_mul(row_elements).ok_or(TensorDeviceError::InvalidBuffer("variable-row output size overflow"))?;
let maximum_peer_rows = matrix.iter().copied().max().unwrap_or(0);
let row_bytes = row_elements.checked_mul(element_type.byte_width()).ok_or(TensorDeviceError::InvalidBuffer("native row byte count overflow"))?;
let peer_rows = if maximum_peer_rows == 0 { 1 } else { max_chunk_bytes / world / row_bytes };
if peer_rows == 0 { return Err(TensorDeviceError::InvalidBuffer("variable-row budget must fit one native row per participant")); }
NativeChunkPlan::new(input_elements, element_type.byte_width(), element_type.byte_width())?;
NativeChunkPlan::new(output_elements, element_type.byte_width(), element_type.byte_width())?;
let value = native_reshape::<B>(value, Shape::new([input_elements]));
let mut output = self.native_empty_like(&value, output_elements, dtype);
let mut state = RowExchangeProgress { completed_rounds: 0, total_rounds: maximum_peer_rows.div_ceil(peer_rows),
sent_rows: 0, total_send_rows: rows, received_rows: 0, total_receive_rows: received_rows };
progress(state);
for round in 0..state.total_rounds {
let offset = round * peer_rows;
let counts = send_counts.iter().map(|count| count.saturating_sub(offset).min(peer_rows)).collect::<Vec<_>>();
let pieces = counts.iter().enumerate().filter(|(_, count)| **count > 0).map(|(destination, count)| {
let start = (send_prefix[destination] + offset) * row_elements;
value.clone().slice(start..start + count * row_elements)
}).collect::<Vec<_>>();
let sent = counts.iter().sum::<usize>();
let packet = if pieces.is_empty() { self.native_empty_like(&value, 0, dtype) } else {
match &value {
Primitive::Float(_) => Primitive::Float(B::float_cat(pieces.into_iter().map(|piece| match piece {
Primitive::Float(value) => value, Primitive::Int(_) => unreachable!("original floating source rows"),
}).collect(), 0)),
Primitive::Int(_) => Primitive::Int(B::int_cat(pieces.into_iter().map(|piece| match piece {
Primitive::Int(value) => value, Primitive::Float(_) => unreachable!("original integer source rows"),
}).collect(), 0)),
}
};
let mut packet_shape = shape.clone(); packet_shape[0] = sent;
let packet = native_reshape::<B>(packet, packet_shape);
let (received, chunk_counts) = match packet {
Primitive::Float(value) => { let result = self.all_to_all_v_float(value, &counts)?;
(Primitive::Float(result.value), result.receive_counts) },
Primitive::Int(value) => { let result = self.all_to_all_v_int(value, &counts)?;
(Primitive::Int(result.value), result.receive_counts) },
};
let expected = receive_counts.iter().map(|count| count.saturating_sub(offset).min(peer_rows)).collect::<Vec<_>>();
if chunk_counts != expected { return Err(TensorDeviceError::InvalidBuffer("native row round receive counts differ from original source matrix")); }
let chunk_prefix = row_prefix(&chunk_counts)?;
let received = native_reshape::<B>(received, Shape::new([chunk_prefix[world] * row_elements]));
for source in 0..world {
let count = chunk_counts[source]; if count == 0 { continue; }
let start = (receive_prefix[source] + offset) * row_elements;
let piece = received.clone().slice(chunk_prefix[source] * row_elements..chunk_prefix[source + 1] * row_elements);
output = output.assign(start..start + count * row_elements, piece);
}
B::sync(self.execution().device())?;
state.completed_rounds += 1;
state.sent_rows += sent;
state.received_rows += chunk_prefix[world];
progress(state);
}
let mut output_shape = shape; output_shape[0] = received_rows;
Ok((native_reshape::<B>(output, output_shape), receive_counts))
}
}
fn row_prefix(counts: &[usize]) -> Result<Vec<usize>, TensorDeviceError> {
let mut result = Vec::with_capacity(counts.len() + 1); result.push(0usize);
for count in counts { result.push(result.last().unwrap().checked_add(*count)
.ok_or(TensorDeviceError::InvalidBuffer("variable-row prefix overflow"))?); }
Ok(result)
}