use alloc::vec::Vec;
use core::{fmt,ops::Range};
use hashbrown::HashMap;
use ruda_model::{module::ParamId,record::Record,tensor::{Tensor,DType,backend::{Backend,AutodiffBackend}}};
use super::{AdaptiveMomentumState,AdamState,AdamWState,Adam,AdamW,Fp32MasterOptimizer,Fp32MasterState,
Sgd,SgdState,AdaGrad,AdaGradState,LrDecayState,RmsProp,RmsPropState,SquareAvgState,CenteredState,RmsPropMomentumState,
Adan,AdanState,AdaptiveNesterovMomentumState,momentum::MomentumState,record::{AdaptorRecord,AdaptorRecordV1}};
pub trait OptimizerCheckpointBuffers<B:Backend,const D:usize>:Record<B>+Clone {
fn visit_checkpoint_buffers<F:FnMut(&Tensor<B,D>)>(&self,visit:&mut F);
fn map_checkpoint_buffers<F:FnMut(Tensor<B,D>)->Tensor<B,D>>(self,map:&mut F) -> Self;
}
pub fn partition_native_optimizer_state<B:Backend,const D:usize,S:OptimizerCheckpointBuffers<B,D>>(state:S,shard:&OptimizerTensorShard)
-> Result<S,OptimizerShardError> {
shard.validate()?;
if shard.global_shape.len() != D {return Err(OptimizerShardError::Shape("native state rank differs from source parameter"));}
let mut buffers = Vec::new();state.visit_checkpoint_buffers(&mut |value|buffers.push(value.clone()));
for value in &buffers {
shard.check(value,"native optimizer buffer")?;
if let Some(reference) = buffers.first() {check_buffer(value,reference,"native optimizer buffer")?;}
}
Ok(state.map_checkpoint_buffers(&mut |value|shard.slice(value)))
}
#[derive(Clone,Debug,PartialEq,Eq)]
pub struct OptimizerTensorShard {
pub global_shape:Vec<usize>,
pub axis:usize,
pub interval:Range<usize>,
}
#[derive(Clone,Debug,PartialEq,Eq)]
pub enum OptimizerShardError {
Placement(&'static str),
Shape(&'static str),
DType(&'static str),
Device(&'static str),
}
impl fmt::Display for OptimizerShardError {
fn fmt(&self,f:&mut fmt::Formatter<'_>) -> fmt::Result {
match self {Self::Placement(value)=>write!(f,"invalid optimizer tensor shard: {value}"),Self::Shape(value)=>write!(f,"optimizer shard shape differs: {value}"),
Self::DType(value)=>write!(f,"optimizer shard precision differs: {value}"),Self::Device(value)=>write!(f,"optimizer shard device differs: {value}")}
}
}
impl core::error::Error for OptimizerShardError {}
impl OptimizerTensorShard {
pub fn new(global_shape:Vec<usize>,axis:usize,interval:Range<usize>) -> Result<Self,OptimizerShardError> {
let shard = Self {global_shape,axis,interval};shard.validate()?;Ok(shard)
}
pub fn validate(&self) -> Result<(),OptimizerShardError> {
if self.axis >= self.global_shape.len() || self.interval.start >= self.interval.end || self.interval.end > self.global_shape[self.axis] {
return Err(OptimizerShardError::Placement("axis/interval must be a nonempty actual source slice"));
}
Ok(())
}
pub fn local_shape(&self) -> Result<Vec<usize>,OptimizerShardError> {
self.validate()?;let mut shape = self.global_shape.clone();shape[self.axis] = self.interval.len();Ok(shape)
}
fn check<B:Backend,const D:usize>(&self,value:&Tensor<B,D>,name:&'static str) -> Result<(),OptimizerShardError> {
self.validate()?;
if value.dims().as_slice() != self.global_shape.as_slice() {return Err(OptimizerShardError::Shape(name));}
if !value.dtype().is_float() {return Err(OptimizerShardError::DType("native moment/master storage must remain floating"));}
Ok(())
}
fn slice<B:Backend,const D:usize>(&self,value:Tensor<B,D>) -> Tensor<B,D> {
if self.interval == (0..self.global_shape[self.axis]) {value} else {value.slice_dim(self.axis,self.interval.clone())}
}
}
impl<B:Backend,const D:usize> AdaptiveMomentumState<B,D> {
pub fn try_into_shard(self,shard:&OptimizerTensorShard) -> Result<Self,OptimizerShardError> {
shard.check(&self.moment_1,"first moment")?;shard.check(&self.moment_2,"second moment")?;
check_buffer(&self.moment_2,&self.moment_1,"second moment")?;
if let Some(maximum) = &self.max_moment_2 {shard.check(maximum,"AMSGrad maximum")?;check_buffer(maximum,&self.moment_1,"AMSGrad maximum")?;}
Ok(Self {time:self.time,moment_1:shard.slice(self.moment_1),moment_2:shard.slice(self.moment_2),max_moment_2:self.max_moment_2.map(|value|shard.slice(value))})
}
}
fn check_buffer<B:Backend,const D:usize>(value:&Tensor<B,D>,reference:&Tensor<B,D>,name:&'static str) -> Result<(),OptimizerShardError> {
if value.dtype() != reference.dtype() {return Err(OptimizerShardError::DType(name));}
if value.device() != reference.device() {return Err(OptimizerShardError::Device(name));}
Ok(())
}
macro_rules! adaptive_state {
($state:ident) => {
impl<B:Backend,const D:usize> $state<B,D> {
pub fn try_into_shard(self,shard:&OptimizerTensorShard) -> Result<Self,OptimizerShardError> {
Ok(Self {momentum:self.momentum.try_into_shard(shard)?})
}
}
};
}
adaptive_state!(AdamState);
adaptive_state!(AdamWState);
impl<B:Backend,const D:usize,S:OptimizerCheckpointBuffers<B,D>> Fp32MasterState<B,D,S> {
pub fn try_into_shard(self,shard:&OptimizerTensorShard) -> Result<Self,OptimizerShardError> {
shard.check(&self.master,"FP32 master")?;
if self.master.dtype() != DType::F32 {return Err(OptimizerShardError::DType("master must remain FP32"));}
let mut buffers = Vec::new();
if let Some(inner) = &self.inner {inner.visit_checkpoint_buffers(&mut |value|buffers.push(value.clone()));}
for value in &buffers {check_buffer(value,&self.master,"native moment/master")?;}
Ok(Self {master:shard.slice(self.master),inner:self.inner.map(|state|partition_native_optimizer_state(state,shard)).transpose()?})
}
}
impl<B:Backend,const D:usize> OptimizerCheckpointBuffers<B,D> for AdaptiveMomentumState<B,D> {
fn visit_checkpoint_buffers<F:FnMut(&Tensor<B,D>)>(&self,visit:&mut F) {
visit(&self.moment_1);visit(&self.moment_2);if let Some(maximum) = &self.max_moment_2 {visit(maximum);}
}
fn map_checkpoint_buffers<F:FnMut(Tensor<B,D>)->Tensor<B,D>>(self,map:&mut F) -> Self {
Self {time:self.time,moment_1:map(self.moment_1),moment_2:map(self.moment_2),max_moment_2:self.max_moment_2.map(map)}
}
}
macro_rules! adaptive_buffers {
($state:ident) => {
impl<B:Backend,const D:usize> OptimizerCheckpointBuffers<B,D> for $state<B,D> {
fn visit_checkpoint_buffers<F:FnMut(&Tensor<B,D>)>(&self,visit:&mut F) {self.momentum.visit_checkpoint_buffers(visit);}
fn map_checkpoint_buffers<F:FnMut(Tensor<B,D>)->Tensor<B,D>>(self,map:&mut F) -> Self {Self {momentum:self.momentum.map_checkpoint_buffers(map)}}
}
};
}
adaptive_buffers!(AdamState);
adaptive_buffers!(AdamWState);
adaptive_buffers!(AdanState);
impl<B:Backend,const D:usize> OptimizerCheckpointBuffers<B,D> for AdaptiveNesterovMomentumState<B,D> {
fn visit_checkpoint_buffers<F:FnMut(&Tensor<B,D>)>(&self,visit:&mut F) {visit(&self.exp_avg);visit(&self.exp_avg_sq);visit(&self.exp_avg_diff);visit(&self.neg_pre_grad);}
fn map_checkpoint_buffers<F:FnMut(Tensor<B,D>)->Tensor<B,D>>(self,map:&mut F) -> Self {
Self {time:self.time,exp_avg:map(self.exp_avg),exp_avg_sq:map(self.exp_avg_sq),exp_avg_diff:map(self.exp_avg_diff),neg_pre_grad:map(self.neg_pre_grad)}
}
}
macro_rules! native_state {
($($state:ident),+) => {$(
impl<B:Backend,const D:usize> $state<B,D> {
pub fn try_into_shard(self,shard:&OptimizerTensorShard) -> Result<Self,OptimizerShardError> {partition_native_optimizer_state(self,shard)}
}
)+};
}
native_state!(SgdState,MomentumState,AdaGradState,LrDecayState,RmsPropState,SquareAvgState,CenteredState,RmsPropMomentumState,AdanState,AdaptiveNesterovMomentumState);
macro_rules! adaptor_partition {
($name:ident,$map:ident,$optimizer:ty) => {
pub fn $name<B:AutodiffBackend>(record:AdaptorRecord<$optimizer,B>,shard:&OptimizerTensorShard)
-> Result<AdaptorRecord<$optimizer,B>,OptimizerShardError> {
Ok(AdaptorRecord::V1(match record {AdaptorRecord::V1(record)=>match record {
AdaptorRecordV1::Rank0(value)=>AdaptorRecordV1::Rank0(value.try_into_shard(shard)?),
AdaptorRecordV1::Rank1(value)=>AdaptorRecordV1::Rank1(value.try_into_shard(shard)?),
AdaptorRecordV1::Rank2(value)=>AdaptorRecordV1::Rank2(value.try_into_shard(shard)?),
AdaptorRecordV1::Rank3(value)=>AdaptorRecordV1::Rank3(value.try_into_shard(shard)?),
AdaptorRecordV1::Rank4(value)=>AdaptorRecordV1::Rank4(value.try_into_shard(shard)?),
AdaptorRecordV1::Rank5(value)=>AdaptorRecordV1::Rank5(value.try_into_shard(shard)?),
AdaptorRecordV1::Rank6(value)=>AdaptorRecordV1::Rank6(value.try_into_shard(shard)?),
AdaptorRecordV1::Rank7(value)=>AdaptorRecordV1::Rank7(value.try_into_shard(shard)?),
AdaptorRecordV1::Rank8(value)=>AdaptorRecordV1::Rank8(value.try_into_shard(shard)?),
}}))
}
pub fn $map<B:AutodiffBackend>(records:HashMap<ParamId,AdaptorRecord<$optimizer,B>>,shards:&HashMap<ParamId,OptimizerTensorShard>)
-> Result<HashMap<ParamId,AdaptorRecord<$optimizer,B>>,OptimizerShardError> {
for shard in shards.values() {shard.validate()?;}
records.into_iter().map(|(id,record)| {
if let Some(shard) = shards.get(&id) {$name::<B>(record,shard).map(|record|(id,record))} else {Ok((id,record))}
}).collect()
}
};
}
adaptor_partition!(partition_adam_record,partition_adam_records,Adam);
adaptor_partition!(partition_adamw_record,partition_adamw_records,AdamW);
adaptor_partition!(partition_adam_master_record,partition_adam_master_records,Fp32MasterOptimizer<Adam>);
adaptor_partition!(partition_adamw_master_record,partition_adamw_master_records,Fp32MasterOptimizer<AdamW>);
adaptor_partition!(partition_sgd_record,partition_sgd_records,Sgd<B::InnerBackend>);
adaptor_partition!(partition_adagrad_record,partition_adagrad_records,AdaGrad);
adaptor_partition!(partition_rmsprop_record,partition_rmsprop_records,RmsProp);
adaptor_partition!(partition_adan_record,partition_adan_records,Adan);
adaptor_partition!(partition_sgd_master_record,partition_sgd_master_records,Fp32MasterOptimizer<Sgd<B::InnerBackend>>);
adaptor_partition!(partition_adagrad_master_record,partition_adagrad_master_records,Fp32MasterOptimizer<AdaGrad>);
adaptor_partition!(partition_rmsprop_master_record,partition_rmsprop_master_records,Fp32MasterOptimizer<RmsProp>);
adaptor_partition!(partition_adan_master_record,partition_adan_master_records,Fp32MasterOptimizer<Adan>);