use alloc::vec::Vec;
use core::fmt;
use serde::{Serialize,Deserialize};
use ruda_model::{record::{Record,PrecisionSettings},tensor::{Tensor,TensorPrimitive,BroadcastTensorCollective,backend::Backend}};
use super::{Muon,MuonState,MuonError,NewtonSchulzParams,LearningRate};
#[derive(Clone,Debug,PartialEq,Eq,Serialize,Deserialize)]
pub struct MuonMatrixShardLayout {
pub axis: usize,
pub lengths: Vec<usize>,
}
impl MuonMatrixShardLayout {
pub fn new(axis: usize,lengths: Vec<usize>) -> Self {Self {axis,lengths}}
pub fn global_shape(&self,rank: u32,world: u32,local: [usize;2]) -> Result<[usize;2],MuonError> {
if self.axis > 1 || world == 0 || rank >= world || self.lengths.len() != world as usize || self.lengths.contains(&0) {
return Err(MuonError::InvalidConfig("Muon matrix shards need a valid axis and actual positive rank lengths"));
}
if local[self.axis] != self.lengths[rank as usize] {return Err(MuonError::ShapeMismatch("rank-local matrix shard"));}
let total = self.lengths.iter().try_fold(0usize,|sum,length|sum.checked_add(*length))
.ok_or(MuonError::InvalidConfig("global Muon matrix axis overflows"))?;
let mut global = local;global[self.axis] = total;
if global.contains(&0) {return Err(MuonError::EmptyMatrix);}
Ok(global)
}
pub fn range(&self,rank: u32) -> Result<core::ops::Range<usize>,MuonError> {
let rank = rank as usize;
if rank >= self.lengths.len() {return Err(MuonError::InvalidConfig("Muon shard rank is outside its layout"));}
let start = self.lengths[..rank].iter().try_fold(0usize,|sum,length|sum.checked_add(*length))
.ok_or(MuonError::InvalidConfig("Muon shard position overflows"))?;
let end = start.checked_add(self.lengths[rank]).ok_or(MuonError::InvalidConfig("Muon shard end overflows"))?;
Ok(start..end)
}
}
#[derive(Debug)]
pub enum MuonShardedError<E: fmt::Debug> {
Muon(MuonError),
Collective(E),
}
impl<E: fmt::Debug> fmt::Display for MuonShardedError<E> {
fn fmt(&self,f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {Self::Muon(error)=>write!(f,"{error}"),Self::Collective(error)=>write!(f,"sharded Muon collective failed: {error:?}")}
}
}
impl<E: fmt::Debug> core::error::Error for MuonShardedError<E> {}
#[derive(Clone)]
pub struct MuonShardedState<B: Backend> {
version: u32,
rank: u32,
layout: MuonMatrixShardLayout,
global_shape: [usize;2],
local: MuonState<B,2>,
}
impl<B: Backend> MuonShardedState<B> {
pub fn from_global_state(state:MuonState<B,2>,layout:&MuonMatrixShardLayout,rank:u32,world:u32) -> Result<Self,MuonError> {
if layout.axis > 1 || rank >= world || layout.lengths.len() != world as usize {return Err(MuonError::InvalidConfig("invalid Muon global-state shard placement"));}
let global_shape = state.momentum.velocity().dims();let mut local_shape = global_shape;
local_shape[layout.axis] = layout.lengths[rank as usize];
if layout.global_shape(rank,world,local_shape)? != global_shape {return Err(MuonError::ShapeMismatch("global momentum"));}
let local = state.momentum.velocity().clone().slice_dim(layout.axis,layout.range(rank)?);
Ok(Self {version:1,rank,layout:layout.clone(),global_shape,local:MuonState::new(super::MomentumState::new(local))})
}
pub fn rank(&self) -> u32 {self.rank}
pub fn layout(&self) -> &MuonMatrixShardLayout {&self.layout}
pub fn global_shape(&self) -> [usize;2] {self.global_shape}
pub fn validate_placement(&self,rank: u32,layout: &MuonMatrixShardLayout,global_shape: [usize;2]) -> Result<(),MuonError> {
if self.version != 1 || self.rank != rank || &self.layout != layout || self.global_shape != global_shape {Err(MuonError::IncompatibleRecord)} else {Ok(())}
}
pub fn momentum(&self) -> &Tensor<B,2> {self.local.momentum.velocity()}
pub fn to_device(mut self,device: &B::Device) -> Self {self.local.momentum = self.local.momentum.to_device(device);self}
}
impl<B: Backend> Record<B> for MuonShardedState<B> {
type Item<S: PrecisionSettings> = (u32,u32,usize,Vec<usize>,[usize;2],<MuonState<B,2> as Record<B>>::Item<S>);
fn into_item<S: PrecisionSettings>(self) -> Self::Item<S> {
(self.version,self.rank,self.layout.axis,self.layout.lengths,self.global_shape,self.local.into_item::<S>())
}
fn from_item<S: PrecisionSettings>(item: Self::Item<S>,device: &B::Device) -> Self {
Self {version:item.0,rank:item.1,layout:MuonMatrixShardLayout::new(item.2,item.3),global_shape:item.4,
local:MuonState::<B,2>::from_item::<S>(item.5,device)}
}
}
fn sum<B,C,const D: usize>(value: Tensor<B,D>,communicator: &C) -> Result<Tensor<B,D>,MuonShardedError<C::Error>>
where B: Backend,C: BroadcastTensorCollective<B> {
let primitive = communicator.all_reduce_sum(value.into_primitive().tensor()).map_err(MuonShardedError::Collective)?;
Ok(Tensor::from_primitive(TensorPrimitive::Float(primitive)))
}
fn maximum<B,C>(value: Tensor<B,1>,communicator: &C) -> Result<Tensor<B,1>,MuonShardedError<C::Error>>
where B: Backend,C: BroadcastTensorCollective<B> {
let primitive = communicator.all_gather_float(value.into_primitive().tensor()).map_err(MuonShardedError::Collective)?;
Ok(Tensor::<B,1>::from_primitive(TensorPrimitive::Float(primitive)).max())
}
fn gather_matrix<B,C>(value: Tensor<B,2>,layout: &MuonMatrixShardLayout,communicator: &C) -> Result<Tensor<B,2>,MuonShardedError<C::Error>>
where B: Backend,C: BroadcastTensorCollective<B> {
let leading = value.swap_dims(0,layout.axis);let [local,other] = leading.dims();
let slots = *layout.lengths.iter().max().expect("validated nonempty Muon rank layout");
let total = slots.checked_mul(layout.lengths.len()).ok_or(MuonShardedError::Muon(MuonError::InvalidConfig("Muon gather padding overflows")))?;
let padded = if local == slots {leading} else {
let device = leading.device();let dtype = leading.dtype();
Tensor::cat(alloc::vec![leading,Tensor::<B,2>::zeros([slots-local,other],(&device,dtype))],0)
};
let gathered = communicator.all_gather_float(padded.into_primitive().tensor()).map_err(MuonShardedError::Collective)?;
let gathered = Tensor::<B,2>::from_primitive(TensorPrimitive::Float(gathered));
assert_eq!(gathered.dims(),[total,other],"Muon transport returned incompatible matrix gather geometry");
let parts = layout.lengths.iter().enumerate().map(|(rank,length)|gathered.clone().slice_dim(0,rank*slots..rank*slots+length)).collect();
Ok(Tensor::cat(parts,0).swap_dims(0,layout.axis))
}
impl<B: Backend> Muon<B> {
pub fn validate_step_sharded<C>(&self,lr: LearningRate,tensor: &Tensor<B,2>,grad: &Tensor<B,2>,state: Option<&MuonShardedState<B>>,
layout: &MuonMatrixShardLayout,communicator: &C) -> Result<[usize;2],MuonShardedError<C::Error>>
where C: BroadcastTensorCollective<B> {
let shape = layout.global_shape(communicator.rank(),communicator.world_size(),tensor.dims()).map_err(MuonShardedError::Muon)?;
if let Some(state) = state {
state.validate_placement(communicator.rank(),layout,shape).map_err(MuonShardedError::Muon)?;
}
self.validate_step_shape(lr,tensor,grad,state.map(|state|&state.local),&shape).map_err(MuonShardedError::Muon)?;
Ok(shape)
}
pub fn try_step_sharded<C>(&self,lr: LearningRate,tensor: Tensor<B,2>,grad: Tensor<B,2>,state: Option<MuonShardedState<B>>,
layout: &MuonMatrixShardLayout,communicator: C) -> Result<(Tensor<B,2>,MuonShardedState<B>),MuonShardedError<C::Error>>
where C: BroadcastTensorCollective<B> {
let rank = communicator.rank();let world = communicator.world_size();
let shape = self.validate_step_sharded(lr,&tensor,&grad,state.as_ref(),layout,&communicator)?;
let (update,momentum) = self.momentum_update(grad,state.map(|state|state.local));
let transposed = shape[0] > shape[1];let oriented_axis = if transposed {1-layout.axis} else {layout.axis};
let update = if world == 1 {self.zeropower_via_newtonschulz(update)} else if oriented_axis == 0 {
let global = gather_matrix(update,layout,&communicator)?;
self.zeropower_via_newtonschulz(global).slice_dim(layout.axis,layout.range(rank).map_err(MuonShardedError::Muon)?)
} else {
let mut x = if transposed {update.swap_dims(0,1)} else {update};
if self.stable_normalization {
let scale = maximum(x.clone().abs().max(),&communicator)?.clamp_min(f32::MIN_POSITIVE);
let scaled = x.div(scale.clone().unsqueeze());let floor = scale.recip().mul_scalar(self.epsilon);
let norm = sum(scaled.clone().square().sum(),&communicator)?.sqrt().max_pair(floor);
x = scaled.div(norm.unsqueeze());
} else {
let norm = sum(x.clone().powf_scalar(2.0).sum(),&communicator)?.sqrt().clamp_min(self.epsilon);
x = x.div(norm.unsqueeze());
}
let NewtonSchulzParams {a,b,c,steps} = self.ns_params;
for _ in 0..steps {
let gram = sum(x.clone().matmul(x.clone().swap_dims(0,1)),&communicator)?;
let squared = gram.clone().matmul(gram.clone());
let polynomial = gram.mul_scalar(b).add(squared.mul_scalar(c));
x = x.clone().mul_scalar(a).add(polynomial.matmul(x));
}
if transposed {x.swap_dims(0,1)} else {x}
};
let adjusted = self.adjust_lr(lr,&shape);
let tensor = match self.weight_decay_penalty {Some(penalty)=>tensor.mul_scalar(1.0-lr*penalty as f64),None=>tensor};
Ok((tensor-update.mul_scalar(adjusted),MuonShardedState {version:1,rank,layout:layout.clone(),global_shape:shape,local:MuonState::new(momentum)}))
}
}