use super::*;
use alloc::{collections::{BTreeMap,BTreeSet},vec::Vec};
use core::fmt;
use ruda_model::{module::ModuleVisitor,record::{Record,PrecisionSettings}};
#[derive(Clone,Debug,PartialEq,Eq)]
pub enum GroupedOptimizerError {
Configuration(&'static str),
UnknownParameter(u64),
State(&'static str),
}
impl fmt::Display for GroupedOptimizerError {
fn fmt(&self,f:&mut fmt::Formatter<'_>) -> fmt::Result {
match self {Self::Configuration(value)=>write!(f,"native optimizer groups: {value}"),
Self::UnknownParameter(id)=>write!(f,"unknown selected native floating parameter {id}"),
Self::State(value)=>write!(f,"native optimizer group record: {value}")}
}
}
impl core::error::Error for GroupedOptimizerError {}
#[derive(Clone)]
pub struct GroupedOptimizerAdaptorRecord<O,B>
where B:AutodiffBackend,O:SimpleOptimizer<B::InnerBackend> {
version:u32,
routes:Vec<(u64,usize)>,
records:Vec<HashMap<ParamId,AdaptorRecord<O,B>>>,
}
impl<O,B> GroupedOptimizerAdaptorRecord<O,B>
where B:AutodiffBackend,O:SimpleOptimizer<B::InnerBackend> {
pub fn routes(&self) -> &[(u64,usize)] {&self.routes}
pub fn group_records(&self) -> &[HashMap<ParamId,AdaptorRecord<O,B>>] {&self.records}
}
impl<O,B> Record<B> for GroupedOptimizerAdaptorRecord<O,B>
where B:AutodiffBackend,O:SimpleOptimizer<B::InnerBackend> {
type Item<P:PrecisionSettings>=(u32,Vec<(u64,usize)>,Vec<<HashMap<ParamId,AdaptorRecord<O,B>> as Record<B>>::Item<P>>);
fn into_item<P:PrecisionSettings>(self) -> Self::Item<P> {
(self.version,self.routes,self.records.into_iter().map(|record|record.into_item::<P>()).collect())
}
fn from_item<P:PrecisionSettings>(item:Self::Item<P>,device:&B::Device) -> Self {
Self {version:item.0,routes:item.1,records:item.2.into_iter()
.map(|record|HashMap::<ParamId,AdaptorRecord<O,B>>::from_item::<P>(record,device)).collect()}
}
}
#[derive(Clone)]
pub struct GroupedOptimizerAdaptor<O,M,B>
where B:AutodiffBackend,M:AutodiffModule<B>,O:SimpleOptimizer<B::InnerBackend> {
groups:Vec<OptimizerAdaptor<O,M,B>>,
routes:BTreeMap<ParamId,usize>,
selection:Vec<Vec<ParamId>>,
}
impl<O,M,B> GroupedOptimizerAdaptor<O,M,B>
where B:AutodiffBackend,M:AutodiffModule<B>,O:SimpleOptimizer<B::InnerBackend> {
pub fn new(module:&M,groups:Vec<OptimizerAdaptor<O,M,B>>,routes:&[(ParamId,usize)]) -> Result<Self,GroupedOptimizerError> {
if groups.is_empty() {return Err(GroupedOptimizerError::Configuration("at least one original native optimizer group required"));}
let mut routing=BTreeMap::new();let mut selection=alloc::vec![Vec::new();groups.len()];
for &(id,index) in routes {
if index>=groups.len() || routing.insert(id,index).is_some() {
return Err(GroupedOptimizerError::Configuration("duplicate canonical parameter route or unknown group"));
}
selection[index].push(id);
}
for ids in &mut selection {ids.sort();}
let result=Self {groups,routes:routing,selection};result.validate_model(module)?;
result.validate_histories(&result.groups.iter().map(|group|&group.records).collect::<Vec<_>>())?;
Ok(result)
}
pub fn groups(&self) -> &[OptimizerAdaptor<O,M,B>] {&self.groups}
pub fn group_for(&self,parameter:ParamId) -> Option<usize> {self.routes.get(¶meter).copied()}
pub fn state_parameter_count(&self) -> usize {self.groups.iter().map(|group|group.records.len()).sum()}
fn route_record(&self) -> Vec<(u64,usize)> {self.routes.iter().map(|(id,index)|(id.val(),*index)).collect()}
pub fn validate_model(&self,module:&M) -> Result<(),GroupedOptimizerError> {
struct Ids {found:BTreeSet<ParamId>}
impl<B:AutodiffBackend> ModuleVisitor<B> for Ids {
fn visit_float<const D:usize>(&mut self,param:&Param<Tensor<B,D>>) {self.found.insert(param.id);}
}
let mut ids=Ids {found:BTreeSet::new()};module.visit(&mut ids);
for id in self.routes.keys() {if !ids.found.contains(id) {return Err(GroupedOptimizerError::UnknownParameter(id.val()));}}
Ok(())
}
fn validate_histories(&self,records:&[&HashMap<ParamId,AdaptorRecord<O,B>>]) -> Result<(),GroupedOptimizerError> {
if records.len()!=self.groups.len() {return Err(GroupedOptimizerError::State("original group count differs"));}
for (index,records) in records.iter().enumerate() {
if records.keys().any(|id|self.routes.get(id)!=Some(&index)) {
return Err(GroupedOptimizerError::State("native history belongs to another or unselected parameter group"));
}
}
Ok(())
}
fn step_with(&mut self,rates:&[LearningRate],mut module:M,mut gradients:GradAdaptor)
-> Result<(M,GradAdaptor),GroupedOptimizerError> {
if rates.len()!=self.groups.len() || rates.iter().any(|rate|!rate.is_finite() || *rate<0.0) {
return Err(GroupedOptimizerError::Configuration("one finite nonnegative learning rate per original group required"));
}
self.validate_model(&module)?;
for ((group,selection),rate) in self.groups.iter_mut().zip(&self.selection).zip(rates) {
let mut mapper=SimpleOptimizerMapper::<B,O>::new(&group.optim,&mut group.records,&mut gradients,*rate,group.grad_clipping.as_ref());
mapper.selection=Some(selection);module=module.map(&mut mapper);
}
Ok((module,gradients))
}
pub fn try_step_with_lrs(&mut self,rates:&[LearningRate],module:M,gradients:GradientsParams)
-> Result<(M,GradientsParams),GroupedOptimizerError> {
let (module,gradients)=self.step_with(rates,module,gradients.into())?;
let GradAdaptor::Single(gradients)=gradients else {unreachable!("original single native gradient input")};
Ok((module,gradients))
}
pub fn try_step_multi_with_lrs(&mut self,rates:&[LearningRate],module:M,gradients:MultiGradientsParams)
-> Result<(M,MultiGradientsParams),GroupedOptimizerError> {
let (module,gradients)=self.step_with(rates,module,gradients.into())?;
let GradAdaptor::Multi(gradients)=gradients else {unreachable!("original multi-device native gradient input")};
Ok((module,gradients))
}
pub fn try_load_record(mut self,record:GroupedOptimizerAdaptorRecord<O,B>) -> Result<Self,GroupedOptimizerError> {
if record.version!=1 || record.routes!=self.route_record() {return Err(GroupedOptimizerError::State("saved original group routing differs"));}
self.validate_histories(&record.records.iter().collect::<Vec<_>>())?;
for (group,records) in self.groups.iter_mut().zip(record.records) {group.records=records;}
Ok(self)
}
pub fn try_load_record_for_model(mut self,module:&M,record:GroupedOptimizerAdaptorRecord<O,B>)
-> Result<Self,OptimizerStatePlacementError> {
if record.version!=1 || record.routes!=self.route_record() {
return Err(OptimizerStatePlacementError::Group(GroupedOptimizerError::State("saved original group routing differs")));
}
self.validate_model(module).map_err(OptimizerStatePlacementError::Group)?;
self.validate_histories(&record.records.iter().collect::<Vec<_>>()).map_err(OptimizerStatePlacementError::Group)?;
let devices=record.records.iter().map(|records|super::placement::record_devices::<B,M,O>(module,records))
.collect::<Result<Vec<_>,_>>()?;
for ((group,records),devices) in self.groups.iter_mut().zip(record.records).zip(devices) {
group.records=records.into_iter().map(|(id,record)| {
let device=devices.get(&id).expect("validated original native group placement");(id,record.to_device(device))
}).collect();
}
Ok(self)
}
}
impl<O,M,B> Optimizer<M,B> for GroupedOptimizerAdaptor<O,M,B>
where B:AutodiffBackend,M:AutodiffModule<B>,O:SimpleOptimizer<B::InnerBackend> {
type Record=GroupedOptimizerAdaptorRecord<O,B>;
fn step(&mut self,lr:LearningRate,module:M,gradients:GradientsParams) -> M {
self.try_step_with_lrs(&alloc::vec![lr;self.groups.len()],module,gradients).unwrap_or_else(|error|panic!("{error}")).0
}
fn step_multi(&mut self,lr:LearningRate,module:M,gradients:MultiGradientsParams) -> M {
self.try_step_multi_with_lrs(&alloc::vec![lr;self.groups.len()],module,gradients).unwrap_or_else(|error|panic!("{error}")).0
}
fn to_record(&self) -> Self::Record {
Self::Record {version:1,routes:self.route_record(),records:self.groups.iter().map(|group|group.records.clone()).collect()}
}
fn load_record(self,record:Self::Record) -> Self {self.try_load_record(record).unwrap_or_else(|error|panic!("{error}"))}
}