1use alloc::{collections::{BTreeMap,BTreeSet},format,string::ToString,vec::Vec};
2use core::marker::PhantomData;
3use ruda_model::{
4 module::{AutodiffModule,ModuleVisitor,Param},
5 record::{PrecisionSettings,Record,Recorder,RecorderError},
6 tensor::{DType,Tensor,backend::{AutodiffBackend,Backend}},
7};
8use serde::{Deserialize,Serialize};
9use crate::{GradientsAccumulator,GradientsParams,GradientsParamsRecord,Optimizer,
10 WeightedAccumulationState,WeightedGradientsAccumulator,lr_scheduler::LrScheduler};
11use super::{RestoredTraining,RestoredWeightedTraining};
12
13#[derive(Clone,Debug,PartialEq,Eq,Serialize,Deserialize)]
15pub struct TrainableParameterContract {
16 entries: Vec<(u64,Vec<usize>,DType)>,
17}
18
19impl<B: Backend> Record<B> for TrainableParameterContract {
20 type Item<S: PrecisionSettings> = Self;
21 fn into_item<S: PrecisionSettings>(self) -> Self { self }
22 fn from_item<S: PrecisionSettings>(item: Self,_device: &B::Device) -> Self { item }
23}
24
25struct CaptureContract {
26 entries: BTreeMap<u64,(Vec<usize>,DType)>,
27 error: bool,
28}
29impl<B: AutodiffBackend> ModuleVisitor<B> for CaptureContract {
30 fn visit_float<const D: usize>(&mut self,param: &Param<Tensor<B,D>>) {
31 let tensor = param.val();
32 if !tensor.is_require_grad() { return; }
33 let metadata = (tensor.dims().to_vec(),tensor.dtype());
34 if let Some(previous) = self.entries.insert(param.id.val(),metadata.clone()) {
35 self.error |= previous != metadata;
36 }
37 }
38}
39
40impl TrainableParameterContract {
41 pub fn capture<B: AutodiffBackend,M: AutodiffModule<B>>(model: &M) -> Result<Self,RecorderError> {
43 let mut visitor = CaptureContract {entries:BTreeMap::new(),error:false};
44 model.visit(&mut visitor);
45 if visitor.error { return Err(invalid("tied trainable parameter geometry/dtype differs")); }
46 Ok(Self {entries:visitor.entries.into_iter().map(|(id,(shape,dtype))|(id,shape,dtype)).collect()})
47 }
48
49 pub fn validate_for<B: AutodiffBackend,M: AutodiffModule<B>>(&self,model: &M) -> Result<(),RecorderError> {
51 let actual = Self::capture::<B,M>(model)?;
52 if *self != actual { return Err(invalid("restored trainable parameter IDs, geometry or storage differs")); }
53 Ok(())
54 }
55
56 pub fn parameters(&self) -> usize { self.entries.len() }
58}
59
60struct PendingCheck<'a> {
61 gradients: &'a GradientsParams,
62 active_ids: BTreeSet<u64>,
63 frozen: bool,
64}
65impl<B: AutodiffBackend> ModuleVisitor<B> for PendingCheck<'_> {
66 fn visit_float<const D: usize>(&mut self,param: &Param<Tensor<B,D>>) {
67 if !param.val().is_require_grad() && !self.active_ids.contains(¶m.id.val())
68 && self.gradients.get::<B::InnerBackend,D>(param.id).is_some() { self.frozen = true; }
69 }
70}
71
72fn check_pending<B: AutodiffBackend,M: AutodiffModule<B>>(model: &M,accumulator: &GradientsAccumulator<M>)
73 -> Result<(),RecorderError> {
74 accumulator.pending().validate_for::<B,M>(model).map_err(|error|RecorderError::Unknown(error.to_string()))?;
75 let active_ids = TrainableParameterContract::capture::<B,M>(model)?.entries.into_iter().map(|(id,_,_)|id).collect();
76 let mut visitor = PendingCheck {gradients:accumulator.pending(),active_ids,frozen:false};
77 model.visit(&mut visitor);
78 if visitor.frozen { return Err(invalid("pending gradients include a frozen parameter")); }
79 Ok(())
80}
81
82fn invalid(reason: &str) -> RecorderError {
83 RecorderError::Unknown(format!("Invalid model-state training record: {reason}"))
84}
85
86pub struct ModelStateTrainingRecord<B,M,O,S,R,U>
93where B: AutodiffBackend,M: AutodiffModule<B>,O: Optimizer<M,B>,S: LrScheduler,R: Record<B>,U: Record<B> {
94 version: u32,
95 model_state: R,
96 contract: TrainableParameterContract,
97 optimizer: O::Record,
98 scheduler: S::Record<B>,
99 gradients: GradientsParamsRecord,
100 state: U,
101 marker: PhantomData<fn()->(B,M,O,S)>,
102}
103
104impl<B,M,O,S,R,U> ModelStateTrainingRecord<B,M,O,S,R,U>
105where B: AutodiffBackend,M: AutodiffModule<B>,O: Optimizer<M,B>,S: LrScheduler,R: Record<B>,U: Record<B> {
106 pub fn capture(model: &M,model_state: R,optimizer: &O,scheduler: &S,
108 accumulator: &GradientsAccumulator<M>,state: U) -> Result<Self,RecorderError> {
109 check_pending::<B,M>(model,accumulator)?;
110 let contract = TrainableParameterContract::capture::<B,M>(model)?;
111 let gradients = accumulator.try_to_record::<B>()?;
112 Ok(Self {version:1,model_state,contract,optimizer:optimizer.to_record(),scheduler:scheduler.to_record::<B>(),
113 gradients,state,marker:PhantomData})
114 }
115
116 pub async fn capture_async(model: &M,model_state: R,optimizer: &O,scheduler: &S,
118 accumulator: &GradientsAccumulator<M>,state: U) -> Result<Self,RecorderError> {
119 check_pending::<B,M>(model,accumulator)?;
120 let contract = TrainableParameterContract::capture::<B,M>(model)?;
121 let gradients = accumulator.to_record_async::<B>().await?;
122 Ok(Self {version:1,model_state,contract,optimizer:optimizer.to_record(),scheduler:scheduler.to_record::<B>(),
123 gradients,state,marker:PhantomData})
124 }
125
126 pub fn capture_weighted(model: &M,model_state: R,optimizer: &O,scheduler: &S,
128 accumulator: &WeightedGradientsAccumulator<M>,state: U)
129 -> Result<ModelStateTrainingRecord<B,M,O,S,R,(WeightedAccumulationState,U)>,RecorderError> {
130 ModelStateTrainingRecord::<B,M,O,S,R,(WeightedAccumulationState,U)>::capture(
131 model,model_state,optimizer,scheduler,accumulator.inner(),(accumulator.state().clone(),state))
132 }
133
134 pub async fn capture_weighted_async(model: &M,model_state: R,optimizer: &O,scheduler: &S,
136 accumulator: &WeightedGradientsAccumulator<M>,state: U)
137 -> Result<ModelStateTrainingRecord<B,M,O,S,R,(WeightedAccumulationState,U)>,RecorderError> {
138 ModelStateTrainingRecord::<B,M,O,S,R,(WeightedAccumulationState,U)>::capture_async(
139 model,model_state,optimizer,scheduler,accumulator.inner(),(accumulator.state().clone(),state)).await
140 }
141
142 pub fn save<C: Recorder<B>>(self,recorder: &C,args: C::RecordArgs) -> Result<C::RecordOutput,RecorderError> {
144 recorder.record(self,args)
145 }
146
147 pub fn load<C: Recorder<B>>(recorder: &C,args: C::LoadArgs,device: &B::Device) -> Result<Self,RecorderError> {
149 recorder.load(args,device)
150 }
151
152 pub fn restore<F>(self,model: M,optimizer: O,scheduler: S,device: &B::Device,restore_model: F)
157 -> Result<RestoredTraining<M,O,S,U>,RecorderError>
158 where F: FnOnce(R,M)->Result<M,RecorderError> {
159 if self.version != 1 { return Err(invalid("unsupported format version")); }
160 let model = restore_model(self.model_state,model)?;
161 if model.devices().iter().any(|actual|actual != device) { return Err(invalid("restored model must already use the requested device")); }
162 self.contract.validate_for::<B,M>(&model)?;
163 let mut accumulator = GradientsAccumulator::new();
164 accumulator.load_record::<B>(self.gradients,device)?;
165 check_pending::<B,M>(&model,&accumulator)?;
166 Ok(RestoredTraining {model,optimizer:optimizer.load_record(self.optimizer),scheduler:scheduler.load_record::<B>(self.scheduler),
167 accumulator,state:self.state})
168 }
169}
170
171impl<B,M,O,S,R,U> ModelStateTrainingRecord<B,M,O,S,R,(WeightedAccumulationState,U)>
172where B: AutodiffBackend,M: AutodiffModule<B>,O: Optimizer<M,B>,S: LrScheduler,R: Record<B>,U: Record<B> {
173 pub fn restore_weighted<F>(self,model: M,optimizer: O,scheduler: S,device: &B::Device,restore_model: F)
175 -> Result<RestoredWeightedTraining<M,O,S,U>,RecorderError>
176 where F: FnOnce(R,M)->Result<M,RecorderError> {
177 super::restore_weighted::<B,M,O,S,U>(self.restore(model,optimizer,scheduler,device,restore_model)?)
178 }
179}
180
181impl<B,M,O,S,R,U> Record<B> for ModelStateTrainingRecord<B,M,O,S,R,U>
182where B: AutodiffBackend,M: AutodiffModule<B>,O: Optimizer<M,B>,S: LrScheduler,R: Record<B>,U: Record<B> {
183 type Item<P: PrecisionSettings> = (u32,R::Item<P>,<TrainableParameterContract as Record<B>>::Item<P>,
184 <O::Record as Record<B>>::Item<P>,<S::Record<B> as Record<B>>::Item<P>,
185 <GradientsParamsRecord as Record<B>>::Item<P>,U::Item<P>);
186 fn into_item<P: PrecisionSettings>(self) -> Self::Item<P> {
187 (self.version,self.model_state.into_item::<P>(),<TrainableParameterContract as Record<B>>::into_item::<P>(self.contract),
188 self.optimizer.into_item::<P>(),self.scheduler.into_item::<P>(),
189 <GradientsParamsRecord as Record<B>>::into_item::<P>(self.gradients),self.state.into_item::<P>())
190 }
191 fn from_item<P: PrecisionSettings>(item: Self::Item<P>,device: &B::Device) -> Self {
192 Self {version:item.0,model_state:R::from_item::<P>(item.1,device),
193 contract:<TrainableParameterContract as Record<B>>::from_item::<P>(item.2,device),
194 optimizer:<O::Record as Record<B>>::from_item::<P>(item.3,device),
195 scheduler:<S::Record<B> as Record<B>>::from_item::<P>(item.4,device),
196 gradients:<GradientsParamsRecord as Record<B>>::from_item::<P>(item.5,device),state:U::from_item::<P>(item.6,device),marker:PhantomData}
197 }
198}