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
18pub 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
40pub struct RestoredTraining<M, O, S, U> {
42 pub model: M,
44 pub optimizer: O,
46 pub scheduler: S,
48 pub accumulator: GradientsAccumulator<M>,
50 pub state: U,
52}
53
54pub struct RestoredWeightedTraining<M,O,S,U> {
56 pub model: M,
58 pub optimizer: O,
60 pub scheduler: S,
62 pub accumulator: WeightedGradientsAccumulator<M>,
64 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 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 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 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 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 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 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 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 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 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 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 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 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 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}