Add an L-BFGS rule with a ring-buffered memory and an mlx dispatch - #162
Merged
jessegrabowski merged 31 commits intoOct 2, 2026
Merged
Conversation
jessegrabowski
force-pushed
the
lbfgs-direction-rule
branch
from
September 19, 2026 13:50
3368eeb to
914fffd
Compare
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Adds L-BFGS at a fixed step size.
lbfgs_updatesinrules.pykeeps per-parameter stacks of the lastmemory_sizeparameter and gradient differences as a ring, written in place withset_subtensor, and admits a pair only wheny . s > eps * y . y, so the inverse-Hessian estimate stays positive definite.LBFGSDirection, aSymbolicOpinoptim/lbfgs.py, applies the two-loop recursion to the gradient with two scans over the ring order and gets a Python-loop dispatch indispatch/mlx/lbfgs.py, sinceScanhas no mlx dispatch. Thelbfgsalias takes a rate or a schedule like the other rules.The ring slot is a one-element index vector rather than a scalar. pytensor's mlx
Subtensordispatch cannot trace a scalar index (pymc-devs/pytensor#2422), and the advanced-indexing ops can.The mlx CI install now requires
mlx>=0.32.2. Below that, mlx's vector matmul ran one threadgroup and a 12.6M-element dot took 155 ms instead of 0.4 ms (ml-explore/mlx#3580).state_fortakeshistory_sizefor the stacked buffers andscalar_statetakes adtypefor the pair count. Tests check the op against the dense BFGS matrix built from its definition and the rule against the secant condition and a quadratic's closed-form minimizer, on numba and on mlx.The line search is the next PR. Part of #58.
📚 Documentation preview 📚: https://pytensor-ml--162.org.readthedocs.build/en/162/