use super::*;
pub struct InnerBackendRecord<B: AutodiffBackend, R: Record<B::InnerBackend>> {
inner: R,
marker: PhantomData<B>,
}
impl<B: AutodiffBackend, R: Record<B::InnerBackend>> InnerBackendRecord<B, R> {
pub fn new(inner: R) -> Self { Self { inner, marker: PhantomData } }
pub fn into_inner(self) -> R { self.inner }
pub fn inner(&self) -> &R { &self.inner }
}
impl<B: AutodiffBackend, R: Record<B::InnerBackend>> Record<B> for InnerBackendRecord<B, R> {
type Item<P: PrecisionSettings> = R::Item<P>;
fn into_item<P: PrecisionSettings>(self) -> Self::Item<P> { self.inner.into_item::<P>() }
fn from_item<P: PrecisionSettings>(item: Self::Item<P>, device: &B::Device) -> Self {
Self::new(R::from_item::<P>(item, device))
}
}
pub struct ModelGroupTrainingRecord<B, M, S, R, G, U>
where
B: AutodiffBackend, M: AutodiffModule<B>, S: LrScheduler,
R: Record<B>, G: Record<B>, U: Record<B>,
{
version: u32,
model_state: R,
contract: TrainableParameterContract,
groups: G,
scheduler: S::Record<B>,
gradients: GradientsParamsRecord,
state: U,
marker: PhantomData<fn() -> (B, M, S)>,
}
impl<B, M, S, R, G, U> ModelGroupTrainingRecord<B, M, S, R, G, U>
where
B: AutodiffBackend, M: AutodiffModule<B>, S: LrScheduler,
R: Record<B>, G: Record<B>, U: Record<B>,
{
pub fn capture(
model: &M, model_state: R, groups: G, scheduler: &S,
accumulator: &GradientsAccumulator<M>, state: U,
) -> Result<Self, RecorderError> {
check_pending::<B, M>(model, accumulator)?;
let contract = TrainableParameterContract::capture::<B, M>(model)?;
let gradients = accumulator.try_to_record::<B>()?;
Ok(Self { version: 1, model_state, contract, groups,
scheduler: scheduler.to_record::<B>(), gradients, state, marker: PhantomData })
}
pub async fn capture_async(
model: &M, model_state: R, groups: G, scheduler: &S,
accumulator: &GradientsAccumulator<M>, state: U,
) -> Result<Self, RecorderError> {
check_pending::<B, M>(model, accumulator)?;
let contract = TrainableParameterContract::capture::<B, M>(model)?;
let gradients = accumulator.to_record_async::<B>().await?;
Ok(Self { version: 1, model_state, contract, groups,
scheduler: scheduler.to_record::<B>(), gradients, state, marker: PhantomData })
}
pub fn capture_weighted(
model: &M, model_state: R, groups: G, scheduler: &S,
accumulator: &WeightedGradientsAccumulator<M>, state: U,
) -> Result<ModelGroupTrainingRecord<B, M, S, R, G, (WeightedAccumulationState, U)>, RecorderError> {
ModelGroupTrainingRecord::capture(model, model_state, groups, scheduler,
accumulator.inner(), (accumulator.state().clone(), state))
}
pub async fn capture_weighted_async(
model: &M, model_state: R, groups: G, scheduler: &S,
accumulator: &WeightedGradientsAccumulator<M>, state: U,
) -> Result<ModelGroupTrainingRecord<B, M, S, R, G, (WeightedAccumulationState, U)>, RecorderError> {
ModelGroupTrainingRecord::capture_async(model, model_state, groups, scheduler,
accumulator.inner(), (accumulator.state().clone(), state)).await
}
pub fn save<C: Recorder<B>>(self, recorder: &C, args: C::RecordArgs)
-> Result<C::RecordOutput, RecorderError> { recorder.record(self, args) }
pub fn load<C: Recorder<B>>(recorder: &C, args: C::LoadArgs, device: &B::Device)
-> Result<Self, RecorderError> { recorder.load(args, device) }
pub fn restore<F, H, Q>(
self, model: M, scheduler: S, restore_model: F, restore_groups: H,
) -> Result<RestoredTraining<M, Q, S, U>, RecorderError>
where
F: FnOnce(R, M) -> Result<M, RecorderError>,
H: FnOnce(M, G) -> Result<(M, Q), RecorderError>,
{
if self.version != 1 { return Err(invalid("unsupported optimizer-group format version")); }
let model = restore_model(self.model_state, model)?;
self.contract.validate_for::<B, M>(&model)?;
let (model, optimizer) = restore_groups(model, self.groups)?;
self.contract.validate_for::<B, M>(&model)?;
let mut accumulator = GradientsAccumulator::new();
accumulator.load_record_for_model::<B>(self.gradients, &model)?;
check_pending::<B, M>(&model, &accumulator)?;
Ok(RestoredTraining { model, optimizer, scheduler: scheduler.load_record::<B>(self.scheduler),
accumulator, state: self.state })
}
}
impl<B, M, S, R, G, U> ModelGroupTrainingRecord<B, M, S, R, G, (WeightedAccumulationState, U)>
where
B: AutodiffBackend, M: AutodiffModule<B>, S: LrScheduler,
R: Record<B>, G: Record<B>, U: Record<B>,
{
pub fn restore_weighted<F, H, Q>(
self, model: M, scheduler: S, restore_model: F, restore_groups: H,
) -> Result<RestoredWeightedTraining<M, Q, S, U>, RecorderError>
where
F: FnOnce(R, M) -> Result<M, RecorderError>,
H: FnOnce(M, G) -> Result<(M, Q), RecorderError>,
{
super::super::restore_weighted::<B, M, Q, S, U>(
self.restore(model, scheduler, restore_model, restore_groups)?)
}
}
impl<B, M, S, R, G, U> Record<B> for ModelGroupTrainingRecord<B, M, S, R, G, U>
where
B: AutodiffBackend, M: AutodiffModule<B>, S: LrScheduler,
R: Record<B>, G: Record<B>, U: Record<B>,
{
type Item<P: PrecisionSettings> = (u32, R::Item<P>, <TrainableParameterContract as Record<B>>::Item<P>,
G::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.version, self.model_state.into_item::<P>(), <TrainableParameterContract as Record<B>>::into_item::<P>(self.contract),
self.groups.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 { version: item.0, model_state: R::from_item::<P>(item.1, device),
contract: <TrainableParameterContract as Record<B>>::from_item::<P>(item.2, device),
groups: G::from_item::<P>(item.3, device), scheduler: <S::Record<B> as Record<B>>::from_item::<P>(item.4, device),
gradients: <GradientsParamsRecord as Record<B>>::from_item::<P>(item.5, device),
state: U::from_item::<P>(item.6, device), marker: PhantomData }
}
}