Skip to main content

Optimizer

Trait Optimizer 

Source
pub trait Optimizer {
    // Required method
    fn step(
        &mut self,
        name: &str,
        shape: &[usize],
        param: &mut [f32],
        grad: &[f32],
    );

    // Provided methods
    fn step_batch(&mut self, items: &mut [OptItem<'_>]) { ... }
    fn end_iteration(&mut self) { ... }
    fn set_lr(&mut self, _lr: f32) { ... }
    fn lr_scale(&self, _name: &str) -> f32 { ... }
    fn state_dict(&self) -> Option<OptimizerState> { ... }
    fn load_state_dict(&mut self, _state: &OptimizerState) -> bool { ... }
}

Required Methods§

Source

fn step(&mut self, name: &str, shape: &[usize], param: &mut [f32], grad: &[f32])

Provided Methods§

Source

fn step_batch(&mut self, items: &mut [OptItem<'_>])

Batched step over ALL parameters in one call. Default: sequential step per item — bit-identical to the per-parameter loop. Optimizers whose parameter groups are independent (e.g. Muon on the 2-D weight matrices vs AdamW on the embeddings/biases/norms) can override this to run the groups on separate threads; because the groups touch disjoint parameters and disjoint optimizer state, the result is bit-for-bit the same as the serial loop — only the wall-clock (the CPU-side optimizer bubble) shrinks toward max(group_times) instead of their sum.

Source

fn end_iteration(&mut self)

Advance the global step counter. Most algorithms increment per call to step, so most implementations leave this a no-op.

Examples found in repository?
examples/step_bench.rs (line 37)
30fn bench(label: &str, n: usize, mut opt: Box<dyn Optimizer>) -> f64 {
31    let shape = [n];
32    let mut param = fill(n, 7);
33    let grad = fill(n, 11);
34    // Warm: the first call allocates the moment buffers.
35    for _ in 0..5 {
36        opt.step("w", &shape, &mut param, &grad);
37        opt.end_iteration();
38    }
39    let iters = (200_000_000 / n).clamp(20, 2000);
40    let start = Instant::now();
41    for _ in 0..iters {
42        opt.step("w", &shape, &mut param, &grad);
43        opt.end_iteration();
44    }
45    let per_step = start.elapsed().as_secs_f64() / iters as f64;
46    let per_elem_ns = per_step * 1e9 / n as f64;
47    println!(
48        "  {label:22} {:>9.1} µs/step  {per_elem_ns:>6.2} ns/element  {:>6.2} GB/s",
49        per_step * 1e6,
50        // Adam touches p, m, v (read+write) and g (read): 7 × 4 bytes per element.
51        (n as f64 * 7.0 * 4.0) / per_step / 1e9
52    );
53    per_elem_ns
54}
Source

fn set_lr(&mut self, _lr: f32)

Set the base learning rate (for LR schedules / warmup). Default is a no-op for algorithms without a scalar lr (e.g. Adafactor’s relative step size); every algorithm in this crate that has an lr field overrides this to update it.

Source

fn lr_scale(&self, _name: &str) -> f32

Per-tensor multiplier on the effective learning rate. Default is 1.0 for every name. Override when wrapping this crate to support per-name LR schedules (e.g. embedding-vs-attention splits, or the Gaussian-splat attribute-typed LR setup). The CPU impls in this crate currently honor this only when the caller passes a pre-scaled lr for the relevant call — backends are encouraged to consult it inside their fused kernel.

Source

fn state_dict(&self) -> Option<OptimizerState>

Snapshot the optimizer’s state for a checkpoint.

None means this algorithm has not opted in, so a run using it cannot be resumed exactly. Callers should say so rather than silently restarting the accumulators at zero.

Source

fn load_state_dict(&mut self, _state: &OptimizerState) -> bool

Restore a snapshot. Returns false when unsupported or when the state does not belong to this algorithm.

Dyn Compatibility§

This trait is dyn compatible.

In older versions of Rust, dyn compatibility was called "object safety".

Implementors§