pub struct WeightedGradientsAccumulator<M> { /* private fields */ }Expand description
Backend-independent unequal-microbatch accumulation with resumable weights.
The caller supplies gradients of either local loss sums or local loss means, plus the effective weight used by that loss. Loss scaling is fixed at setup; it is never inferred from values. No optimizer, scheduler, data iterator, dynamic scaler, clipping or distributed reduction is advanced implicitly.
Implementations§
Source§impl<M> WeightedGradientsAccumulator<M>
impl<M> WeightedGradientsAccumulator<M>
Sourcepub fn new(
dtype: FloatDType,
loss_scale: f64,
) -> Result<Self, WeightedAccumulationError>
pub fn new( dtype: FloatDType, loss_scale: f64, ) -> Result<Self, WeightedAccumulationError>
Start an empty window. Use one loss scale for every forward/backward in it.
Sourcepub fn state(&self) -> &WeightedAccumulationState
pub fn state(&self) -> &WeightedAccumulationState
Current actual counters; inspecting them does not clear gradients.
Sourcepub fn from_accumulator<B: AutodiffBackend>(
module: &M,
accumulator: GradientsAccumulator<M>,
state: WeightedAccumulationState,
) -> Result<Self, WeightedAccumulationError>where
M: AutodiffModule<B>,
pub fn from_accumulator<B: AutodiffBackend>(
module: &M,
accumulator: GradientsAccumulator<M>,
state: WeightedAccumulationState,
) -> Result<Self, WeightedAccumulationError>where
M: AutodiffModule<B>,
Reattach saved normalization state to already-restored pending gradients. The supplied module must have the same IDs, dimensions and device.
Sourcepub fn accumulate_sum<B: AutodiffBackend>(
&mut self,
module: &M,
gradients: &GradientsParams,
weight: f64,
) -> Result<(), WeightedAccumulationError>where
M: AutodiffModule<B>,
pub fn accumulate_sum<B: AutodiffBackend>(
&mut self,
module: &M,
gradients: &GradientsParams,
weight: f64,
) -> Result<(), WeightedAccumulationError>where
M: AutodiffModule<B>,
Accumulate gradients of a caller’s weighted loss SUM.
weight is its effective token/sample count, not the loss multiplier.
An explicit zero-weight batch contributes no gradients but counts as issued.
Sourcepub fn accumulate_mean<B: AutodiffBackend>(
&mut self,
module: &M,
gradients: &GradientsParams,
weight: f64,
) -> Result<(), WeightedAccumulationError>where
M: AutodiffModule<B>,
pub fn accumulate_mean<B: AutodiffBackend>(
&mut self,
module: &M,
gradients: &GradientsParams,
weight: f64,
) -> Result<(), WeightedAccumulationError>where
M: AutodiffModule<B>,
Accumulate a MEAN loss by multiplying gradients by its actual weight after work-dtype conversion. Unequal batch/token counts are not averaged equally.
Sourcepub fn finish_sums(&mut self) -> AccumulatedGradients
pub fn finish_sums(&mut self) -> AccumulatedGradients
Return the actual scaled loss sums and their local weight, then reset.
Useful before DDP’s existing sum/weight reduction; no communication occurs.
Divide the reduced gradients by state.loss_scale exactly once yourself.
Sourcepub fn finish_mean<B: AutodiffBackend>(
&mut self,
module: &M,
) -> Result<AccumulatedGradients, WeightedAccumulationError>where
M: AutodiffModule<B>,
pub fn finish_mean<B: AutodiffBackend>(
&mut self,
module: &M,
) -> Result<AccumulatedGradients, WeightedAccumulationError>where
M: AutodiffModule<B>,
Normalize by loss scale and total effective weight, then clear the window. Division is sequential, avoiding an overflowing scale*weight product. Zero total weight returns an error without clearing anything.
Sourcepub fn try_to_record<B: AutodiffBackend>(
&self,
) -> Result<WeightedGradientsRecord, RecorderError>where
M: AutodiffModule<B>,
pub fn try_to_record<B: AutodiffBackend>(
&self,
) -> Result<WeightedGradientsRecord, RecorderError>where
M: AutodiffModule<B>,
Snapshot actual gradient values/options/counters without resetting.
Sourcepub async fn to_record_async<B: AutodiffBackend>(
&self,
) -> Result<WeightedGradientsRecord, RecorderError>where
M: AutodiffModule<B>,
pub async fn to_record_async<B: AutodiffBackend>(
&self,
) -> Result<WeightedGradientsRecord, RecorderError>where
M: AutodiffModule<B>,
Snapshot with asynchronous device readback; source position remains caller state.
Sourcepub fn load_record<B: AutodiffBackend>(
&mut self,
module: &M,
record: WeightedGradientsRecord,
device: &B::Device,
) -> Result<(), RecorderError>where
M: AutodiffModule<B>,
pub fn load_record<B: AutodiffBackend>(
&mut self,
module: &M,
record: WeightedGradientsRecord,
device: &B::Device,
) -> Result<(), RecorderError>where
M: AutodiffModule<B>,
Restore counters and pending values together for the checkpoint’s model IDs. Invalid state/membership leaves this accumulator unchanged. No batch replay.