use super::*;
use crate::ElementwiseShardOptimizer;
pub trait FlatShardElementwiseOptimizer<B:Backend>:ElementwiseShardOptimizer<B> {
fn partition_native_history(record:AdaptorRecordV1<Self,B>,shard:&FlatOptimizerTensorShard) -> Result<Self::State<1>,OptimizerShardError>
where Self:Sized;
}
impl<B:Backend,O:ElementwiseShardOptimizer<B>> FlatShardElementwiseOptimizer<B> for O
where O::State<0>:FlatOptimizerCheckpointState<B,0,FlatState=O::State<1>>,
O::State<1>:FlatOptimizerCheckpointState<B,1,FlatState=O::State<1>>,
O::State<2>:FlatOptimizerCheckpointState<B,2,FlatState=O::State<1>>,
O::State<3>:FlatOptimizerCheckpointState<B,3,FlatState=O::State<1>>,
O::State<4>:FlatOptimizerCheckpointState<B,4,FlatState=O::State<1>>,
O::State<5>:FlatOptimizerCheckpointState<B,5,FlatState=O::State<1>>,
O::State<6>:FlatOptimizerCheckpointState<B,6,FlatState=O::State<1>>,
O::State<7>:FlatOptimizerCheckpointState<B,7,FlatState=O::State<1>>,
O::State<8>:FlatOptimizerCheckpointState<B,8,FlatState=O::State<1>> {
fn partition_native_history(record:AdaptorRecordV1<Self,B>,shard:&FlatOptimizerTensorShard) -> Result<Self::State<1>,OptimizerShardError> {
match record {
AdaptorRecordV1::Rank0(state)=>state.into_flat_shard(shard),
AdaptorRecordV1::Rank1(state)=>state.into_flat_shard(shard),
AdaptorRecordV1::Rank2(state)=>state.into_flat_shard(shard),
AdaptorRecordV1::Rank3(state)=>state.into_flat_shard(shard),
AdaptorRecordV1::Rank4(state)=>state.into_flat_shard(shard),
AdaptorRecordV1::Rank5(state)=>state.into_flat_shard(shard),
AdaptorRecordV1::Rank6(state)=>state.into_flat_shard(shard),
AdaptorRecordV1::Rank7(state)=>state.into_flat_shard(shard),
AdaptorRecordV1::Rank8(state)=>state.into_flat_shard(shard),
}
}
}
pub trait OptimizerCheckpointScalars {
type Scalars:PartialEq;
fn checkpoint_scalars(&self) -> Self::Scalars;
}
pub trait FlatOptimizerCheckpointState<B:Backend,const D:usize>:OptimizerCheckpointBuffers<B,D> {
type FlatState:OptimizerCheckpointBuffers<B,1>;
fn into_flat_shard(self,shard:&FlatOptimizerTensorShard) -> Result<Self::FlatState,OptimizerShardError>;
}
#[derive(Clone,Debug,PartialEq,Eq)]
pub struct FlatOptimizerTensorShard {
pub global_shape:Vec<usize>,
pub rank:u32,
pub world_size:u32,
}
impl FlatOptimizerTensorShard {
pub fn new(global_shape:Vec<usize>,rank:u32,world_size:u32) -> Result<Self,OptimizerShardError> {
let value=Self {global_shape,rank,world_size};value.geometry()?;Ok(value)
}
pub fn geometry(&self) -> Result<(usize,usize),OptimizerShardError> {
if self.global_shape.contains(&0) || self.world_size==0 || self.rank>=self.world_size {
return Err(OptimizerShardError::Placement("positive original FSDP axes and valid data rank/world required"));
}
let count=self.global_shape.iter().try_fold(1usize,|total,axis|total.checked_mul(*axis)).ok_or(OptimizerShardError::Placement("original flat parameter size overflows"))?;
let slots=count.div_ceil(self.world_size as usize);
slots.checked_mul(self.world_size as usize).ok_or(OptimizerShardError::Placement("padded flat parameter size overflows"))?;
Ok((count,slots))
}
pub fn interval(&self) -> Result<Range<usize>,OptimizerShardError> {
let (count,slots)=self.geometry()?;let start=(self.rank as usize*slots).min(count);
Ok(start..start.saturating_add(slots).min(count))
}
pub fn partition<B:Backend,const D:usize>(&self,value:Tensor<B,D>) -> Result<Tensor<B,1>,OptimizerShardError> {
let (count,slots)=self.geometry()?;
if value.dims().as_slice()!=self.global_shape.as_slice() {return Err(OptimizerShardError::Shape("original flat optimizer buffer geometry"));}
if !value.dtype().is_float() || matches!(value.dtype(),DType::QFloat(_)) {return Err(OptimizerShardError::DType("native full optimizer buffer required"));}
if self.world_size==1 {return Ok(value.reshape([count]));}
let start=self.rank as usize*slots;let real=count.saturating_sub(start).min(slots);
let mut local=Tensor::zeros([slots],(&value.device(),value.dtype()));
if real>0 {local=local.slice_assign([0..real],value.reshape([count]).slice([start..start+real]));}Ok(local)
}
pub fn validate_state<B:Backend,const D:usize,S:OptimizerCheckpointBuffers<B,D>>(&self,state:&S) -> Result<(),OptimizerShardError> {
self.geometry()?;
if D!=self.global_shape.len() {return Err(OptimizerShardError::Shape("source optimizer history rank differs from original parameter"));}
let mut error=None;let mut device=None;
state.visit_checkpoint_buffers(&mut |value| {
if value.dims().as_slice()!=self.global_shape.as_slice() {error=Some(OptimizerShardError::Shape("original native history dimensions"));}
if !value.dtype().is_float() || matches!(value.dtype(),DType::QFloat(_)) {error=Some(OptimizerShardError::DType("original native history precision"));}
if device.as_ref().is_some_and(|previous|*previous!=value.device()) {error=Some(OptimizerShardError::Device("one native parameter history spans different devices"));}
device=Some(value.device());
});
if let Some(error)=error {return Err(error);}Ok(())
}
}
impl<B:Backend,const D:usize> FlatOptimizerCheckpointState<B,D> for AdaptiveMomentumState<B,D> {
type FlatState=AdaptiveMomentumState<B,1>;
fn into_flat_shard(self,shard:&FlatOptimizerTensorShard) -> Result<Self::FlatState,OptimizerShardError> {
shard.validate_state(&self)?;
Ok(AdaptiveMomentumState {time:self.time,moment_1:shard.partition(self.moment_1)?,moment_2:shard.partition(self.moment_2)?,
max_moment_2:self.max_moment_2.map(|value|shard.partition(value)).transpose()?})
}
}
impl<B:Backend,const D:usize> OptimizerCheckpointScalars for AdaptiveMomentumState<B,D> {
type Scalars=(usize,bool);
fn checkpoint_scalars(&self) -> Self::Scalars {(self.time,self.max_moment_2.is_some())}
}
impl<B:Backend,const D:usize> FlatOptimizerCheckpointState<B,D> for AdaptiveNesterovMomentumState<B,D> {
type FlatState=AdaptiveNesterovMomentumState<B,1>;
fn into_flat_shard(self,shard:&FlatOptimizerTensorShard) -> Result<Self::FlatState,OptimizerShardError> {
shard.validate_state(&self)?;
Ok(AdaptiveNesterovMomentumState {time:self.time,exp_avg:shard.partition(self.exp_avg)?,exp_avg_sq:shard.partition(self.exp_avg_sq)?,
exp_avg_diff:shard.partition(self.exp_avg_diff)?,neg_pre_grad:shard.partition(self.neg_pre_grad)?})
}
}
impl<B:Backend,const D:usize> OptimizerCheckpointScalars for AdaptiveNesterovMomentumState<B,D> {
type Scalars=usize;
fn checkpoint_scalars(&self) -> usize {self.time}
}
macro_rules! adaptive_flat {
($(($state:ident,$momentum:ident)),+) => {$(
impl<B:Backend,const D:usize> FlatOptimizerCheckpointState<B,D> for $state<B,D> {
type FlatState=$state<B,1>;
fn into_flat_shard(self,shard:&FlatOptimizerTensorShard) -> Result<Self::FlatState,OptimizerShardError> {
Ok($state {momentum:self.momentum.into_flat_shard(shard)?})
}
}
impl<B:Backend,const D:usize> OptimizerCheckpointScalars for $state<B,D> {
type Scalars=<$momentum<B,D> as OptimizerCheckpointScalars>::Scalars;
fn checkpoint_scalars(&self) -> Self::Scalars {self.momentum.checkpoint_scalars()}
}
)+};
}
adaptive_flat!((AdamState,AdaptiveMomentumState),(AdamWState,AdaptiveMomentumState),(AdanState,AdaptiveNesterovMomentumState));
impl<B:Backend,const D:usize> FlatOptimizerCheckpointState<B,D> for SgdState<B,D> {
type FlatState=SgdState<B,1>;
fn into_flat_shard(self,shard:&FlatOptimizerTensorShard) -> Result<Self::FlatState,OptimizerShardError> {
shard.validate_state(&self)?;Ok(SgdState {momentum:self.momentum.map(|value|value.into_flat_shard(shard)).transpose()?})
}
}
impl<B:Backend,const D:usize> OptimizerCheckpointScalars for SgdState<B,D> {
type Scalars=bool;
fn checkpoint_scalars(&self) -> bool {self.momentum.is_some()}
}
impl<B:Backend,const D:usize> FlatOptimizerCheckpointState<B,D> for SquareAvgState<B,D> {
type FlatState=SquareAvgState<B,1>;
fn into_flat_shard(self,shard:&FlatOptimizerTensorShard) -> Result<Self::FlatState,OptimizerShardError> {Ok(SquareAvgState {square_avg:shard.partition(self.square_avg)?})}
}
impl<B:Backend,const D:usize> FlatOptimizerCheckpointState<B,D> for CenteredState<B,D> {
type FlatState=CenteredState<B,1>;
fn into_flat_shard(self,shard:&FlatOptimizerTensorShard) -> Result<Self::FlatState,OptimizerShardError> {
shard.validate_state(&self)?;Ok(CenteredState {avg:shard.partition(self.avg)?,grad_avg:self.grad_avg.map(|value|shard.partition(value)).transpose()?})
}
}
impl<B:Backend,const D:usize> FlatOptimizerCheckpointState<B,D> for RmsPropState<B,D> {
type FlatState=RmsPropState<B,1>;
fn into_flat_shard(self,shard:&FlatOptimizerTensorShard) -> Result<Self::FlatState,OptimizerShardError> {
shard.validate_state(&self)?;Ok(RmsPropState {square_avg:self.square_avg.into_flat_shard(shard)?,centered:self.centered.into_flat_shard(shard)?,
momentum:self.momentum.map(|value|value.into_flat_shard(shard)).transpose()?})
}
}
impl<B:Backend,const D:usize> OptimizerCheckpointScalars for RmsPropState<B,D> {
type Scalars=(bool,bool);
fn checkpoint_scalars(&self) -> Self::Scalars {(self.centered.grad_avg.is_some(),self.momentum.is_some())}
}
impl<B:Backend,const D:usize,S:FlatOptimizerCheckpointState<B,D>> FlatOptimizerCheckpointState<B,D> for Fp32MasterState<B,D,S> {
type FlatState=Fp32MasterState<B,1,S::FlatState>;
fn into_flat_shard(self,shard:&FlatOptimizerTensorShard) -> Result<Self::FlatState,OptimizerShardError> {
shard.validate_state(&self)?;
let mut fp32=true;self.visit_checkpoint_buffers(&mut |value|fp32&=value.dtype()==DType::F32);
if !fp32 {return Err(OptimizerShardError::DType("authoritative native master and inner history must retain FP32"));}
Ok(Fp32MasterState {master:shard.partition(self.master)?,inner:self.inner.map(|state|state.into_flat_shard(shard)).transpose()?})
}
}
impl<B:Backend,const D:usize,S:Record<B>+Clone+OptimizerCheckpointScalars> OptimizerCheckpointScalars for Fp32MasterState<B,D,S> {
type Scalars=Option<S::Scalars>;
fn checkpoint_scalars(&self) -> Self::Scalars {self.inner.as_ref().map(|state|state.checkpoint_scalars())}
}