use core::marker::PhantomData;
use ruda_model::{
module::AutodiffModule,
record::{PrecisionSettings, Record, Recorder, RecorderError},
tensor::backend::AutodiffBackend,
};
use crate::{GradientsAccumulator, GradientsParamsRecord, Optimizer, lr_scheduler::LrScheduler};
#[cfg(test)]
mod tests;
pub struct TrainingRecord<B, M, O, S, U>
where
B: AutodiffBackend,
M: AutodiffModule<B>,
O: Optimizer<M, B>,
S: LrScheduler,
U: Record<B>,
{
model: M::Record,
optimizer: O::Record,
scheduler: S::Record<B>,
gradients: GradientsParamsRecord,
state: U,
marker: PhantomData<fn() -> (B, M, O, S)>,
}
pub struct RestoredTraining<M, O, S, U> {
pub model: M,
pub optimizer: O,
pub scheduler: S,
pub accumulator: GradientsAccumulator<M>,
pub state: U,
}
impl<B, M, O, S, U> TrainingRecord<B, M, O, S, U>
where
B: AutodiffBackend,
M: AutodiffModule<B>,
O: Optimizer<M, B>,
S: LrScheduler,
U: Record<B>,
{
pub fn capture(
model: &M,
optimizer: &O,
scheduler: &S,
accumulator: &GradientsAccumulator<M>,
state: U,
) -> Result<Self, RecorderError> {
let gradients = accumulator.try_to_record::<B>()?;
Ok(Self {
model: model.clone().into_record(),
optimizer: optimizer.to_record(),
scheduler: scheduler.to_record::<B>(),
gradients,
state,
marker: PhantomData,
})
}
pub async fn capture_async(
model: &M,
optimizer: &O,
scheduler: &S,
accumulator: &GradientsAccumulator<M>,
state: U,
) -> Result<Self, RecorderError> {
let gradients = accumulator.to_record_async::<B>().await?;
Ok(Self {
model: model.clone().into_record(),
optimizer: optimizer.to_record(),
scheduler: scheduler.to_record::<B>(),
gradients,
state,
marker: PhantomData,
})
}
pub fn save<R: Recorder<B>>(
self,
recorder: &R,
args: R::RecordArgs,
) -> Result<R::RecordOutput, RecorderError> {
recorder.record(self, args)
}
pub fn load<R: Recorder<B>>(
recorder: &R,
args: R::LoadArgs,
device: &B::Device,
) -> Result<Self, RecorderError> {
recorder.load(args, device)
}
pub fn restore(
self,
model: M,
optimizer: O,
scheduler: S,
device: &B::Device,
) -> Result<RestoredTraining<M, O, S, U>, RecorderError> {
let mut accumulator = GradientsAccumulator::new();
accumulator.load_record::<B>(self.gradients, device)?;
let model = model.load_record(self.model).fork(device);
let optimizer = optimizer.load_record(self.optimizer);
let scheduler = scheduler.load_record::<B>(self.scheduler);
Ok(RestoredTraining {
model,
optimizer,
scheduler,
accumulator,
state: self.state,
})
}
}
impl<B, M, O, S, U> Record<B> for TrainingRecord<B, M, O, S, U>
where
B: AutodiffBackend,
M: AutodiffModule<B>,
O: Optimizer<M, B>,
S: LrScheduler,
U: Record<B>,
{
type Item<P: PrecisionSettings> = (
<M::Record as Record<B>>::Item<P>,
<O::Record as Record<B>>::Item<P>,
<S::Record<B> as Record<B>>::Item<P>,
<GradientsParamsRecord as Record<B>>::Item<P>,
U::Item<P>,
);
fn into_item<P: PrecisionSettings>(self) -> Self::Item<P> {
(
self.model.into_item::<P>(),
self.optimizer.into_item::<P>(),
self.scheduler.into_item::<P>(),
<GradientsParamsRecord as Record<B>>::into_item::<P>(self.gradients),
self.state.into_item::<P>(),
)
}
fn from_item<P: PrecisionSettings>(item: Self::Item<P>, device: &B::Device) -> Self {
Self {
model: <M::Record as Record<B>>::from_item::<P>(item.0, device),
optimizer: <O::Record as Record<B>>::from_item::<P>(item.1, device),
scheduler: <S::Record<B> as Record<B>>::from_item::<P>(item.2, device),
gradients: <GradientsParamsRecord as Record<B>>::from_item::<P>(item.3, device),
state: U::from_item::<P>(item.4, device),
marker: PhantomData,
}
}
}