pub struct GradientsAccumulator<M> { /* private fields */ }Expand description
Accumulate gradients into a single Gradients object.
Implementations§
Source§impl<M> GradientsAccumulator<M>
impl<M> GradientsAccumulator<M>
Source§impl<M> GradientsAccumulator<M>
impl<M> GradientsAccumulator<M>
Sourcepub fn try_to_record<B: AutodiffBackend>(
&self,
) -> Result<GradientsParamsRecord, RecorderError>where
M: AutodiffModule<B>,
pub fn try_to_record<B: AutodiffBackend>(
&self,
) -> Result<GradientsParamsRecord, RecorderError>where
M: AutodiffModule<B>,
Snapshot pending gradients without resetting the accumulation window.
Sourcepub async fn to_record_async<B: AutodiffBackend>(
&self,
) -> Result<GradientsParamsRecord, RecorderError>where
M: AutodiffModule<B>,
pub async fn to_record_async<B: AutodiffBackend>(
&self,
) -> Result<GradientsParamsRecord, RecorderError>where
M: AutodiffModule<B>,
Asynchronously snapshot pending gradients without resetting the accumulator.
Sourcepub fn load_record<B: AutodiffBackend>(
&mut self,
record: GradientsParamsRecord,
device: &B::Device,
) -> Result<(), RecorderError>where
M: AutodiffModule<B>,
pub fn load_record<B: AutodiffBackend>(
&mut self,
record: GradientsParamsRecord,
device: &B::Device,
) -> Result<(), RecorderError>where
M: AutodiffModule<B>,
Replace pending gradients with checkpoint state on the given device.
The model IDs and the caller’s accumulation count must be restored from the same checkpoint. A rejected record leaves the accumulator unchanged.
Sourcepub fn accumulate<B: AutodiffBackend>(
&mut self,
module: &M,
grads: GradientsParams,
)where
M: AutodiffModule<B>,
pub fn accumulate<B: AutodiffBackend>(
&mut self,
module: &M,
grads: GradientsParams,
)where
M: AutodiffModule<B>,
Accumulate the given gradients for each parameter in the given module.
Sourcepub fn grads(&mut self) -> GradientsParams
pub fn grads(&mut self) -> GradientsParams
Return the accumulated gradients and reset the accumulator state.
Trait Implementations§
Auto Trait Implementations§
impl<M> !RefUnwindSafe for GradientsAccumulator<M>
impl<M> !Sync for GradientsAccumulator<M>
impl<M> !UnwindSafe for GradientsAccumulator<M>
impl<M> Freeze for GradientsAccumulator<M>where
PhantomData<M>: Freeze,
impl<M> Send for GradientsAccumulator<M>where
PhantomData<M>: Send,
impl<M> Unpin for GradientsAccumulator<M>where
PhantomData<M>: Unpin,
impl<M> UnsafeUnpin for GradientsAccumulator<M>where
PhantomData<M>: UnsafeUnpin,
Blanket Implementations§
Source§impl<T> BorrowMut<T> for Twhere
T: ?Sized,
impl<T> BorrowMut<T> for Twhere
T: ?Sized,
Source§fn borrow_mut(&mut self) -> &mut T
fn borrow_mut(&mut self) -> &mut T
Mutably borrows from an owned value. Read more