Skip to main content

ruda_optim/
training.rs

1use core::marker::PhantomData;
2use alloc::string::ToString;
3use ruda_model::{
4    module::{AutodiffModule, ModuleDTypeRecord},
5    record::{PrecisionSettings, Record, Recorder, RecorderError},
6    tensor::backend::AutodiffBackend,
7};
8
9use crate::{GradientsAccumulator, GradientsParamsRecord, Optimizer, WeightedGradientsAccumulator,
10    WeightedAccumulationState, lr_scheduler::LrScheduler};
11
12#[cfg(test)]
13mod tests;
14
15mod model_state;
16pub use model_state::{ModelStateTrainingRecord,TrainableParameterContract};
17
18/// One record containing the trainable state and a caller-defined continuation record.
19///
20/// Capture at a training boundary with no concurrent updates. The caller's state
21/// carries its counters and any available input/RNG records; this type does not
22/// infer them or snapshot a DataLoader. Recorder precision settings apply to all
23/// component records and must preserve their values for exact continuation.
24pub struct TrainingRecord<B, M, O, S, U>
25where
26    B: AutodiffBackend,
27    M: AutodiffModule<B>,
28    O: Optimizer<M, B>,
29    S: LrScheduler,
30    U: Record<B>,
31{
32    model: M::Record,
33    optimizer: O::Record,
34    scheduler: S::Record<B>,
35    gradients: GradientsParamsRecord,
36    state: U,
37    marker: PhantomData<fn() -> (B, M, O, S)>,
38}
39
40/// Components restored together, including pending gradients and caller state.
41pub struct RestoredTraining<M, O, S, U> {
42    /// Model with the recorded parameter IDs.
43    pub model: M,
44    /// Optimizer with the recorded parameter state.
45    pub optimizer: O,
46    /// Scheduler at the recorded position.
47    pub scheduler: S,
48    /// Pending gradients, without an implicit optimizer update or reset.
49    pub accumulator: GradientsAccumulator<M>,
50    /// Caller-defined continuation state.
51    pub state: U,
52}
53
54/// Training components restored with actual unequal-microbatch normalization state.
55pub struct RestoredWeightedTraining<M,O,S,U> {
56    /// Model with the recorded IDs and optional per-parameter storage dtypes.
57    pub model: M,
58    /// Recorded optimizer state, without a parameter update.
59    pub optimizer: O,
60    /// Scheduler at the saved position, not advanced during restore.
61    pub scheduler: S,
62    /// Pending gradients, effective weight, microbatch count and loss scale.
63    pub accumulator: WeightedGradientsAccumulator<M>,
64    /// Caller-owned source/sampler/RNG state from the same training checkpoint.
65    pub state: U,
66}
67
68impl<B, M, O, S, U> TrainingRecord<B, M, O, S, U>
69where
70    B: AutodiffBackend,
71    M: AutodiffModule<B>,
72    O: Optimizer<M, B>,
73    S: LrScheduler,
74    U: Record<B>,
75{
76    /// Capture weighted accumulation using the existing combined record format.
77    /// Counters/options are kept with caller state, not inferred on resume.
78    pub fn capture_weighted(
79        model: &M,optimizer: &O,scheduler: &S,
80        accumulator: &WeightedGradientsAccumulator<M>,state: U,
81    ) -> Result<TrainingRecord<B,M,O,S,(WeightedAccumulationState,U)>,RecorderError> {
82        TrainingRecord::<B,M,O,S,(WeightedAccumulationState,U)>::capture(
83            model,optimizer,scheduler,accumulator.inner(),(accumulator.state().clone(),state))
84    }
85
86    /// Asynchronous combined capture of weighted pending gradients and counters.
87    pub async fn capture_weighted_async(
88        model: &M,optimizer: &O,scheduler: &S,
89        accumulator: &WeightedGradientsAccumulator<M>,state: U,
90    ) -> Result<TrainingRecord<B,M,O,S,(WeightedAccumulationState,U)>,RecorderError> {
91        TrainingRecord::<B,M,O,S,(WeightedAccumulationState,U)>::capture_async(
92            model,optimizer,scheduler,accumulator.inner(),(accumulator.state().clone(),state)).await
93    }
94
95    /// Capture weighted accumulation plus mixed floating parameter storage dtypes.
96    pub fn capture_weighted_with_dtypes(
97        model: &M,optimizer: &O,scheduler: &S,
98        accumulator: &WeightedGradientsAccumulator<M>,state: U,
99    ) -> Result<TrainingRecord<B,M,O,S,(ModuleDTypeRecord,(WeightedAccumulationState,U))>,RecorderError> {
100        TrainingRecord::<B,M,O,S,(WeightedAccumulationState,U)>::capture_with_dtypes(
101            model,optimizer,scheduler,accumulator.inner(),(accumulator.state().clone(),state))
102    }
103
104    /// Capture the components without consuming the model or clearing gradients.
105    pub fn capture(
106        model: &M,
107        optimizer: &O,
108        scheduler: &S,
109        accumulator: &GradientsAccumulator<M>,
110        state: U,
111    ) -> Result<Self, RecorderError> {
112        let gradients = accumulator.try_to_record::<B>()?;
113        Ok(Self {
114            model: model.clone().into_record(),
115            optimizer: optimizer.to_record(),
116            scheduler: scheduler.to_record::<B>(),
117            gradients,
118            state,
119            marker: PhantomData,
120        })
121    }
122
123    /// Capture with asynchronous gradient readback; no recorder I/O is performed.
124    pub async fn capture_async(
125        model: &M,
126        optimizer: &O,
127        scheduler: &S,
128        accumulator: &GradientsAccumulator<M>,
129        state: U,
130    ) -> Result<Self, RecorderError> {
131        let gradients = accumulator.to_record_async::<B>().await?;
132        Ok(Self {
133            model: model.clone().into_record(),
134            optimizer: optimizer.to_record(),
135            scheduler: scheduler.to_record::<B>(),
136            gradients,
137            state,
138            marker: PhantomData,
139        })
140    }
141
142    /// Capture with opt-in per-parameter floating storage dtype metadata.
143    ///
144    /// The metadata is recorded with caller state; ordinary captures and their
145    /// serialization format are unchanged. Use `restore_with_dtypes` to restore
146    /// mixed storage dtypes independently from the recorder's value precision.
147    pub fn capture_with_dtypes(
148        model: &M,
149        optimizer: &O,
150        scheduler: &S,
151        accumulator: &GradientsAccumulator<M>,
152        state: U,
153    ) -> Result<TrainingRecord<B, M, O, S, (ModuleDTypeRecord, U)>, RecorderError> {
154        let dtypes = ModuleDTypeRecord::capture(model)?;
155        TrainingRecord::<B, M, O, S, (ModuleDTypeRecord, U)>::capture(
156            model,
157            optimizer,
158            scheduler,
159            accumulator,
160            (dtypes, state),
161        )
162    }
163
164    /// Asynchronously capture pending gradients and per-parameter storage dtypes.
165    pub async fn capture_async_with_dtypes(
166        model: &M,
167        optimizer: &O,
168        scheduler: &S,
169        accumulator: &GradientsAccumulator<M>,
170        state: U,
171    ) -> Result<TrainingRecord<B, M, O, S, (ModuleDTypeRecord, U)>, RecorderError> {
172        let dtypes = ModuleDTypeRecord::capture(model)?;
173        TrainingRecord::<B, M, O, S, (ModuleDTypeRecord, U)>::capture_async(
174            model,
175            optimizer,
176            scheduler,
177            accumulator,
178            (dtypes, state),
179        )
180        .await
181    }
182
183    /// Save all components in one recorder payload.
184    pub fn save<R: Recorder<B>>(
185        self,
186        recorder: &R,
187        args: R::RecordArgs,
188    ) -> Result<R::RecordOutput, RecorderError> {
189        recorder.record(self, args)
190    }
191
192    /// Read a combined record using the selected recorder and device.
193    pub fn load<R: Recorder<B>>(
194        recorder: &R,
195        args: R::LoadArgs,
196        device: &B::Device,
197    ) -> Result<Self, RecorderError> {
198        recorder.load(args, device)
199    }
200
201    /// Restore onto compatible model, optimizer and scheduler configurations.
202    ///
203    /// No scheduler step, optimizer step or gradient reset is performed. Apply
204    /// the returned caller state before consuming the next training batch.
205    pub fn restore(
206        self,
207        model: M,
208        optimizer: O,
209        scheduler: S,
210        device: &B::Device,
211    ) -> Result<RestoredTraining<M, O, S, U>, RecorderError> {
212        let mut accumulator = GradientsAccumulator::new();
213        accumulator.load_record::<B>(self.gradients, device)?;
214        let model = model.load_record(self.model).fork(device);
215        let optimizer = optimizer.load_record(self.optimizer);
216        let scheduler = scheduler.load_record::<B>(self.scheduler);
217        Ok(RestoredTraining {
218            model,
219            optimizer,
220            scheduler,
221            accumulator,
222            state: self.state,
223        })
224    }
225}
226
227impl<B,M,O,S,U> TrainingRecord<B,M,O,S,(WeightedAccumulationState,U)>
228where B: AutodiffBackend,M: AutodiffModule<B>,O: Optimizer<M,B>,S: LrScheduler,U: Record<B> {
229    /// Restore pending gradients and their weight/loss-scale state without replay.
230    pub fn restore_weighted(
231        self,model: M,optimizer: O,scheduler: S,device: &B::Device,
232    ) -> Result<RestoredWeightedTraining<M,O,S,U>,RecorderError> {
233        restore_weighted::<B,M,O,S,U>(self.restore(model,optimizer,scheduler,device)?)
234    }
235}
236
237impl<B,M,O,S,U> TrainingRecord<B,M,O,S,(ModuleDTypeRecord,(WeightedAccumulationState,U))>
238where B: AutodiffBackend,M: AutodiffModule<B>,O: Optimizer<M,B>,S: LrScheduler,U: Record<B> {
239    /// Restore both mixed storage and the exact saved accumulation window.
240    pub fn restore_weighted_with_dtypes(
241        self,model: M,optimizer: O,scheduler: S,device: &B::Device,
242    ) -> Result<RestoredWeightedTraining<M,O,S,U>,RecorderError> {
243        restore_weighted::<B,M,O,S,U>(self.restore_with_dtypes(model,optimizer,scheduler,device)?)
244    }
245}
246
247fn restore_weighted<B,M,O,S,U>(
248    restored: RestoredTraining<M,O,S,(WeightedAccumulationState,U)>,
249) -> Result<RestoredWeightedTraining<M,O,S,U>,RecorderError>
250where B: AutodiffBackend,M: AutodiffModule<B> {
251    let (normalization,state) = restored.state;
252    let accumulator = WeightedGradientsAccumulator::from_accumulator::<B>(
253        &restored.model,restored.accumulator,normalization)
254        .map_err(|error|RecorderError::Unknown(error.to_string()))?;
255    Ok(RestoredWeightedTraining {model:restored.model,optimizer:restored.optimizer,
256        scheduler:restored.scheduler,accumulator,state})
257}
258
259impl<B, M, O, S, U> TrainingRecord<B, M, O, S, (ModuleDTypeRecord, U)>
260where
261    B: AutodiffBackend,
262    M: AutodiffModule<B>,
263    O: Optimizer<M, B>,
264    S: LrScheduler,
265    U: Record<B>,
266{
267    /// Restore captured floating storage dtypes, pending gradients and trainable state.
268    /// No update, scheduler step, gradient reset or requantization is performed.
269    pub fn restore_with_dtypes(
270        self,
271        model: M,
272        optimizer: O,
273        scheduler: S,
274        device: &B::Device,
275    ) -> Result<RestoredTraining<M, O, S, U>, RecorderError> {
276        let restored = self.restore(model, optimizer, scheduler, device)?;
277        let (dtypes, state) = restored.state;
278        Ok(RestoredTraining {
279            model: dtypes.apply(restored.model)?,
280            optimizer: restored.optimizer,
281            scheduler: restored.scheduler,
282            accumulator: restored.accumulator,
283            state,
284        })
285    }
286}
287
288impl<B, M, O, S, U> Record<B> for TrainingRecord<B, M, O, S, U>
289where
290    B: AutodiffBackend,
291    M: AutodiffModule<B>,
292    O: Optimizer<M, B>,
293    S: LrScheduler,
294    U: Record<B>,
295{
296    type Item<P: PrecisionSettings> = (
297        <M::Record as Record<B>>::Item<P>,
298        <O::Record as Record<B>>::Item<P>,
299        <S::Record<B> as Record<B>>::Item<P>,
300        <GradientsParamsRecord as Record<B>>::Item<P>,
301        U::Item<P>,
302    );
303
304    fn into_item<P: PrecisionSettings>(self) -> Self::Item<P> {
305        (
306            self.model.into_item::<P>(),
307            self.optimizer.into_item::<P>(),
308            self.scheduler.into_item::<P>(),
309            <GradientsParamsRecord as Record<B>>::into_item::<P>(self.gradients),
310            self.state.into_item::<P>(),
311        )
312    }
313
314    fn from_item<P: PrecisionSettings>(item: Self::Item<P>, device: &B::Device) -> Self {
315        Self {
316            model: <M::Record as Record<B>>::from_item::<P>(item.0, device),
317            optimizer: <O::Record as Record<B>>::from_item::<P>(item.1, device),
318            scheduler: <S::Record<B> as Record<B>>::from_item::<P>(item.2, device),
319            gradients: <GradientsParamsRecord as Record<B>>::from_item::<P>(item.3, device),
320            state: U::from_item::<P>(item.4, device),
321            marker: PhantomData,
322        }
323    }
324}