use alloc::vec::Vec;
use ruda_autodiff::{Autodiff,checkpoint::strategy::CheckpointStrategy,collective};
use ruda_model::{
module::{Module,Param},
tensor::{Tensor,Int,DType,FloatDType,backend::Backend,module::{linear,embedding}},
};
use ruda_autodiff::tensor_parallel::BroadcastTensorCollective;
#[derive(Module,Debug)]
pub struct ShardedParameter<B:Backend> {
pub local:Param<Tensor<B,1>>,
pub logical_shape:Vec<usize>,
pub rank:usize,
pub world_size:usize,
}
impl<B:Backend> ShardedParameter<B> {
pub fn from_local(local:Param<Tensor<B,1>>,logical_shape:Vec<usize>,rank:usize,world_size:usize)->Self {
assert!(world_size>0 && rank<world_size,"invalid parameter shard topology");
assert!(!logical_shape.is_empty() && logical_shape.iter().all(|&n|n>0),"logical parameter dimensions must be positive");
let elements=logical_shape.iter().try_fold(1usize,|n,&d|n.checked_mul(d)).expect("logical parameter size overflow");
assert_eq!(local.val().dims(),[elements.div_ceil(world_size)],"local parameter slice length differs");
assert!(matches!(local.val().dtype(),DType::F32|DType::F16|DType::BF16),"floating parameter storage required");
Self{local,logical_shape,rank,world_size}
}
/// Slice a caller-loaded logical value on its current backend/device.
/// Existing parameter IDs are retained for explicitly shared/tied weights.
pub fn from_full<const D:usize>(parameter:Param<Tensor<B,D>>,rank:usize,world_size:usize)->Self {
assert!(world_size>0 && rank<world_size,"invalid parameter shard topology");
let value=parameter.val();
let shape=value.dims().to_vec();
let elements=shape.iter().try_fold(1usize,|n,&d|n.checked_mul(d)).expect("logical parameter size overflow");
assert!(elements>0,"logical parameter cannot be empty");
let size=elements.div_ceil(world_size);
let start=rank*size;
let end=(start+size).min(elements);
let mut local=Tensor::<B,1>::zeros([size],&value.device()).cast(value.dtype());
if end>start {
local=local.slice_assign([0..end-start],value.clone().reshape([elements]).slice([start..end]));
}
// A shard is a new leaf: its optimizer must not keep the full source graph.
let trainable=value.is_require_grad();
let local=local.detach().set_require_grad(trainable);
Self::from_local(Param::initialized(parameter.id,local),shape,rank,world_size)
}
}
impl<B:Backend,S:CheckpointStrategy> ShardedParameter<Autodiff<B,S>> {
/// Gather a logical value; backward uses the data-parallel SUM derivative.
/// All ranks must execute the same collective-bearing layer order.
pub fn gather<C:BroadcastTensorCollective<B>,const D:usize>(&self,communicator:C)
->Result<Tensor<Autodiff<B,S>,D>,C::Error> {
self.gather_inner(communicator,None)
}
/// Explicit gather/arithmetic precision without changing parameter storage.
/// Casting before gather also keeps its collective backward in this dtype.
pub fn gather_with_compute_dtype<C:BroadcastTensorCollective<B>,const D:usize>(
&self,communicator:C,dtype:FloatDType)->Result<Tensor<Autodiff<B,S>,D>,C::Error> {
self.gather_inner(communicator,Some(dtype.into()))
}
fn gather_inner<C:BroadcastTensorCollective<B>,const D:usize>(&self,communicator:C,dtype:Option<DType>)
->Result<Tensor<Autodiff<B,S>,D>,C::Error> {
assert_eq!(communicator.rank() as usize,self.rank,"parameter shard rank differs");
assert_eq!(communicator.world_size() as usize,self.world_size,"parameter shard world size differs");
let shape:[usize;D]=self.logical_shape.clone().try_into().expect("logical parameter rank differs");
let elements=self.logical_shape.iter().product();
let local=self.local.val();
let local=if let Some(dtype)=dtype {local.cast(dtype)} else {local};
collective::all_gather(local,communicator).map(|full|full.slice([0..elements]).reshape(shape))
}
}
/// Linear projection with local flattened weight/bias and ordinary global output.
/// Weight logical layout is RUDA's `[input,output]`, not `[output,input]`.
#[derive(Module,Debug)]
pub struct FullyShardedLinear<B:Backend> {
/// Full logical matrix metadata and its local storage.
pub weight:ShardedParameter<B>,
/// Optional full logical output-bias metadata and local storage.
pub bias:Option<ShardedParameter<B>>,
}
impl<B:Backend> FullyShardedLinear<B> {
/// Convert caller-loaded dense weights into local slices before optimizer creation.
pub fn from_full(layer:crate::Linear<B>,rank:usize,world_size:usize)->Self {
Self{weight:ShardedParameter::from_full(layer.weight,rank,world_size),
bias:layer.bias.map(|bias|ShardedParameter::from_full(bias,rank,world_size))}
}
}
impl<B:Backend,S:CheckpointStrategy> FullyShardedLinear<Autodiff<B,S>> {
/// Materialize this layer's parameters, project local data, and reduce-scatter in backward.
/// Full gathered weights may be retained by the selected autodiff checkpoint strategy.
pub fn forward<C:BroadcastTensorCollective<B>,const D:usize>(
&self,input:Tensor<Autodiff<B,S>,D>,communicator:C)->Result<Tensor<Autodiff<B,S>,D>,C::Error> {
let weight=self.weight.gather::<C,2>(communicator.clone())?;
let bias=match &self.bias {Some(bias)=>Some(bias.gather::<C,1>(communicator)?),None=>None};
Ok(linear(input,weight,bias))
}
}
/// Embedding with local flattened table slices and globally indexed token lookup.
#[derive(Module,Debug)]
pub struct FullyShardedEmbedding<B:Backend> {
/// Full `[vocabulary,features]` metadata and local table storage.
pub weight:ShardedParameter<B>,
}
impl<B:Backend> FullyShardedEmbedding<B> {
/// Split a loaded ordinary embedding on its current backend.
pub fn from_full(layer:crate::Embedding<B>,rank:usize,world_size:usize)->Self {
Self{weight:ShardedParameter::from_full(layer.weight,rank,world_size)}
}
}
impl<B:Backend,S:CheckpointStrategy> FullyShardedEmbedding<Autodiff<B,S>> {
/// Gather this table, preserving the existing integer lookup and padding semantics.
pub fn forward<C:BroadcastTensorCollective<B>>(&self,input:Tensor<Autodiff<B,S>,2,Int>,communicator:C)
->Result<Tensor<Autodiff<B,S>,3>,C::Error> {
Ok(embedding(self.weight.gather::<C,2>(communicator)?,input))
}
}