pub struct GradientsParams { /* private fields */ }Expand description
Data type that contains gradients for parameters.
Implementations§
Source§impl GradientsParams
impl GradientsParams
Sourcepub fn try_to_record<B: Backend>(
&self,
) -> Result<GradientsParamsRecord, RecorderError>
pub fn try_to_record<B: Backend>( &self, ) -> Result<GradientsParamsRecord, RecorderError>
Snapshot all gradients without clearing them, including accumulated gradients.
B must be the backend used to register the gradients (normally the
autodiff backend’s inner backend). Device reads complete before returning.
Sourcepub async fn to_record_async<B: Backend>(
&self,
) -> Result<GradientsParamsRecord, RecorderError>
pub async fn to_record_async<B: Backend>( &self, ) -> Result<GradientsParamsRecord, RecorderError>
Asynchronously snapshot all gradients without clearing the container.
B must match the backend used to register the gradients.
Sourcepub fn from_record<B: Backend>(
record: GradientsParamsRecord,
device: &B::Device,
) -> Result<Self, RecorderError>
pub fn from_record<B: Backend>( record: GradientsParamsRecord, device: &B::Device, ) -> Result<Self, RecorderError>
Restore recorded gradients to a device without summing or clearing entries.
Restore the model’s parameter IDs from the same checkpoint before using
these gradients. B is the gradient backend, normally the inner backend.
Source§impl GradientsParams
impl GradientsParams
Sourcepub fn new() -> Self
pub fn new() -> Self
Creates a new GradientsParams.
Sourcepub fn from_grads<B: AutodiffBackend, M: AutodiffModule<B>>(
grads: B::Gradients,
module: &M,
) -> Self
pub fn from_grads<B: AutodiffBackend, M: AutodiffModule<B>>( grads: B::Gradients, module: &M, ) -> Self
Extract each tensor gradients for the given module.
Note: This consumes the gradients. See [‘from_module’] to extract gradients only for a specific module.
Sourcepub fn from_module<B: AutodiffBackend, M: AutodiffModule<B>>(
grads: &mut B::Gradients,
module: &M,
) -> Self
pub fn from_module<B: AutodiffBackend, M: AutodiffModule<B>>( grads: &mut B::Gradients, module: &M, ) -> Self
Extract each tensor gradients for the given module.
Sourcepub fn from_params<B: AutodiffBackend, M: AutodiffModule<B>>(
grads: &mut B::Gradients,
module: &M,
params: &[ParamId],
) -> Self
pub fn from_params<B: AutodiffBackend, M: AutodiffModule<B>>( grads: &mut B::Gradients, module: &M, params: &[ParamId], ) -> Self
Extract tensor gradients for the given module and given parameters.
Sourcepub fn get<B, const D: usize>(&self, id: ParamId) -> Option<Tensor<B, D>>where
B: Backend,
pub fn get<B, const D: usize>(&self, id: ParamId) -> Option<Tensor<B, D>>where
B: Backend,
Get the gradients for the given parameter id.
§Notes
You should use remove if you want to get the gradients only one time.
Sourcepub fn remove<B, const D: usize>(&mut self, id: ParamId) -> Option<Tensor<B, D>>where
B: Backend,
pub fn remove<B, const D: usize>(&mut self, id: ParamId) -> Option<Tensor<B, D>>where
B: Backend,
Remove the gradients for the given parameter id.
Sourcepub fn register<B, const D: usize>(&mut self, id: ParamId, value: Tensor<B, D>)where
B: Backend,
pub fn register<B, const D: usize>(&mut self, id: ParamId, value: Tensor<B, D>)where
B: Backend,
Register a gradients tensor for the given parameter id.
§Notes
If a tensor is already registered for the given parameter id, it will be replaced.
Sourcepub fn to_device<B: AutodiffBackend, M: AutodiffModule<B>>(
self,
device: &B::Device,
module: &M,
) -> Self
pub fn to_device<B: AutodiffBackend, M: AutodiffModule<B>>( self, device: &B::Device, module: &M, ) -> Self
Change the device of each tensor gradients registered for the given module.