Skip to main content

WeightedGradientsAccumulator

Struct WeightedGradientsAccumulator 

Source
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>

Source

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.

Source

pub fn state(&self) -> &WeightedAccumulationState

Current actual counters; inspecting them does not clear gradients.

Source

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.

Source

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.

Source

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.

Source

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.

Source

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.

Source

pub fn try_to_record<B: AutodiffBackend>( &self, ) -> Result<WeightedGradientsRecord, RecorderError>
where M: AutodiffModule<B>,

Snapshot actual gradient values/options/counters without resetting.

Source

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.

Source

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.

Auto Trait Implementations§

Blanket Implementations§

Source§

impl<T> Any for T
where T: 'static + ?Sized,

Source§

fn type_id(&self) -> TypeId

Gets the TypeId of self. Read more
Source§

impl<T> Borrow<T> for T
where T: ?Sized,

Source§

fn borrow(&self) -> &T

Immutably borrows from an owned value. Read more
Source§

impl<T> BorrowMut<T> for T
where T: ?Sized,

Source§

fn borrow_mut(&mut self) -> &mut T

Mutably borrows from an owned value. Read more
Source§

impl<ST, DT> CastableFrom<ST, Initialized, Initialized> for DT
where ST: ?Sized, DT: ?Sized,

Source§

impl<ST, DT> CastableFrom<ST, Uninit, Uninit> for DT
where ST: ?Sized, DT: ?Sized,

Source§

impl<T> From<T> for T

Source§

fn from(t: T) -> T

Returns the argument unchanged.

Source§

impl<T, U> Into<U> for T
where U: From<T>,

Source§

fn into(self) -> U

Calls U::from(self).

That is, this conversion is whatever the implementation of From<T> for U chooses to do.

Source§

impl<T> Read<Exclusive, BecauseExclusive> for T
where T: ?Sized,

Source§

impl<T, U> TryFrom<U> for T
where U: Into<T>,

Source§

type Error = !

The type returned in the event of a conversion error.
Source§

fn try_from(value: U) -> Result<T, !>

Performs the conversion.
Source§

impl<T, U> TryInto<U> for T
where U: TryFrom<T>,

Source§

type Error = <U as TryFrom<T>>::Error

The type returned in the event of a conversion error.
Source§

fn try_into(self) -> Result<U, <U as TryFrom<T>>::Error>

Performs the conversion.