use super::*;
use crate::{FullyShardedAccumulationContract,FullyShardedWeightedAccumulationContract,FullyShardedGradientsAccumulator,FullyShardedWeightedGradientsAccumulator};
pub struct RestoredFullyShardedTraining<M,O,S,U> {
pub model:M,
pub optimizer:O,
pub scheduler:S,
pub accumulator:FullyShardedGradientsAccumulator<M>,
pub state:U,
}
pub struct RestoredFullyShardedWeightedTraining<B:AutodiffBackend,M:AutodiffModule<B>,O,S,U> {
pub model:M,
pub optimizer:O,
pub scheduler:S,
pub accumulator:FullyShardedWeightedGradientsAccumulator<M,B>,
pub state:U,
}
impl<B,M,O,S,R,U> ModelStateTrainingRecord<B,M,O,S,R,U>
where B:AutodiffBackend,M:AutodiffModule<B>,O:Optimizer<M,B>,S:LrScheduler,R:Record<B>,U:Record<B> {
pub fn capture_fully_sharded(model:&M,model_state:R,optimizer:&O,scheduler:&S,accumulator:&FullyShardedGradientsAccumulator<M>,state:U)
-> Result<ModelStateTrainingRecord<B,M,O,S,R,(FullyShardedAccumulationContract,U)>,RecorderError> {
let continuation=accumulator.continuation();continuation.validate_for::<B,M>(model).map_err(|error|RecorderError::Unknown(error.to_string()))?;
ModelStateTrainingRecord::<B,M,O,S,R,(FullyShardedAccumulationContract,U)>::capture(model,model_state,optimizer,scheduler,accumulator.inner(),(continuation,state))
}
pub async fn capture_fully_sharded_async(model:&M,model_state:R,optimizer:&O,scheduler:&S,accumulator:&FullyShardedGradientsAccumulator<M>,state:U)
-> Result<ModelStateTrainingRecord<B,M,O,S,R,(FullyShardedAccumulationContract,U)>,RecorderError> {
let continuation=accumulator.continuation();continuation.validate_for::<B,M>(model).map_err(|error|RecorderError::Unknown(error.to_string()))?;
ModelStateTrainingRecord::<B,M,O,S,R,(FullyShardedAccumulationContract,U)>::capture_async(model,model_state,optimizer,scheduler,accumulator.inner(),(continuation,state)).await
}
pub fn capture_fully_sharded_weighted(model:&M,model_state:R,optimizer:&O,scheduler:&S,accumulator:&FullyShardedWeightedGradientsAccumulator<M,B>,state:U)
-> Result<ModelStateTrainingRecord<B,M,O,S,R,(FullyShardedWeightedAccumulationContract<B>,U)>,RecorderError> {
ModelStateTrainingRecord::<B,M,O,S,R,(FullyShardedWeightedAccumulationContract<B>,U)>::capture(model,model_state,optimizer,scheduler,accumulator.inner(),(accumulator.continuation(),state))
}
pub async fn capture_fully_sharded_weighted_async(model:&M,model_state:R,optimizer:&O,scheduler:&S,accumulator:&FullyShardedWeightedGradientsAccumulator<M,B>,state:U)
-> Result<ModelStateTrainingRecord<B,M,O,S,R,(FullyShardedWeightedAccumulationContract<B>,U)>,RecorderError> {
ModelStateTrainingRecord::<B,M,O,S,R,(FullyShardedWeightedAccumulationContract<B>,U)>::capture_async(model,model_state,optimizer,scheduler,accumulator.inner(),(accumulator.continuation(),state)).await
}
}
impl<B,M,O,S,R,U> ModelStateTrainingRecord<B,M,O,S,R,(FullyShardedAccumulationContract,U)>
where B:AutodiffBackend,M:AutodiffModule<B>,O:Optimizer<M,B>,S:LrScheduler,R:Record<B>,U:Record<B> {
pub fn restore_fully_sharded<F>(self,model:M,optimizer:O,scheduler:S,device:&B::Device,restore_model:F)
-> Result<RestoredFullyShardedTraining<M,O,S,U>,RecorderError> where F:FnOnce(R,M)->Result<M,RecorderError> {
let restored=self.restore(model,optimizer,scheduler,device,restore_model)?;let (continuation,state)=restored.state;
let accumulator=FullyShardedGradientsAccumulator::from_accumulator::<B>(&restored.model,restored.accumulator,continuation)
.map_err(|error|RecorderError::Unknown(error.to_string()))?;
Ok(RestoredFullyShardedTraining {model:restored.model,optimizer:restored.optimizer,scheduler:restored.scheduler,accumulator,state})
}
}
impl<B,M,O,S,R,U> ModelStateTrainingRecord<B,M,O,S,R,(FullyShardedWeightedAccumulationContract<B>,U)>
where B:AutodiffBackend,M:AutodiffModule<B>,O:Optimizer<M,B>,S:LrScheduler,R:Record<B>,U:Record<B> {
pub fn restore_fully_sharded_weighted<F>(self,model:M,optimizer:O,scheduler:S,device:&B::Device,restore_model:F)
-> Result<RestoredFullyShardedWeightedTraining<B,M,O,S,U>,RecorderError> where F:FnOnce(R,M)->Result<M,RecorderError> {
let restored=self.restore(model,optimizer,scheduler,device,restore_model)?;let (continuation,state)=restored.state;
let accumulator=FullyShardedWeightedGradientsAccumulator::from_accumulator(&restored.model,restored.accumulator,continuation)
.map_err(|error|RecorderError::Unknown(error.to_string()))?;
Ok(RestoredFullyShardedWeightedTraining {model:restored.model,optimizer:restored.optimizer,scheduler:restored.scheduler,accumulator,state})
}
}