use super::*;
use core::fmt;
#[derive(Debug)]
pub struct FullyShardedMuonStepFailure<M,E:fmt::Debug> {
pub module:M,
pub gradients:GradientsParams,
pub error:MuonShardedError<E>,
}
impl<M,E:fmt::Debug> fmt::Display for FullyShardedMuonStepFailure<M,E> {
fn fmt(&self,f:&mut fmt::Formatter<'_>) -> fmt::Result {fmt::Display::fmt(&self.error,f)}
}
impl<M:fmt::Debug,E:fmt::Debug> core::error::Error for FullyShardedMuonStepFailure<M,E> {}
impl<M,B,C> FullyShardedMuonAdamW<M,B,C>
where B:AutodiffBackend,M:AutodiffModule<B>,C:BroadcastTensorCollective<B::InnerBackend> {
pub fn try_step_recoverable_with_lrs(&mut self,muon_lr:LearningRate,adamw_lr:LearningRate,module:M,gradients:GradientsParams)
-> Result<M,FullyShardedMuonStepFailure<M,C::Error>> {
let input=gradients.clone_native::<B::InnerBackend>();
let (states,mut mapper)=match self.prepare_step_with_lrs(muon_lr,adamw_lr,&module,gradients) {
Ok(proposed)=>proposed,
Err(error)=>return Err(FullyShardedMuonStepFailure {module,gradients:input,error}),
};
let module=module.map(&mut mapper);self.states=states;Ok(module)
}
}