ruda-nn 0.21.32

Ruda neural network layers, activation modules and losses.
//! Rust model layers whose persistent parameters/gradients are equal flat shards.
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;

/// A logical parameter backed only by this rank's padded element slice.
///
/// Tied parameters may share the same local Param ID. Gathered values are
/// transient tensors, not registered full parameters. Backward sums independent
/// rank-local losses into a reduce-scattered gradient. Normalize the loss by the
/// global sample/token weight BEFORE backward; do not all-reduce these gradients
/// a second time. Select an elementwise optimizer for these flattened records.
#[derive(Module,Debug)]
pub struct ShardedParameter<B:Backend> {
    /// Rank-local storage; padding lies only after the logical element count.
    pub local:Param<Tensor<B,1>>,
    /// Full logical dimensions, excluding padding.
    pub logical_shape:Vec<usize>,
    /// Owner of this equal contiguous slice.
    pub rank:usize,
    /// Number of equal padded slices.
    pub world_size:usize,
}

impl<B:Backend> ShardedParameter<B> {
    /// Construct from an already loaded local checkpoint slice; no full tensor is needed.
    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))
    }
}