All pages
Writing optimizers
An optimizer is a function from a parameter, its gradient, its states and
ctx to the new parameter and the new states. The body is traced once per
parameter into that parameter's update graph.
(require "sexpgpu/optim" schedule-value update-diagnostics)
(defoptimizer momentum-sgd (param grad ctx &key (lr 0.01) (momentum 0.9))
(setq lr (schedule-value lr ctx))
(defstate velocity (zeros-like param))
(let* ((v (+ (* momentum velocity) grad))
(next (- param (* lr v))))
(update-diagnostics param grad next)
(optimizer-update next :velocity v)))
defoptimizer
(defoptimizer name (param grad ctx &key ...) body...) defines an
optimizer. Calling it with keyword values configures it:
(momentum-sgd :lr 0.05) is the value the contract's optimizer holds,
and what group :optimizer accepts; see
optimizer groups.
- Traced once per parameter, so a shape-dependent choice is a plain
compile-time
if: Muon transposes a tall matrix by asking(> (dim x 0) (dim x 1)). - Each update graph owns one parameter; there is no cross-parameter state.
- A body reports with
diagnostic, nevermetric(E-CONTRACT-011); see optimizer diagnostics.
States
(defstate name init) declares one per-parameter state.
initevaluates normally and must be a zero constant:zeros,zeros-likeor a zerofull, including casts, reshapes, transposes and broadcasts of those. Helpers, conditionals and macros may return it.- Nonzero or computed values (even ones that happen to be zero),
duplicate states and a
defstateoutside an optimizer areE-OPT-003. - Ordinary states are named in the IR and checkpoints by their names;
gensym states get stable
state.Nnames in declaration order.
optimizer-update
(optimizer-update new-param :state new-value ...) is the body's result.
- Every declared state is returned exactly once. A body that does not end
in
optimizer-update, or drops a state, isE-OPT-004. - A state is named by keyword,
:velocity, or by its quoted declaration symbol,',velocity, which is how a macro returns a gensym state. Never turn a gensym into a keyword to name its state. - It is an ordinary function, so
applycan pass a computed list of updates.
Schedules
A schedule is a function of ctx returning a scalar. Called while the
update graph is traced, it builds ordinary nodes in it.
(defun stable-then-decay (peak &key (decay-fraction 0.7))
(lambda (ctx)
(let ((progress (getf ctx :progress)))
(* peak (minimum 1.0 (/ (- 1.0 progress) decay-fraction))))))
To accept a number or a schedule, call (schedule-value lr ctx) from
sexpgpu/optim in the body: it calls a function with ctx and returns
anything else unchanged. The standard optimizers do this for every declared
hyperparameter. Other function arguments, a gradient transform for example,
stay callable by the body. The ctx keys are in
curriculum and ctx.
Related: sexpgpu/optim, optimizer groups, macros.