Skip to main content

ruda_optim/
training.rs

1use core::marker::PhantomData;
2use ruda_model::{
3    module::AutodiffModule,
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    /// Save all components in one recorder payload.
96    pub fn save<R: Recorder<B>>(
97        self,
98        recorder: &R,
99        args: R::RecordArgs,
100    ) -> Result<R::RecordOutput, RecorderError> {
101        recorder.record(self, args)
102    }
103
104    /// Read a combined record using the selected recorder and device.
105    pub fn load<R: Recorder<B>>(
106        recorder: &R,
107        args: R::LoadArgs,
108        device: &B::Device,
109    ) -> Result<Self, RecorderError> {
110        recorder.load(args, device)
111    }
112
113    /// Restore onto compatible model, optimizer and scheduler configurations.
114    ///
115    /// No scheduler step, optimizer step or gradient reset is performed. Apply
116    /// the returned caller state before consuming the next training batch.
117    pub fn restore(
118        self,
119        model: M,
120        optimizer: O,
121        scheduler: S,
122        device: &B::Device,
123    ) -> Result<RestoredTraining<M, O, S, U>, RecorderError> {
124        let mut accumulator = GradientsAccumulator::new();
125        accumulator.load_record::<B>(self.gradients, device)?;
126        let model = model.load_record(self.model).fork(device);
127        let optimizer = optimizer.load_record(self.optimizer);
128        let scheduler = scheduler.load_record::<B>(self.scheduler);
129        Ok(RestoredTraining {
130            model,
131            optimizer,
132            scheduler,
133            accumulator,
134            state: self.state,
135        })
136    }
137}
138
139impl<B, M, O, S, U> Record<B> for TrainingRecord<B, M, O, S, U>
140where
141    B: AutodiffBackend,
142    M: AutodiffModule<B>,
143    O: Optimizer<M, B>,
144    S: LrScheduler,
145    U: Record<B>,
146{
147    type Item<P: PrecisionSettings> = (
148        <M::Record as Record<B>>::Item<P>,
149        <O::Record as Record<B>>::Item<P>,
150        <S::Record<B> as Record<B>>::Item<P>,
151        <GradientsParamsRecord as Record<B>>::Item<P>,
152        U::Item<P>,
153    );
154
155    fn into_item<P: PrecisionSettings>(self) -> Self::Item<P> {
156        (
157            self.model.into_item::<P>(),
158            self.optimizer.into_item::<P>(),
159            self.scheduler.into_item::<P>(),
160            <GradientsParamsRecord as Record<B>>::into_item::<P>(self.gradients),
161            self.state.into_item::<P>(),
162        )
163    }
164
165    fn from_item<P: PrecisionSettings>(item: Self::Item<P>, device: &B::Device) -> Self {
166        Self {
167            model: <M::Record as Record<B>>::from_item::<P>(item.0, device),
168            optimizer: <O::Record as Record<B>>::from_item::<P>(item.1, device),
169            scheduler: <S::Record<B> as Record<B>>::from_item::<P>(item.2, device),
170            gradients: <GradientsParamsRecord as Record<B>>::from_item::<P>(item.3, device),
171            state: U::from_item::<P>(item.4, device),
172            marker: PhantomData,
173        }
174    }
175}