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§
Provided Methods§
Sourcefn step_batch(&mut self, items: &mut [OptItem<'_>])
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.
Sourcefn end_iteration(&mut self)
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?
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}Sourcefn set_lr(&mut self, _lr: f32)
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.
Sourcefn lr_scale(&self, _name: &str) -> f32
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.
Sourcefn state_dict(&self) -> Option<OptimizerState>
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.
Sourcefn load_state_dict(&mut self, _state: &OptimizerState) -> bool
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".