Skip to main content

ModelStateTrainingRecord

Struct ModelStateTrainingRecord 

Source
pub struct 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>,
{ /* private fields */ }
Available on crate feature std only.
Expand description

Caller-selected model state plus actual optimizer/scheduler/pending gradients.

R can be RUDA’s A/B-only native adapter record. It must restore every trainable parameter, including IDs/dtypes, and identify any omitted frozen state. Capture R and these components at the same boundary, without concurrent updates. No full model record or hidden frozen tensor copy is created here.

Implementations§

Source§

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>,

Source

pub fn capture( model: &M, model_state: R, optimizer: &O, scheduler: &S, accumulator: &GradientsAccumulator<M>, state: U, ) -> Result<Self, RecorderError>

Capture provided actual model state, optimizer, scheduler and pending values.

Source

pub async fn capture_async( model: &M, model_state: R, optimizer: &O, scheduler: &S, accumulator: &GradientsAccumulator<M>, state: U, ) -> Result<Self, RecorderError>

Same capture with asynchronous readback of actual pending gradients.

Source

pub fn capture_weighted( model: &M, model_state: R, optimizer: &O, scheduler: &S, accumulator: &WeightedGradientsAccumulator<M>, state: U, ) -> Result<ModelStateTrainingRecord<B, M, O, S, R, (WeightedAccumulationState, U)>, RecorderError>

Capture actual weighted-window counts/scale alongside caller source/RNG state.

Source

pub async fn capture_weighted_async( model: &M, model_state: R, optimizer: &O, scheduler: &S, accumulator: &WeightedGradientsAccumulator<M>, state: U, ) -> Result<ModelStateTrainingRecord<B, M, O, S, R, (WeightedAccumulationState, U)>, RecorderError>

Asynchronous weighted capture, with no implicit window reset or update.

Source

pub fn save<C: Recorder<B>>( self, recorder: &C, args: C::RecordArgs, ) -> Result<C::RecordOutput, RecorderError>

Save provided model state and live training components in one recorder payload.

Source

pub fn load<C: Recorder<B>>( recorder: &C, args: C::LoadArgs, device: &B::Device, ) -> Result<Self, RecorderError>

Read the combined state on the requested device; recreate omitted base state separately.

Source

pub fn restore<F>( self, model: M, optimizer: O, scheduler: S, device: &B::Device, restore_model: F, ) -> Result<RestoredTraining<M, O, S, U>, RecorderError>
where F: FnOnce(R, M) -> Result<M, RecorderError>,

Restore caller-selected model state first, then validate actual trainable identities. restore_model may call a native adapter record’s restore_into with an independently supplied frozen base identity. The prepared model must already use the requested device; the omitted base is never copied or moved implicitly.

Source§

impl<B, M, O, S, R, U> ModelStateTrainingRecord<B, M, O, S, R, (WeightedAccumulationState, U)>
where B: AutodiffBackend, M: AutodiffModule<B>, O: Optimizer<M, B>, S: LrScheduler, R: Record<B>, U: Record<B>,

Source

pub fn restore_weighted<F>( self, model: M, optimizer: O, scheduler: S, device: &B::Device, restore_model: F, ) -> Result<RestoredWeightedTraining<M, O, S, U>, RecorderError>
where F: FnOnce(R, M) -> Result<M, RecorderError>,

Restore actual A/B/model state and the weighted accumulation window together.

Trait Implementations§

Source§

impl<B, M, O, S, R, U> Record<B> for 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>,

Source§

type Item<P: PrecisionSettings> = (u32, <R as Record<B>>::Item<P>, <TrainableParameterContract as Record<B>>::Item<P>, <<O as Optimizer<M, B>>::Record as Record<B>>::Item<P>, <<S as LrScheduler>::Record<B> as Record<B>>::Item<P>, <GradientsParamsRecord as Record<B>>::Item<P>, <U as Record<B>>::Item<P>)

Type of the item that can be serialized and deserialized.
Source§

fn into_item<P: PrecisionSettings>(self) -> Self::Item<P>

Convert the current record into the corresponding item that follows the given settings.
Source§

fn from_item<P: PrecisionSettings>( item: Self::Item<P>, device: &B::Device, ) -> Self

Convert the given item into a record.

Auto Trait Implementations§

§

impl<B, M, O, S, R, U> !RefUnwindSafe for ModelStateTrainingRecord<B, M, O, S, R, U>

§

impl<B, M, O, S, R, U> !UnwindSafe for ModelStateTrainingRecord<B, M, O, S, R, U>

§

impl<B, M, O, S, R, U> Freeze for ModelStateTrainingRecord<B, M, O, S, R, U>
where R: Freeze, <O as Optimizer<M, B>>::Record: Freeze, <S as LrScheduler>::Record<B>: Freeze, U: Freeze, PhantomData<fn() -> (B, M, O, S)>: Freeze,

§

impl<B, M, O, S, R, U> Send for ModelStateTrainingRecord<B, M, O, S, R, U>
where <O as Optimizer<M, B>>::Record: Send, <S as LrScheduler>::Record<B>: Send, PhantomData<fn() -> (B, M, O, S)>: Send,

§

impl<B, M, O, S, R, U> Sync for ModelStateTrainingRecord<B, M, O, S, R, U>
where R: Sync, <O as Optimizer<M, B>>::Record: Sync, <S as LrScheduler>::Record<B>: Sync, U: Sync, PhantomData<fn() -> (B, M, O, S)>: Sync,

§

impl<B, M, O, S, R, U> Unpin for ModelStateTrainingRecord<B, M, O, S, R, U>
where R: Unpin, <O as Optimizer<M, B>>::Record: Unpin, <S as LrScheduler>::Record<B>: Unpin, U: Unpin, PhantomData<fn() -> (B, M, O, S)>: Unpin,

§

impl<B, M, O, S, R, U> UnsafeUnpin for ModelStateTrainingRecord<B, M, O, S, R, U>

Blanket Implementations§

Source§

impl<T> Any for T
where T: 'static + ?Sized,

Source§

fn type_id(&self) -> TypeId

Gets the TypeId of self. Read more
Source§

impl<T> Borrow<T> for T
where T: ?Sized,

Source§

fn borrow(&self) -> &T

Immutably borrows from an owned value. Read more
Source§

impl<T> BorrowMut<T> for T
where T: ?Sized,

Source§

fn borrow_mut(&mut self) -> &mut T

Mutably borrows from an owned value. Read more
Source§

impl<ST, DT> CastableFrom<ST, Initialized, Initialized> for DT
where ST: ?Sized, DT: ?Sized,

Source§

impl<ST, DT> CastableFrom<ST, Uninit, Uninit> for DT
where ST: ?Sized, DT: ?Sized,

Source§

impl<T> From<T> for T

Source§

fn from(t: T) -> T

Returns the argument unchanged.

Source§

impl<T, U> Into<U> for T
where U: From<T>,

Source§

fn into(self) -> U

Calls U::from(self).

That is, this conversion is whatever the implementation of From<T> for U chooses to do.

Source§

impl<T> Read<Exclusive, BecauseExclusive> for T
where T: ?Sized,

Source§

impl<T, U> TryFrom<U> for T
where U: Into<T>,

Source§

type Error = !

The type returned in the event of a conversion error.
Source§

fn try_from(value: U) -> Result<T, !>

Performs the conversion.
Source§

impl<T, U> TryInto<U> for T
where U: TryFrom<T>,

Source§

type Error = <U as TryFrom<T>>::Error

The type returned in the event of a conversion error.
Source§

fn try_into(self) -> Result<U, <U as TryFrom<T>>::Error>

Performs the conversion.