1use alloc::string::ToString;
2use core::fmt;
3use ruda_model::{
4 module::AutodiffModule,
5 record::{PrecisionSettings,Record,RecorderError},
6 tensor::{DType,FloatDType,TensorMetadata,backend::{AutodiffBackend,Backend}},
7};
8use serde::{Deserialize,Serialize};
9use super::{GradientsAccumulator,GradientsParams,GradientsParamsRecord,GradientTransformError};
10use super::gradient_transform::{representable,validate_work_dtype};
11
12#[derive(Clone,Debug,PartialEq,Serialize,Deserialize)]
14pub struct WeightedAccumulationState {
15 #[serde(serialize_with="serialize_dtype",deserialize_with="deserialize_dtype")]
17 pub dtype: FloatDType,
18 pub loss_scale: f64,
20 pub total_weight: f64,
22 pub microbatches: u64,
24}
25
26fn serialize_dtype<S: serde::Serializer>(dtype: &FloatDType,serializer: S) -> Result<S::Ok,S::Error> {
27 DType::from(*dtype).serialize(serializer)
28}
29
30fn deserialize_dtype<'de,D: serde::Deserializer<'de>>(deserializer: D) -> Result<FloatDType,D::Error> {
31 match DType::deserialize(deserializer)? {
32 DType::F32 => Ok(FloatDType::F32),
33 DType::F64 => Ok(FloatDType::F64),
34 _ => Err(serde::de::Error::custom("weighted accumulation dtype must be F32 or F64")),
35 }
36}
37
38impl<B: Backend> Record<B> for WeightedAccumulationState {
39 type Item<S: PrecisionSettings> = Self;
40 fn into_item<S: PrecisionSettings>(self) -> Self { self }
41 fn from_item<S: PrecisionSettings>(item: Self,_device: &B::Device) -> Self { item }
42}
43
44#[derive(Clone,Debug)]
46pub struct WeightedGradientsRecord {
47 pub gradients: GradientsParamsRecord,
49 pub state: WeightedAccumulationState,
51}
52
53impl<B: Backend> Record<B> for WeightedGradientsRecord {
54 type Item<S: PrecisionSettings> = (GradientsParamsRecord,WeightedAccumulationState);
55 fn into_item<S: PrecisionSettings>(self) -> Self::Item<S> {
56 (<GradientsParamsRecord as Record<B>>::into_item::<S>(self.gradients),self.state)
57 }
58 fn from_item<S: PrecisionSettings>(item: Self::Item<S>,device: &B::Device) -> Self {
59 Self { gradients:<GradientsParamsRecord as Record<B>>::from_item::<S>(item.0,device),state:item.1 }
60 }
61}
62
63#[derive(Clone,Debug,PartialEq,Eq)]
65pub enum WeightedAccumulationError {
66 Gradient(GradientTransformError),
68 InvalidWeight,
70 CounterOverflow,
72 EmptyWeight,
74 InvalidState,
76}
77
78impl From<GradientTransformError> for WeightedAccumulationError {
79 fn from(error: GradientTransformError) -> Self { Self::Gradient(error) }
80}
81
82impl fmt::Display for WeightedAccumulationError {
83 fn fmt(&self,f: &mut fmt::Formatter<'_>) -> fmt::Result {
84 match self {
85 Self::Gradient(error) => fmt::Display::fmt(error,f),
86 Self::InvalidWeight => f.write_str("effective accumulation weight must be finite, nonnegative and representable"),
87 Self::CounterOverflow => f.write_str("accumulation microbatch counter overflow"),
88 Self::EmptyWeight => f.write_str("cannot normalize an accumulation window with zero effective weight"),
89 Self::InvalidState => f.write_str("inconsistent weighted accumulation checkpoint state"),
90 }
91 }
92}
93
94#[cfg(feature = "std")]
95impl std::error::Error for WeightedAccumulationError {}
96
97pub struct AccumulatedGradients {
99 pub gradients: GradientsParams,
101 pub state: WeightedAccumulationState,
103}
104
105pub struct WeightedGradientsAccumulator<M> {
112 accumulator: GradientsAccumulator<M>,
113 state: WeightedAccumulationState,
114}
115
116impl<M> WeightedGradientsAccumulator<M> {
117 pub fn new(dtype: FloatDType,loss_scale: f64) -> Result<Self,WeightedAccumulationError> {
119 let state = WeightedAccumulationState { dtype,loss_scale,total_weight:0.,microbatches:0 };
120 validate_state(&state)?;
121 Ok(Self { accumulator:GradientsAccumulator::new(),state })
122 }
123
124 pub fn state(&self) -> &WeightedAccumulationState { &self.state }
126
127 pub(crate) fn inner(&self) -> &GradientsAccumulator<M> { &self.accumulator }
128
129 pub fn from_accumulator<B: AutodiffBackend>(
132 module: &M,accumulator: GradientsAccumulator<M>,state: WeightedAccumulationState,
133 ) -> Result<Self,WeightedAccumulationError> where M: AutodiffModule<B> {
134 validate_state(&state)?;
135 if state.total_weight == 0. && !accumulator.pending().is_empty() {
136 return Err(WeightedAccumulationError::InvalidState);
137 }
138 validate_pending_dtype::<B>(accumulator.pending(),state.dtype)?;
139 let values = accumulator.pending().cast_for::<B,M>(module,state.dtype)?;
140 let mut restored = GradientsAccumulator::new();
141 restored.accumulate_with_dtype::<B>(module,values,state.dtype);
142 Ok(Self {accumulator:restored,state})
143 }
144
145 pub fn accumulate_sum<B: AutodiffBackend>(
149 &mut self,module: &M,gradients: &GradientsParams,weight: f64,
150 ) -> Result<(),WeightedAccumulationError> where M: AutodiffModule<B> {
151 self.accumulate::<B>(module,gradients,weight,false)
152 }
153
154 pub fn accumulate_mean<B: AutodiffBackend>(
157 &mut self,module: &M,gradients: &GradientsParams,weight: f64,
158 ) -> Result<(),WeightedAccumulationError> where M: AutodiffModule<B> {
159 self.accumulate::<B>(module,gradients,weight,true)
160 }
161
162 fn accumulate<B: AutodiffBackend>(
163 &mut self,module: &M,gradients: &GradientsParams,weight: f64,mean: bool,
164 ) -> Result<(),WeightedAccumulationError> where M: AutodiffModule<B> {
165 if weight < 0. || !representable(weight,self.state.dtype) ||
166 (weight > 0. && self.state.dtype == FloatDType::F32 && weight as f32 == 0.) {
167 return Err(WeightedAccumulationError::InvalidWeight);
168 }
169 let total = self.state.total_weight + weight;
170 if !representable(total,self.state.dtype) { return Err(WeightedAccumulationError::InvalidWeight); }
171 let count = self.state.microbatches.checked_add(1).ok_or(WeightedAccumulationError::CounterOverflow)?;
172 if weight > 0. {
175 let multiplier = if mean { weight } else { 1. };
176 let values = gradients.scaled_for::<B,M>(module,multiplier,self.state.dtype)?;
177 self.accumulator.accumulate_with_dtype::<B>(module,values,self.state.dtype);
178 } else {
179 gradients.validate_for::<B,M>(module)?;
180 }
181 self.state.total_weight = total;
182 self.state.microbatches = count;
183 Ok(())
184 }
185
186 pub fn finish_sums(&mut self) -> AccumulatedGradients {
190 let state = self.state.clone();
191 let gradients = self.accumulator.grads();
192 self.state.total_weight = 0.;
193 self.state.microbatches = 0;
194 AccumulatedGradients {gradients,state}
195 }
196
197 pub fn finish_mean<B: AutodiffBackend>(
201 &mut self,module: &M,
202 ) -> Result<AccumulatedGradients,WeightedAccumulationError> where M: AutodiffModule<B> {
203 if self.state.total_weight <= 0. { return Err(WeightedAccumulationError::EmptyWeight); }
204 let gradients = self.accumulator.pending().unscaled_for::<B,M>(
205 module,self.state.loss_scale,self.state.dtype)?.unscaled_for::<B,M>(
206 module,self.state.total_weight,self.state.dtype)?;
207 let state = self.state.clone();
208 self.accumulator.grads();
209 self.state.total_weight = 0.;
210 self.state.microbatches = 0;
211 Ok(AccumulatedGradients {gradients,state})
212 }
213
214 pub fn try_to_record<B: AutodiffBackend>(&self) -> Result<WeightedGradientsRecord,RecorderError>
216 where M: AutodiffModule<B> {
217 Ok(WeightedGradientsRecord {gradients:self.accumulator.try_to_record::<B>()?,state:self.state.clone()})
218 }
219
220 pub async fn to_record_async<B: AutodiffBackend>(&self) -> Result<WeightedGradientsRecord,RecorderError>
222 where M: AutodiffModule<B> {
223 Ok(WeightedGradientsRecord {gradients:self.accumulator.to_record_async::<B>().await?,state:self.state.clone()})
224 }
225
226 pub fn load_record<B: AutodiffBackend>(
229 &mut self,module: &M,record: WeightedGradientsRecord,device: &B::Device,
230 ) -> Result<(),RecorderError> where M: AutodiffModule<B> {
231 validate_state(&record.state).map_err(|error|RecorderError::Unknown(error.to_string()))?;
232 let gradients = GradientsParams::from_record::<B::InnerBackend>(record.gradients,device)?;
233 if record.state.total_weight == 0. && !gradients.is_empty() {
234 return Err(RecorderError::Unknown(WeightedAccumulationError::InvalidState.to_string()));
235 }
236 validate_pending_dtype::<B>(&gradients,record.state.dtype)
237 .map_err(|error|RecorderError::Unknown(error.to_string()))?;
238 gradients.validate_for::<B,M>(module).map_err(|error|RecorderError::Unknown(error.to_string()))?;
239 let mut restored = GradientsAccumulator::new();
240 let values = gradients.cast_for::<B,M>(module,record.state.dtype)
241 .map_err(|error|RecorderError::Unknown(error.to_string()))?;
242 restored.accumulate_with_dtype::<B>(module,values,record.state.dtype);
243 self.accumulator = restored;
244 self.state = record.state;
245 Ok(())
246 }
247}
248
249fn validate_pending_dtype<B: AutodiffBackend>(gradients: &GradientsParams,dtype: FloatDType)
250 -> Result<(),WeightedAccumulationError> {
251 for id in gradients.container.ids() {
252 let primitive = gradients.container.get::<B::InnerBackend>(id)
253 .ok_or(WeightedAccumulationError::InvalidState)?;
254 if primitive.dtype() != DType::from(dtype) { return Err(WeightedAccumulationError::InvalidState); }
255 }
256 Ok(())
257}
258
259fn validate_state(state: &WeightedAccumulationState) -> Result<(),WeightedAccumulationError> {
260 validate_work_dtype(state.dtype)?;
261 if !representable(state.loss_scale,state.dtype) || state.loss_scale <= 0. ||
262 (state.dtype == FloatDType::F32 && state.loss_scale as f32 == 0.) {
263 return Err(GradientTransformError::InvalidScalar.into());
264 }
265 if state.total_weight < 0. || !representable(state.total_weight,state.dtype) ||
266 (state.total_weight > 0. && state.dtype == FloatDType::F32 && state.total_weight as f32 == 0.) ||
267 (state.microbatches == 0 && state.total_weight != 0.) {
268 return Err(WeightedAccumulationError::InvalidState);
269 }
270 Ok(())
271}