use super::*;
use super::sharded::ShardedReductions;
use hashbrown::HashMap;
use ruda_model::{record::PrecisionSettings,tensor::BroadcastTensorCollective};
mod partition;
pub use partition::LBFGSMasterTensorShard;
#[derive(Clone,Debug,PartialEq,Eq,Serialize,Deserialize)]
pub struct LBFGSMasterParameter {
pub id:u64,
pub shape:Vec<usize>,
pub storage:DType,
}
fn parameter_length(parameters:&[LBFGSMasterParameter]) -> Result<usize,LBFGSShardError> {
let mut seen = HashSet::new();
parameters.iter().try_fold(0usize,|total,parameter| {
if !seen.insert(parameter.id) {return Err(LBFGSShardError::Record);}
if !matches!(parameter.storage,DType::F16|DType::BF16|DType::F32) {return Err(LBFGSShardError::DType);}
let length = parameter.shape.iter().try_fold(1usize,|total,axis|total.checked_mul(*axis))
.ok_or(LBFGSShardError::Shape("master parameter shape overflows"))?;
total.checked_add(length).ok_or(LBFGSShardError::Shape("master vector length overflows"))
})
}
struct MasterParameters<B:AutodiffBackend> {
parameters:Vec<LBFGSMasterParameter>,
tensors:Vec<Tensor<B::InnerBackend,1>>,
seen:HashMap<ParamId,usize>,
device:Option<B::Device>,
materialize:bool,
error:Option<LBFGSShardError>,
}
impl<B:AutodiffBackend> ModuleVisitor<B> for MasterParameters<B> {
fn visit_float<const D:usize>(&mut self,parameter:&Param<Tensor<B,D>>) {
if self.error.is_some() {return;}
let value = parameter.val();
if !value.is_require_grad() {return;}
if !matches!(value.dtype(),DType::F16|DType::BF16|DType::F32) {
self.error = Some(LBFGSShardError::DType);return;
}
if let Some(device) = &self.device {
if &value.device() != device {self.error = Some(LBFGSShardError::Device);return;}
} else {self.device = Some(value.device());}
let metadata = LBFGSMasterParameter {id:parameter.id.val(),shape:value.dims().to_vec(),storage:value.dtype()};
if let Some(index) = self.seen.get(¶meter.id) {
if self.parameters[*index] != metadata {self.error = Some(LBFGSShardError::Record);}
return;
}
self.seen.insert(parameter.id,self.parameters.len());
self.parameters.push(metadata);
if self.materialize {
let length = value.shape().num_elements();
self.tensors.push(value.inner().cast(DType::F32).reshape([length]));
}
}
}
#[derive(Clone)]
pub struct LBFGSFp32MasterState<B:Backend> {
version:u32,
parameters:Vec<LBFGSMasterParameter>,
master:Option<Tensor<B,1>>,
optimizer:LBFGSState<B>,
placement:Option<(u32,LBFGSShardLayout)>,
}
impl<B:Backend> LBFGSFp32MasterState<B> {
pub fn master(&self) -> Option<&Tensor<B,1>> {self.master.as_ref()}
pub fn parameters(&self) -> &[LBFGSMasterParameter] {&self.parameters}
pub fn optimizer(&self) -> &LBFGSState<B> {&self.optimizer}
pub fn placement(&self) -> Option<(u32,&LBFGSShardLayout)> {
self.placement.as_ref().map(|(rank,layout)|(*rank,layout))
}
pub fn validate(&self) -> Result<(),LBFGSShardError> {
if self.version != 1 {return Err(LBFGSShardError::Record);}
let length = parameter_length(&self.parameters)?;
if let Some(master) = &self.master {
if length == 0 || master.dims() != [length] {return Err(LBFGSShardError::Shape("FP32 master vector"));}
if master.dtype() != DType::F32 {return Err(LBFGSShardError::DType);}
self.optimizer.validate_vectors(length,Some(master))?;
} else {
if length != 0 || self.placement.is_some() {return Err(LBFGSShardError::Record);}
self.optimizer.validate_vectors(0,None)?;
}
if let Some((rank,layout)) = &self.placement {
let world = u32::try_from(layout.lengths.len()).map_err(|_|LBFGSShardError::Layout("master rank count overflows"))?;
layout.validate(*rank,world)?;
if layout.lengths[*rank as usize] != length {return Err(LBFGSShardError::Shape("FP32 master shard interval"));}
}
Ok(())
}
pub fn to_device(mut self,device:&B::Device) -> Self {
self.master = self.master.map(|master|master.to_device(device));
self.optimizer = self.optimizer.to_device(device);self
}
}
impl<B:Backend> Record<B> for LBFGSFp32MasterState<B> {
type Item<P:PrecisionSettings> = (u32,Vec<LBFGSMasterParameter>,Option<(u32,Vec<usize>)>,
<(Option<Tensor<B,1>>,LBFGSState<B>) as Record<B>>::Item<P>);
fn into_item<P:PrecisionSettings>(self) -> Self::Item<P> {
(self.version,self.parameters,self.placement.map(|(rank,layout)|(rank,layout.lengths)),
(self.master,self.optimizer).into_item::<P>())
}
fn from_item<P:PrecisionSettings>(item:Self::Item<P>,device:&B::Device) -> Self {
let (master,optimizer) = <(Option<Tensor<B,1>>,LBFGSState<B>) as Record<B>>::from_item::<P>(item.3,device);
Self {version:item.0,parameters:item.1,placement:item.2.map(|(rank,lengths)|(rank,LBFGSShardLayout::new(lengths))),master,optimizer}
}
}
#[derive(Clone)]
pub struct LBFGSFp32Master<B:AutodiffBackend> {
optimizer:LBFGS<B>,
master:Option<Tensor<B::InnerBackend,1>>,
parameters:Vec<LBFGSMasterParameter>,
placement:Option<(u32,LBFGSShardLayout)>,
}
impl LBFGSConfig {
pub fn init_fp32_master<B:AutodiffBackend>(&self) -> LBFGSFp32Master<B> {
LBFGSFp32Master {optimizer:self.init(),master:None,parameters:Vec::new(),placement:None}
}
}
impl<B:AutodiffBackend> LBFGSFp32Master<B> {
pub fn to_record(&self) -> LBFGSFp32MasterState<B::InnerBackend> {
LBFGSFp32MasterState {version:1,parameters:self.parameters.clone(),master:self.master.clone(),
optimizer:self.optimizer.to_record(),placement:self.placement.clone()}
}
pub fn load_record(mut self,record:LBFGSFp32MasterState<B::InnerBackend>) -> Result<Self,LBFGSShardError> {
record.validate()?;self.optimizer = self.optimizer.load_record(record.optimizer);
self.master = record.master;self.parameters = record.parameters;self.placement = record.placement;Ok(self)
}
pub fn to_device(mut self,device:&B::Device) -> Self {
self.optimizer = self.optimizer.to_device(device);self.master = self.master.map(|master|master.to_device(device));self
}
fn prepare<M:Module<B>>(&self,module:&M) -> Result<(Vec<LBFGSMasterParameter>,Tensor<B::InnerBackend,1>),LBFGSShardError> {
let mut visitor = MasterParameters::<B> {parameters:Vec::new(),tensors:Vec::new(),seen:HashMap::new(),device:None,
materialize:self.master.is_none(),error:None};
module.visit(&mut visitor);
if let Some(error) = visitor.error {return Err(error);}
let length = parameter_length(&visitor.parameters)?;
if length == 0 {return Err(LBFGSShardError::Shape("FP32 master requires a nonempty trainable vector"));}
let master = if let Some(master) = &self.master {
if self.parameters != visitor.parameters {return Err(LBFGSShardError::Record);}
if master.dims() != [length] {return Err(LBFGSShardError::Shape("FP32 master parameter length"));}
if master.dtype() != DType::F32 {return Err(LBFGSShardError::DType);}
if Some(master.device()) != visitor.device {return Err(LBFGSShardError::Device);}
master.clone()
} else {Tensor::cat(visitor.tensors,0)};
self.optimizer.state.validate_vectors(length,Some(&master))?;
Ok((visitor.parameters,master))
}
pub fn step<M,F>(&mut self,lr:LearningRate,module:M,mut closure:F) -> Result<(M,f64),LBFGSShardError>
where M:AutodiffModule<B>+Clone,F:FnMut(M)->(f64,GradientsParams) {
if self.placement.is_some() {return Err(LBFGSShardError::Record);}
let (parameters,master) = self.prepare(&module)?;
let (model,loss,master) = self.optimizer.try_step_with_reductions(lr,module,|model|Ok(closure(model)),
&mut LocalReductions,Some(master),true).unwrap_or_else(|error|match error {});
self.master = master;self.parameters = parameters;Ok((model,loss))
}
pub fn step_sharded<M,F,C>(&mut self,lr:LearningRate,module:M,mut closure:F,layout:&LBFGSShardLayout,communicator:&C)
-> Result<(M,f64),LBFGSShardedError<C::Error>>
where M:AutodiffModule<B>+Clone,F:FnMut(M)->(f64,GradientsParams),C:BroadcastTensorCollective<B::InnerBackend> {
self.try_step_sharded(lr,module,|model|Ok(closure(model)),layout,communicator)
}
pub fn try_step_sharded<M,F,C>(&mut self,lr:LearningRate,module:M,closure:F,layout:&LBFGSShardLayout,communicator:&C)
-> Result<(M,f64),LBFGSShardedError<C::Error>>
where M:AutodiffModule<B>+Clone,F:FnMut(M)->Result<(f64,GradientsParams),LBFGSShardedError<C::Error>>,
C:BroadcastTensorCollective<B::InnerBackend> {
layout.validate(communicator.rank(),communicator.world_size())?;
let placement = Some((communicator.rank(),layout.clone()));
if self.master.is_some() && self.placement != placement {return Err(LBFGSShardError::Record.into());}
let (parameters,master) = self.prepare(&module)?;
if master.dims() != [layout.lengths[communicator.rank() as usize]] {
return Err(LBFGSShardError::Shape("FP32 master local shard length").into());
}
let (model,loss,master) = self.optimizer.try_step_with_reductions(lr,module,closure,
&mut ShardedReductions {communicator},Some(master),true)?;
self.master = master;self.parameters = parameters;self.placement = placement;Ok((model,loss))
}
}