Skip to main content

ruda_optim/optim/
weighted_accum.rs

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/// Explicit normalization and continuation state of an accumulation window.
13#[derive(Clone,Debug,PartialEq,Serialize,Deserialize)]
14pub struct WeightedAccumulationState {
15    /// F32/F64 arithmetic used before adding each microbatch's gradients.
16    #[serde(serialize_with="serialize_dtype",deserialize_with="deserialize_dtype")]
17    pub dtype: FloatDType,
18    /// Loss multiplier actually applied by the caller throughout this window.
19    pub loss_scale: f64,
20    /// Sum of caller-supplied effective token/sample weights, not batch count.
21    pub total_weight: f64,
22    /// Number of accepted microbatches, including explicit zero-weight batches.
23    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/// Pending gradients and their actual normalization state, saved together.
45#[derive(Clone,Debug)]
46pub struct WeightedGradientsRecord {
47    /// Values keyed by the original model IDs; full precision settings recommended.
48    pub gradients: GradientsParamsRecord,
49    /// Work dtype, loss scale, effective weight and pending microbatch count.
50    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/// Explicit accumulation argument, geometry or checkpoint error.
64#[derive(Clone,Debug,PartialEq,Eq)]
65pub enum WeightedAccumulationError {
66    /// Invalid work dtype, scalar or gradient metadata.
67    Gradient(GradientTransformError),
68    /// Effective weights must be nonnegative, finite and representable.
69    InvalidWeight,
70    /// The actual microbatch counter cannot accept another increment.
71    CounterOverflow,
72    /// A mean requires positive total weight; the pending window remains intact.
73    EmptyWeight,
74    /// Stored counters/options describe an inconsistent accumulation window.
75    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
97/// Result of an explicitly completed accumulation window.
98pub struct AccumulatedGradients {
99    /// Present gradients only; globally/local unused parameters are not invented.
100    pub gradients: GradientsParams,
101    /// Actual counts and loss scale of the window that produced these gradients.
102    pub state: WeightedAccumulationState,
103}
104
105/// Backend-independent unequal-microbatch accumulation with resumable weights.
106///
107/// The caller supplies gradients of either local loss sums or local loss means,
108/// plus the effective weight used by that loss. Loss scaling is fixed at setup;
109/// it is never inferred from values. No optimizer, scheduler, data iterator,
110/// dynamic scaler, clipping or distributed reduction is advanced implicitly.
111pub struct WeightedGradientsAccumulator<M> {
112    accumulator: GradientsAccumulator<M>,
113    state: WeightedAccumulationState,
114}
115
116impl<M> WeightedGradientsAccumulator<M> {
117    /// Start an empty window. Use one loss scale for every forward/backward in it.
118    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    /// Current actual counters; inspecting them does not clear gradients.
125    pub fn state(&self) -> &WeightedAccumulationState { &self.state }
126
127    pub(crate) fn inner(&self) -> &GradientsAccumulator<M> { &self.accumulator }
128
129    /// Reattach saved normalization state to already-restored pending gradients.
130    /// The supplied module must have the same IDs, dimensions and device.
131    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    /// Accumulate gradients of a caller's weighted loss SUM.
146    /// `weight` is its effective token/sample count, not the loss multiplier.
147    /// An explicit zero-weight batch contributes no gradients but counts as issued.
148    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    /// Accumulate a MEAN loss by multiplying gradients by its actual weight
155    /// after work-dtype conversion. Unequal batch/token counts are not averaged equally.
156    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        // Validate every supplied ID before any scaling/addition. A rejected
173        // argument leaves pending gradients and normalization counters untouched.
174        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    /// Return the actual scaled loss sums and their local weight, then reset.
187    /// Useful before DDP's existing sum/weight reduction; no communication occurs.
188    /// Divide the reduced gradients by `state.loss_scale` exactly once yourself.
189    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    /// Normalize by loss scale and total effective weight, then clear the window.
198    /// Division is sequential, avoiding an overflowing scale*weight product.
199    /// Zero total weight returns an error without clearing anything.
200    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    /// Snapshot actual gradient values/options/counters without resetting.
215    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    /// Snapshot with asynchronous device readback; source position remains caller state.
221    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    /// Restore counters and pending values together for the checkpoint's model IDs.
227    /// Invalid state/membership leaves this accumulator unchanged. No batch replay.
228    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}