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
13pub 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
35pub struct RestoredTraining<M, O, S, U> {
37 pub model: M,
39 pub optimizer: O,
41 pub scheduler: S,
43 pub accumulator: GradientsAccumulator<M>,
45 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 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 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 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 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 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 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 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 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}