Skip to main content

Optimizer

Trait Optimizer 

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

Source

fn step( &mut self, params: &Array<A, D>, gradients: &Array<A, D>, ) -> Result<Array<A, D>>

Updates parameters using the given gradients

§Arguments
  • params - The current parameter values
  • gradients - The gradients of the parameters
§Returns

The updated parameters

Source

fn get_learning_rate(&self) -> A

Gets the current learning rate

Source

fn set_learning_rate(&mut self, learning_rate: A)

Sets a new learning rate

Provided Methods§

Source

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 arrays
  • gradients_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§

Source§

impl Optimizer<f32, Dim<[usize; 1]>> for SimdSGD<f32>

Source§

impl Optimizer<f64, Dim<[usize; 1]>> for SimdSGD<f64>

Source§

impl<A, D> Optimizer<A, D> for Adagrad<A>
where A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension,

Source§

impl<A, D> Optimizer<A, D> for Adam<A>
where A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension,

Source§

impl<A, D> Optimizer<A, D> for AdamW<A>
where A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension,

Source§

impl<A, D> Optimizer<A, D> for ChainedOptimizer<A, D>

Source§

impl<A, D> Optimizer<A, D> for HybridQuantumClassical<A>
where A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension,

Source§

impl<A, D> Optimizer<A, D> for LAMB<A>
where A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension,

Source§

impl<A, D> Optimizer<A, D> for LBFGS<A>
where A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension,

Source§

impl<A, D> Optimizer<A, D> for Lion<A>
where A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension,

Source§

impl<A, D> Optimizer<A, D> for MAML<A>

Source§

impl<A, D> Optimizer<A, D> for MetaSGD<A>

Source§

impl<A, D> Optimizer<A, D> for NtmOptimizer<A>

Source§

impl<A, D> Optimizer<A, D> for ParallelOptimizer<A, D>

Source§

impl<A, D> Optimizer<A, D> for QuantumAnnealing<A>
where A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension,

Source§

impl<A, D> Optimizer<A, D> for RAdam<A>
where A: Float + ScalarOperand + Debug + Send + Sync + From<f64>, D: Dimension,

Source§

impl<A, D> Optimizer<A, D> for RMSprop<A>
where A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension,

Source§

impl<A, D> Optimizer<A, D> for ReptileOptimizer<A>

Source§

impl<A, D> Optimizer<A, D> for SGD<A>
where A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension,

Source§

impl<A, D> Optimizer<A, D> for SequentialOptimizer<A, D>

Source§

impl<A, D> Optimizer<A, D> for VariationalQuantumOptimizer<A>
where A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension,

Source§

impl<A, D> Optimizer<A, D> for WeightedOptimizer<A, D>

Source§

impl<A, O, D> Optimizer<A, D> for Lookahead<A, O, D>
where A: Float + ScalarOperand + Debug + Send + Sync, O: Optimizer<A, D> + Clone + Send + Sync, D: Dimension,

Source§

impl<A, O, D> Optimizer<A, D> for SAM<A, O, D>
where A: Float + ScalarOperand + Debug + Send + Sync, O: Optimizer<A, D> + Clone + Send + Sync, D: Dimension,

Source§

impl<A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension + Send + Sync> Optimizer<A, D> for GroupedAdam<A, D>

Source§

impl<A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension + Send + Sync> Optimizer<A, D> for LARS<A>

Source§

impl<A> Optimizer<A, Dim<[usize; 1]>> for SparseAdam<A>
where A: Float + ScalarOperand + Debug + Send + Sync,

Source§

impl<F, D> Optimizer<F, D> for MetricOptimizer<F, D>
where F: Float + Debug + Display + FromPrimitive + ScalarOperand + 'static, D: Dimension + 'static,

Available on crate feature metrics-integration only.
Source§

impl<T> Optimizer<T, Dim<[usize; 1]>> for AdaBound<T>
where T: Float + ScalarOperand + Debug + Send + Sync,

Source§

impl<T> Optimizer<T, Dim<[usize; 1]>> for AdaDelta<T>
where T: Float + ScalarOperand + Debug + Send + Sync,

Source§

impl<T> Optimizer<T, Dim<[usize; 1]>> for Ranger<T>
where T: Float + ScalarOperand + Debug + Send + Sync,