Skip to main content

ruda_optim/training/
model_state.rs

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/// Exact trainable IDs, logical shapes and storage, without frozen weight values.
14#[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    /// Capture actual trainable metadata once per tied ID; no tensor values are read.
42    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    /// Match original trainable IDs/dtypes after the caller's model-state restoration.
50    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    /// Actual unique trainable parameter count, not the total base-model size.
57    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(&param.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
86/// Caller-selected model state plus actual optimizer/scheduler/pending gradients.
87///
88/// `R` can be RUDA's A/B-only native adapter record. It must restore every
89/// trainable parameter, including IDs/dtypes, and identify any omitted frozen
90/// state. Capture `R` and these components at the same boundary, without concurrent
91/// updates. No full model record or hidden frozen tensor copy is created here.
92pub 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    /// Capture provided actual model state, optimizer, scheduler and pending values.
107    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    /// Same capture with asynchronous readback of actual pending gradients.
117    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    /// Capture actual weighted-window counts/scale alongside caller source/RNG state.
127    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    /// Asynchronous weighted capture, with no implicit window reset or update.
135    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    /// Save provided model state and live training components in one recorder payload.
143    pub fn save<C: Recorder<B>>(self,recorder: &C,args: C::RecordArgs) -> Result<C::RecordOutput,RecorderError> {
144        recorder.record(self,args)
145    }
146
147    /// Read the combined state on the requested device; recreate omitted base state separately.
148    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    /// Restore caller-selected model state first, then validate actual trainable identities.
153    /// `restore_model` may call a native adapter record's restore_into with an independently
154    /// supplied frozen base identity. The prepared model must already use the requested
155    /// device; the omitted base is never copied or moved implicitly.
156    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    /// Restore actual A/B/model state and the weighted accumulation window together.
174    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}