Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
@@ -0,0 +1,438 @@
//! Redistribution operations for distributed tensors.
use crate::comm::ReduceOp;
use crate::error::{DistributedError, Result};
use crate::group::ProcessGroup;
use parking_lot::RwLock;
use rtx_tensor::Tensor;
use std::sync::Arc;
use std::sync::atomic::Ordering;
use super::dtensor_core::{DTENSOR_ID_COUNTER, DTensor};
use super::placement::{PartialReduceOp, Placement};
use super::spec::TensorSpec;
impl DTensor {
/// Redistribute tensor to a new placement specification.
///
/// This inserts necessary communication (AllGather, ReduceScatter, etc.)
/// to convert the tensor to the target placement.
pub fn redistribute(&self, target_spec: TensorSpec, pg: &ProcessGroup) -> Result<DTensor> {
// Check if redistribution is needed
if self.spec.placements == target_spec.placements {
return Ok(self.clone());
}
// For each mesh dimension, determine required communication
let mut current = self.clone();
for (mesh_dim, (src_placement, dst_placement)) in self
.spec
.placements
.iter()
.zip(target_spec.placements.iter())
.enumerate()
{
if src_placement == dst_placement {
continue;
}
current = match (src_placement, dst_placement) {
// Shard -> Replicate: AllGather
(Placement::Shard { tensor_dim }, Placement::Replicate) => {
current.all_gather(mesh_dim, *tensor_dim, pg)?
}
// Replicate -> Shard: Local slice
(Placement::Replicate, Placement::Shard { tensor_dim }) => {
current.local_shard_for_placement(mesh_dim, *tensor_dim)?
}
// Partial -> Replicate: AllReduce
(Placement::Partial { reduce_op }, Placement::Replicate) => {
current.all_reduce(mesh_dim, (*reduce_op).into(), pg)?
}
// Shard -> Shard (different dim): AllToAll
(
Placement::Shard {
tensor_dim: src_dim,
},
Placement::Shard {
tensor_dim: dst_dim,
},
) if src_dim != dst_dim => current.reshard(mesh_dim, *src_dim, *dst_dim, pg)?,
_ => {
return Err(DistributedError::configuration(format!(
"Unsupported redistribution: {:?} -> {:?}",
src_placement, dst_placement
)));
}
};
}
// Update spec
let new_spec = TensorSpec {
placements: target_spec.placements,
..current.spec.clone()
};
Ok(DTensor {
spec: new_spec,
..current
})
}
/// AllGather: Collect sharded tensor to full tensor on all devices.
///
/// Gathers shards from all devices along the specified tensor dimension
/// and replicates the full tensor to all devices.
pub(crate) fn all_gather(
&self,
mesh_dim: usize,
tensor_dim: usize,
pg: &ProcessGroup,
) -> Result<DTensor> {
// Get the local shard data
let local_tensor = self.local_shard.read().clone();
// Use ProcessGroup's all_gather to collect tensors from all ranks
let gathered_tensors = pg.all_gather(&local_tensor)?;
// Concatenate gathered tensors along tensor_dim
let gathered_tensor = if gathered_tensors.is_empty() {
local_tensor
} else if gathered_tensors.len() == 1 {
gathered_tensors.into_iter().next().unwrap()
} else {
// Concatenate all tensors along tensor_dim
Self::concatenate_tensors(&gathered_tensors, tensor_dim)?
};
// Create new spec with Replicate placement for this mesh dimension
let mut new_placements = self.spec.placements.clone();
if mesh_dim < new_placements.len() {
new_placements[mesh_dim] = Placement::Replicate;
}
let new_spec = TensorSpec {
global_shape: self.spec.global_shape.clone(),
placements: new_placements,
dtype: self.spec.dtype,
requires_grad: self.spec.requires_grad,
name: self.spec.name.clone(),
};
Ok(DTensor {
id: DTENSOR_ID_COUNTER.fetch_add(1, Ordering::SeqCst),
local_shard: Arc::new(RwLock::new(gathered_tensor)),
spec: new_spec,
mesh: self.mesh.clone(),
local_coord: self.local_coord.clone(),
grad: Arc::new(RwLock::new(None)),
})
}
/// Helper: Concatenate tensors along a dimension.
pub(crate) fn concatenate_tensors(tensors: &[Tensor], dim: usize) -> Result<Tensor> {
if tensors.is_empty() {
return Err(DistributedError::tensor(
"Cannot concatenate empty tensor list",
));
}
let first = &tensors[0];
let first_shape = first.shape().dims().to_vec();
let ndims = first_shape.len();
if dim >= ndims {
return Err(DistributedError::tensor(format!(
"Concatenation dim {} out of range for {}D tensor",
dim, ndims
)));
}
// Calculate output shape
let mut out_shape = first_shape.clone();
out_shape[dim] = tensors.iter().map(|t| t.shape().dims()[dim]).sum();
// Calculate total elements
let out_numel: usize = out_shape.iter().product();
let mut out_data = vec![0.0f32; out_numel];
// Calculate stride for the concatenation dimension
let inner_size: usize = first_shape[dim + 1..].iter().product::<usize>().max(1);
let outer_size: usize = first_shape[..dim].iter().product::<usize>().max(1);
let mut offset_in_dim = 0;
for tensor in tensors {
let t_data = tensor
.to_vec()
.map_err(|e| DistributedError::tensor(e.to_string()))?;
let t_shape = tensor.shape().dims().to_vec();
let t_dim_size = t_shape[dim];
// Copy data with proper striding
for outer in 0..outer_size {
for d in 0..t_dim_size {
for inner in 0..inner_size {
let src_idx = outer * t_dim_size * inner_size + d * inner_size + inner;
let dst_dim_idx = offset_in_dim + d;
let dst_idx =
outer * out_shape[dim] * inner_size + dst_dim_idx * inner_size + inner;
if src_idx < t_data.len() && dst_idx < out_data.len() {
out_data[dst_idx] = t_data[src_idx];
}
}
}
}
offset_in_dim += t_dim_size;
}
Tensor::from_data(out_data, out_shape, &crate::Device::cpu())
.map_err(|e| DistributedError::tensor(e.to_string()))
}
/// AllReduce: Reduce partial results across devices.
///
/// Reduces partial tensor values across all devices using the specified
/// reduction operation, then replicates the result to all devices.
pub(crate) fn all_reduce(
&self,
mesh_dim: usize,
reduce_op: ReduceOp,
pg: &ProcessGroup,
) -> Result<DTensor> {
// Get the local tensor data
let mut local_tensor = self.local_shard.read().clone();
// Use ProcessGroup's all_reduce - works with NCCL, RNCCL, or CPU fallback
// ProcessGroup::all_reduce is async, so we use block_on for sync context
futures::executor::block_on(pg.all_reduce(&mut local_tensor, reduce_op))?;
// If this was a Partial(Mean), divide by world_size after sum
let world_size = pg.world_size();
if let Some(Placement::Partial {
reduce_op: PartialReduceOp::Mean,
}) = self.spec.placements.get(mesh_dim)
{
let mut data = local_tensor
.to_vec()
.map_err(|e| DistributedError::tensor(e.to_string()))?;
for value in &mut data {
*value /= world_size as f32;
}
local_tensor = Tensor::from_data(
data,
local_tensor.shape().dims().to_vec(),
local_tensor.device(),
)
.map_err(|e| DistributedError::tensor(e.to_string()))?;
}
// The tensor has been reduced in-place
let reduced_tensor = local_tensor;
// Create new spec with Replicate placement for this mesh dimension
let mut new_placements = self.spec.placements.clone();
if mesh_dim < new_placements.len() {
new_placements[mesh_dim] = Placement::Replicate;
}
let new_spec = TensorSpec {
global_shape: self.spec.global_shape.clone(),
placements: new_placements,
dtype: self.spec.dtype,
requires_grad: self.spec.requires_grad,
name: self.spec.name.clone(),
};
Ok(DTensor {
id: DTENSOR_ID_COUNTER.fetch_add(1, Ordering::SeqCst),
local_shard: Arc::new(RwLock::new(reduced_tensor)),
spec: new_spec,
mesh: self.mesh.clone(),
local_coord: self.local_coord.clone(),
grad: Arc::new(RwLock::new(None)),
})
}
/// Extract local shard from replicated tensor.
///
/// Slices the tensor along `tensor_dim` to get the shard belonging to this
/// device based on its coordinate in the mesh.
pub(crate) fn local_shard_for_placement(
&self,
mesh_dim: usize,
tensor_dim: usize,
) -> Result<DTensor> {
let mesh_size = self.mesh.shape().get(mesh_dim).copied().unwrap_or(1);
let mesh_coord = self.local_coord.get(mesh_dim).copied().unwrap_or(0);
let global_size = self.spec.global_shape.get(tensor_dim).copied().unwrap_or(0);
let shard_size = global_size.div_ceil(mesh_size);
let start = mesh_coord * shard_size;
let end = ((mesh_coord + 1) * shard_size).min(global_size);
// Get the local tensor data
let local_tensor = self.local_shard.read().clone();
let local_shape = local_tensor.shape().dims().to_vec();
let data = local_tensor
.to_vec()
.map_err(|e| DistributedError::tensor(e.to_string()))?;
// Calculate new shape after slicing
let mut new_shape = local_shape.clone();
if tensor_dim < new_shape.len() {
new_shape[tensor_dim] = end - start;
}
// Perform the slicing
let sliced_data = Self::slice_tensor_data(&data, &local_shape, tensor_dim, start, end)?;
// Create sliced tensor
let sliced_tensor = Tensor::from_data(sliced_data, new_shape, local_tensor.device())
.map_err(|e| DistributedError::tensor(e.to_string()))?;
// Create new spec with Shard placement
let mut new_placements = self.spec.placements.clone();
while new_placements.len() <= mesh_dim {
new_placements.push(Placement::Replicate);
}
new_placements[mesh_dim] = Placement::Shard { tensor_dim };
let new_spec = TensorSpec {
global_shape: self.spec.global_shape.clone(),
placements: new_placements,
dtype: self.spec.dtype,
requires_grad: self.spec.requires_grad,
name: self.spec.name.clone(),
};
Ok(DTensor {
id: DTENSOR_ID_COUNTER.fetch_add(1, Ordering::SeqCst),
local_shard: Arc::new(RwLock::new(sliced_tensor)),
spec: new_spec,
mesh: self.mesh.clone(),
local_coord: self.local_coord.clone(),
grad: Arc::new(RwLock::new(None)),
})
}
/// Helper: Slice tensor data along a dimension.
pub(crate) fn slice_tensor_data(
data: &[f32],
shape: &[usize],
dim: usize,
start: usize,
end: usize,
) -> Result<Vec<f32>> {
if dim >= shape.len() {
return Err(DistributedError::tensor(format!(
"Slice dim {} out of range for {}D tensor",
dim,
shape.len()
)));
}
let slice_size = end - start;
let mut out_shape = shape.to_vec();
out_shape[dim] = slice_size;
// Calculate strides
let inner_size: usize = shape[dim + 1..].iter().product::<usize>().max(1);
let outer_size: usize = shape[..dim].iter().product::<usize>().max(1);
let dim_size = shape[dim];
let out_numel: usize = out_shape.iter().product();
let mut out_data = vec![0.0f32; out_numel];
// Copy sliced data
for outer in 0..outer_size {
for d in 0..slice_size {
for inner in 0..inner_size {
let src_d = start + d;
let src_idx = outer * dim_size * inner_size + src_d * inner_size + inner;
let dst_idx = outer * slice_size * inner_size + d * inner_size + inner;
if src_idx < data.len() && dst_idx < out_data.len() {
out_data[dst_idx] = data[src_idx];
}
}
}
}
Ok(out_data)
}
/// Reshard: Change sharding dimension (requires AllToAll).
///
/// Redistributes the tensor from being sharded on `src_tensor_dim` to being
/// sharded on `dst_tensor_dim`. This is implemented as AllGather followed
/// by local slicing (AllToAll pattern).
pub(crate) fn reshard(
&self,
mesh_dim: usize,
src_tensor_dim: usize,
dst_tensor_dim: usize,
pg: &ProcessGroup,
) -> Result<DTensor> {
// Step 1: AllGather to get full tensor (shard src_tensor_dim -> replicate)
let gathered = self.all_gather(mesh_dim, src_tensor_dim, pg)?;
// Step 2: Take local slice along dst_tensor_dim (replicate -> shard dst_tensor_dim)
let local_slice = gathered.local_shard_for_placement(mesh_dim, dst_tensor_dim)?;
// Return with updated spec
Ok(local_slice)
}
/// ReduceScatter: Reduce across devices and scatter results.
///
/// This is the inverse of AllGather. Used in FSDP backward pass to reduce
/// gradients and scatter them back to the appropriate shards.
pub fn reduce_scatter(
&self,
mesh_dim: usize,
tensor_dim: usize,
reduce_op: ReduceOp,
pg: &ProcessGroup,
) -> Result<DTensor> {
// Get the local tensor data
let local_tensor = self.local_shard.read().clone();
// Use ProcessGroup's reduce_scatter
let scattered_tensor = pg.reduce_scatter(&local_tensor, reduce_op)?;
// Calculate new local shape after scatter
let mesh_size = self.mesh.shape().get(mesh_dim).copied().unwrap_or(1);
let mut new_local_shape = self.spec.global_shape.clone();
if tensor_dim < new_local_shape.len() {
new_local_shape[tensor_dim] = new_local_shape[tensor_dim].div_ceil(mesh_size);
}
// Create new spec with Shard placement
let mut new_placements = self.spec.placements.clone();
while new_placements.len() <= mesh_dim {
new_placements.push(Placement::Replicate);
}
new_placements[mesh_dim] = Placement::Shard { tensor_dim };
let new_spec = TensorSpec {
global_shape: self.spec.global_shape.clone(),
placements: new_placements,
dtype: self.spec.dtype,
requires_grad: self.spec.requires_grad,
name: self.spec.name.clone(),
};
Ok(DTensor {
id: DTENSOR_ID_COUNTER.fetch_add(1, Ordering::SeqCst),
local_shard: Arc::new(RwLock::new(scattered_tensor)),
spec: new_spec,
mesh: self.mesh.clone(),
local_coord: self.local_coord.clone(),
grad: Arc::new(RwLock::new(None)),
})
}
}