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 */ }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>,
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>,
Sourcepub fn capture(
model: &M,
model_state: R,
optimizer: &O,
scheduler: &S,
accumulator: &GradientsAccumulator<M>,
state: U,
) -> Result<Self, RecorderError>
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.
Sourcepub async fn capture_async(
model: &M,
model_state: R,
optimizer: &O,
scheduler: &S,
accumulator: &GradientsAccumulator<M>,
state: U,
) -> Result<Self, RecorderError>
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.
Sourcepub 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>
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.
Sourcepub 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>
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.
Sourcepub fn save<C: Recorder<B>>(
self,
recorder: &C,
args: C::RecordArgs,
) -> Result<C::RecordOutput, RecorderError>
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.
Sourcepub fn load<C: Recorder<B>>(
recorder: &C,
args: C::LoadArgs,
device: &B::Device,
) -> Result<Self, RecorderError>
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.
Sourcepub fn restore<F>(
self,
model: M,
optimizer: O,
scheduler: S,
device: &B::Device,
restore_model: F,
) -> Result<RestoredTraining<M, O, S, U>, RecorderError>
pub fn restore<F>( self, model: M, optimizer: O, scheduler: S, device: &B::Device, restore_model: F, ) -> Result<RestoredTraining<M, O, S, U>, 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>,
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>,
Sourcepub fn restore_weighted<F>(
self,
model: M,
optimizer: O,
scheduler: S,
device: &B::Device,
restore_model: F,
) -> Result<RestoredWeightedTraining<M, O, S, U>, RecorderError>
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>
Restore actual A/B/model state and the weighted accumulation window together.