Skip to main content

ruda_optim/
training.rs

1use core::marker::PhantomData;
2use ruda_model::{
3    module::{AutodiffModule, ModuleDTypeRecord},
4    record::{PrecisionSettings, Record, Recorder, RecorderError},
5    tensor::backend::AutodiffBackend,
6};
7
8use crate::{GradientsAccumulator, GradientsParamsRecord, Optimizer, lr_scheduler::LrScheduler};
9
10#[cfg(test)]
11mod tests;
12
13/// One record containing the trainable state and a caller-defined continuation record.
14///
15/// Capture at a training boundary with no concurrent updates. The caller's state
16/// carries its counters and any available input/RNG records; this type does not
17/// infer them or snapshot a DataLoader. Recorder precision settings apply to all
18/// component records and must preserve their values for exact continuation.
19pub struct TrainingRecord<B, M, O, S, U>
20where
21    B: AutodiffBackend,
22    M: AutodiffModule<B>,
23    O: Optimizer<M, B>,
24    S: LrScheduler,
25    U: Record<B>,
26{
27    model: M::Record,
28    optimizer: O::Record,
29    scheduler: S::Record<B>,
30    gradients: GradientsParamsRecord,
31    state: U,
32    marker: PhantomData<fn() -> (B, M, O, S)>,
33}
34
35/// Components restored together, including pending gradients and caller state.
36pub struct RestoredTraining<M, O, S, U> {
37    /// Model with the recorded parameter IDs.
38    pub model: M,
39    /// Optimizer with the recorded parameter state.
40    pub optimizer: O,
41    /// Scheduler at the recorded position.
42    pub scheduler: S,
43    /// Pending gradients, without an implicit optimizer update or reset.
44    pub accumulator: GradientsAccumulator<M>,
45    /// Caller-defined continuation state.
46    pub state: U,
47}
48
49impl<B, M, O, S, U> TrainingRecord<B, M, O, S, U>
50where
51    B: AutodiffBackend,
52    M: AutodiffModule<B>,
53    O: Optimizer<M, B>,
54    S: LrScheduler,
55    U: Record<B>,
56{
57    /// Capture the components without consuming the model or clearing gradients.
58    pub fn capture(
59        model: &M,
60        optimizer: &O,
61        scheduler: &S,
62        accumulator: &GradientsAccumulator<M>,
63        state: U,
64    ) -> Result<Self, RecorderError> {
65        let gradients = accumulator.try_to_record::<B>()?;
66        Ok(Self {
67            model: model.clone().into_record(),
68            optimizer: optimizer.to_record(),
69            scheduler: scheduler.to_record::<B>(),
70            gradients,
71            state,
72            marker: PhantomData,
73        })
74    }
75
76    /// Capture with asynchronous gradient readback; no recorder I/O is performed.
77    pub async fn capture_async(
78        model: &M,
79        optimizer: &O,
80        scheduler: &S,
81        accumulator: &GradientsAccumulator<M>,
82        state: U,
83    ) -> Result<Self, RecorderError> {
84        let gradients = accumulator.to_record_async::<B>().await?;
85        Ok(Self {
86            model: model.clone().into_record(),
87            optimizer: optimizer.to_record(),
88            scheduler: scheduler.to_record::<B>(),
89            gradients,
90            state,
91            marker: PhantomData,
92        })
93    }
94
95    /// Capture with opt-in per-parameter floating storage dtype metadata.
96    ///
97    /// The metadata is recorded with caller state; ordinary captures and their
98    /// serialization format are unchanged. Use `restore_with_dtypes` to restore
99    /// mixed storage dtypes independently from the recorder's value precision.
100    pub fn capture_with_dtypes(
101        model: &M,
102        optimizer: &O,
103        scheduler: &S,
104        accumulator: &GradientsAccumulator<M>,
105        state: U,
106    ) -> Result<TrainingRecord<B, M, O, S, (ModuleDTypeRecord, U)>, RecorderError> {
107        let dtypes = ModuleDTypeRecord::capture(model)?;
108        TrainingRecord::<B, M, O, S, (ModuleDTypeRecord, U)>::capture(
109            model,
110            optimizer,
111            scheduler,
112            accumulator,
113            (dtypes, state),
114        )
115    }
116
117    /// Asynchronously capture pending gradients and per-parameter storage dtypes.
118    pub async fn capture_async_with_dtypes(
119        model: &M,
120        optimizer: &O,
121        scheduler: &S,
122        accumulator: &GradientsAccumulator<M>,
123        state: U,
124    ) -> Result<TrainingRecord<B, M, O, S, (ModuleDTypeRecord, U)>, RecorderError> {
125        let dtypes = ModuleDTypeRecord::capture(model)?;
126        TrainingRecord::<B, M, O, S, (ModuleDTypeRecord, U)>::capture_async(
127            model,
128            optimizer,
129            scheduler,
130            accumulator,
131            (dtypes, state),
132        )
133        .await
134    }
135
136    /// Save all components in one recorder payload.
137    pub fn save<R: Recorder<B>>(
138        self,
139        recorder: &R,
140        args: R::RecordArgs,
141    ) -> Result<R::RecordOutput, RecorderError> {
142        recorder.record(self, args)
143    }
144
145    /// Read a combined record using the selected recorder and device.
146    pub fn load<R: Recorder<B>>(
147        recorder: &R,
148        args: R::LoadArgs,
149        device: &B::Device,
150    ) -> Result<Self, RecorderError> {
151        recorder.load(args, device)
152    }
153
154    /// Restore onto compatible model, optimizer and scheduler configurations.
155    ///
156    /// No scheduler step, optimizer step or gradient reset is performed. Apply
157    /// the returned caller state before consuming the next training batch.
158    pub fn restore(
159        self,
160        model: M,
161        optimizer: O,
162        scheduler: S,
163        device: &B::Device,
164    ) -> Result<RestoredTraining<M, O, S, U>, RecorderError> {
165        let mut accumulator = GradientsAccumulator::new();
166        accumulator.load_record::<B>(self.gradients, device)?;
167        let model = model.load_record(self.model).fork(device);
168        let optimizer = optimizer.load_record(self.optimizer);
169        let scheduler = scheduler.load_record::<B>(self.scheduler);
170        Ok(RestoredTraining {
171            model,
172            optimizer,
173            scheduler,
174            accumulator,
175            state: self.state,
176        })
177    }
178}
179
180impl<B, M, O, S, U> TrainingRecord<B, M, O, S, (ModuleDTypeRecord, U)>
181where
182    B: AutodiffBackend,
183    M: AutodiffModule<B>,
184    O: Optimizer<M, B>,
185    S: LrScheduler,
186    U: Record<B>,
187{
188    /// Restore captured floating storage dtypes, pending gradients and trainable state.
189    /// No update, scheduler step, gradient reset or requantization is performed.
190    pub fn restore_with_dtypes(
191        self,
192        model: M,
193        optimizer: O,
194        scheduler: S,
195        device: &B::Device,
196    ) -> Result<RestoredTraining<M, O, S, U>, RecorderError> {
197        let restored = self.restore(model, optimizer, scheduler, device)?;
198        let (dtypes, state) = restored.state;
199        Ok(RestoredTraining {
200            model: dtypes.apply(restored.model)?,
201            optimizer: restored.optimizer,
202            scheduler: restored.scheduler,
203            accumulator: restored.accumulator,
204            state,
205        })
206    }
207}
208
209impl<B, M, O, S, U> Record<B> for TrainingRecord<B, M, O, S, U>
210where
211    B: AutodiffBackend,
212    M: AutodiffModule<B>,
213    O: Optimizer<M, B>,
214    S: LrScheduler,
215    U: Record<B>,
216{
217    type Item<P: PrecisionSettings> = (
218        <M::Record as Record<B>>::Item<P>,
219        <O::Record as Record<B>>::Item<P>,
220        <S::Record<B> as Record<B>>::Item<P>,
221        <GradientsParamsRecord as Record<B>>::Item<P>,
222        U::Item<P>,
223    );
224
225    fn into_item<P: PrecisionSettings>(self) -> Self::Item<P> {
226        (
227            self.model.into_item::<P>(),
228            self.optimizer.into_item::<P>(),
229            self.scheduler.into_item::<P>(),
230            <GradientsParamsRecord as Record<B>>::into_item::<P>(self.gradients),
231            self.state.into_item::<P>(),
232        )
233    }
234
235    fn from_item<P: PrecisionSettings>(item: Self::Item<P>, device: &B::Device) -> Self {
236        Self {
237            model: <M::Record as Record<B>>::from_item::<P>(item.0, device),
238            optimizer: <O::Record as Record<B>>::from_item::<P>(item.1, device),
239            scheduler: <S::Record<B> as Record<B>>::from_item::<P>(item.2, device),
240            gradients: <GradientsParamsRecord as Record<B>>::from_item::<P>(item.3, device),
241            state: U::from_item::<P>(item.4, device),
242            marker: PhantomData,
243        }
244    }
245}