//! Placement types for distributed tensors. use crate::comm::ReduceOp; use serde::{Deserialize, Serialize}; /// Specifies how a tensor dimension is placed across a mesh dimension. #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub enum Placement { /// The tensor dimension is replicated across this mesh dimension. /// All devices in this mesh dimension hold the same data. Replicate, /// The tensor dimension is sharded across this mesh dimension. /// Each device holds a slice of the data. Shard { /// Which tensor dimension is sharded (index into tensor shape) tensor_dim: usize, }, /// Partial placement - tensor holds partial results that need reduction. /// Used after operations like matmul where results are distributed. Partial { /// Reduction operation needed to combine partial results reduce_op: PartialReduceOp, }, /// Interleaved sharding for better memory access patterns. /// Data is distributed in an interleaved fashion rather than contiguous chunks. InterleavedShard { /// Which tensor dimension is sharded tensor_dim: usize, /// Interleave factor (default: 1) interleave_factor: usize, }, } /// Reduction operation for partial placements. #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] pub enum PartialReduceOp { /// Sum reduction Sum, /// Mean reduction Mean, /// Max reduction Max, /// Min reduction Min, } impl From for ReduceOp { fn from(op: PartialReduceOp) -> Self { match op { PartialReduceOp::Sum => ReduceOp::Sum, PartialReduceOp::Mean => ReduceOp::Sum, // Mean is sum then divide PartialReduceOp::Max => ReduceOp::Max, PartialReduceOp::Min => ReduceOp::Min, } } }