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
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 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 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 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}