pub trait Optimizer<A, D>{
// Required methods
fn step(
&mut self,
params: &Array<A, D>,
gradients: &Array<A, D>,
) -> Result<Array<A, D>>;
fn get_learning_rate(&self) -> A;
fn set_learning_rate(&mut self, learning_rate: A);
// Provided method
fn step_list(
&mut self,
params_list: &[&Array<A, D>],
gradients_list: &[&Array<A, D>],
) -> Result<Vec<Array<A, D>>> { ... }
}Expand description
Trait that defines the interface for optimization algorithms
Required Methods§
Sourcefn get_learning_rate(&self) -> A
fn get_learning_rate(&self) -> A
Gets the current learning rate
Sourcefn set_learning_rate(&mut self, learning_rate: A)
fn set_learning_rate(&mut self, learning_rate: A)
Sets a new learning rate
Provided Methods§
Sourcefn step_list(
&mut self,
params_list: &[&Array<A, D>],
gradients_list: &[&Array<A, D>],
) -> Result<Vec<Array<A, D>>>
fn step_list( &mut self, params_list: &[&Array<A, D>], gradients_list: &[&Array<A, D>], ) -> Result<Vec<Array<A, D>>>
Updates multiple parameter arrays at once
§State contract
Position i in params_list identifies parameter tensor i and must get
its own optimizer state (moments, accumulators, velocities and any per-tensor
timestep). The caller is expected to pass the tensors in a stable order across
calls, exactly like PyTorch’s parameter groups.
The default implementation below simply forwards to Optimizer::step, which
is only correct for stateless optimizers. Every stateful optimizer in this
crate overrides step_list and routes each index to a dedicated state slot
(see e.g. Adam::step_indexed). Implementors of new stateful optimizers must
do the same: relying on the default makes all tensors share one state slot, so
they reset each other on every shape change and their bias correction advances
once per tensor instead of once per step.
§Arguments
params_list- List of parameter arraysgradients_list- List of gradient arrays corresponding to the parameters
§Returns
Updated parameter arrays
Dyn Compatibility§
This trait is dyn compatible.
In older versions of Rust, dyn compatibility was called "object safety".
Implementors§
impl Optimizer<f32, Dim<[usize; 1]>> for SimdSGD<f32>
impl Optimizer<f64, Dim<[usize; 1]>> for SimdSGD<f64>
impl<A, D> Optimizer<A, D> for Adagrad<A>
impl<A, D> Optimizer<A, D> for Adam<A>
impl<A, D> Optimizer<A, D> for AdamW<A>
impl<A, D> Optimizer<A, D> for ChainedOptimizer<A, D>
impl<A, D> Optimizer<A, D> for HybridQuantumClassical<A>
impl<A, D> Optimizer<A, D> for LAMB<A>
impl<A, D> Optimizer<A, D> for LBFGS<A>
impl<A, D> Optimizer<A, D> for Lion<A>
impl<A, D> Optimizer<A, D> for MAML<A>
impl<A, D> Optimizer<A, D> for MetaSGD<A>
impl<A, D> Optimizer<A, D> for NtmOptimizer<A>
impl<A, D> Optimizer<A, D> for ParallelOptimizer<A, D>
impl<A, D> Optimizer<A, D> for QuantumAnnealing<A>
impl<A, D> Optimizer<A, D> for RAdam<A>
impl<A, D> Optimizer<A, D> for RMSprop<A>
impl<A, D> Optimizer<A, D> for ReptileOptimizer<A>
impl<A, D> Optimizer<A, D> for SGD<A>
impl<A, D> Optimizer<A, D> for SequentialOptimizer<A, D>
impl<A, D> Optimizer<A, D> for VariationalQuantumOptimizer<A>
impl<A, D> Optimizer<A, D> for WeightedOptimizer<A, D>
impl<A, O, D> Optimizer<A, D> for Lookahead<A, O, D>
impl<A, O, D> Optimizer<A, D> for SAM<A, O, D>
impl<A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension + Send + Sync> Optimizer<A, D> for GroupedAdam<A, D>
impl<A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension + Send + Sync> Optimizer<A, D> for LARS<A>
impl<A> Optimizer<A, Dim<[usize; 1]>> for SparseAdam<A>
impl<F, D> Optimizer<F, D> for MetricOptimizer<F, D>
metrics-integration only.