Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
31 commits
Select commit Hold shift + click to select a range
f8f3ad3
Let state_for stack a history axis
jessegrabowski Sep 17, 2026
2d8e5b7
Add an LBFGSDirection op with a scan inner graph
jessegrabowski Sep 17, 2026
e4ebd9f
Take the recursion's dot products through BLAS
jessegrabowski Sep 17, 2026
e4893cf
Dispatch LBFGSDirection to mlx as a loop over ring rows
jessegrabowski Sep 17, 2026
63c5d6f
Test the direction op against the dense BFGS matrix
jessegrabowski Sep 17, 2026
959b004
Keep each parameter's dtype through the recursion
jessegrabowski Sep 17, 2026
75af337
Compute the curvatures with the dot the rule's guard uses
jessegrabowski Sep 17, 2026
5e86542
Allocate optimizer state at the parameter's declared dtype
jessegrabowski Sep 17, 2026
fceed2a
Add lbfgs_updates with a ring memory and curvature guard
jessegrabowski Sep 17, 2026
8b5a8de
Raise the mlx floor to 0.32.1 and dot through its vector matmul
jessegrabowski Sep 17, 2026
13eba45
Set the mlx floor to 0.32.2
jessegrabowski Sep 17, 2026
05d2ffb
Add the lbfgs alias
jessegrabowski Sep 17, 2026
a20f465
Build literal count and gamma at their dtype instead of casting
jessegrabowski Sep 17, 2026
914fffd
Require pytensor 3.3.2
jessegrabowski Sep 17, 2026
41274c2
Cast the mlx recursion to each parameter's dtype
jessegrabowski Oct 2, 2026
5177de8
Cache each admitted pair's curvature in rule state
jessegrabowski Oct 2, 2026
215e6ba
Require pytensor 3.3.3
jessegrabowski Oct 2, 2026
08d6215
Move the L-BFGS test references into their own module
jessegrabowski Oct 2, 2026
c6d8338
Require static slot counts in LBFGSDirection
jessegrabowski Oct 2, 2026
f443bc5
Build gamma at the curvature dtype
jessegrabowski Oct 2, 2026
7744434
Carry the L-BFGS identity scale in rule state
jessegrabowski Oct 2, 2026
2a29785
Refuse L-BFGS pairs whose curvature overflows
jessegrabowski Oct 2, 2026
20e231d
Test L-BFGS on a one-slot ring
jessegrabowski Oct 2, 2026
3e036f3
Test that L-BFGS rejects a negligible positive curvature
jessegrabowski Oct 2, 2026
e5629c2
Test L-BFGS with parameters of mixed dtypes
jessegrabowski Oct 2, 2026
883e742
Test that L-BFGS resumes from a checkpoint
jessegrabowski Oct 2, 2026
218b635
Run lbfgs at its default rate in the loss-reduction test
jessegrabowski Oct 2, 2026
e2fbdf3
Test lbfgs with its step clipped after it
jessegrabowski Oct 2, 2026
a6fb4d1
Test that skip_if holds back the L-BFGS state
jessegrabowski Oct 2, 2026
18ee016
Document the fixed-step contract of lbfgs
jessegrabowski Oct 2, 2026
0bf07b0
List LBFGSDirection in the optim API docs
jessegrabowski Oct 2, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .github/workflows/run_tests.yml
Original file line number Diff line number Diff line change
Expand Up @@ -152,7 +152,7 @@ jobs:
run: |
conda activate pytensor_ml
if [[ $INSTALL_JAX == "1" ]]; then pip install "jax>=0.8,<0.9.1" jaxlib; fi
if [[ $INSTALL_MLX == "1" ]]; then pip install "mlx>=0.30,<0.32"; fi
if [[ $INSTALL_MLX == "1" ]]; then pip install "mlx>=0.32.2,<0.33"; fi
if [[ $INSTALL_TORCH == "1" ]]; then pip install torch --index-url https://download.pytorch.org/whl/cpu; fi
env:
INSTALL_JAX: ${{ matrix.install-jax }}
Expand Down
2 changes: 1 addition & 1 deletion conda_envs/environment-docs.yml
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ channels:
dependencies:
- python>=3.12
# Runtime deps: autodoc imports pytensor_ml, so the full runtime stack has to be in scope.
- pytensor>=3.3.0,<3.4.0
- pytensor>=3.3.3,<3.4.0
- numpy
- safetensors
# The gallery extension renders notebook thumbnails with matplotlib.
Expand Down
2 changes: 1 addition & 1 deletion conda_envs/pytensor_ml-gpu_jax.yml
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@ channels:

dependencies:
- python>=3.12
- pytensor>=3.2.3,<4.0.0
- pytensor>=3.3.3,<4.0.0
- numpy
- scikit-learn

Expand Down
2 changes: 1 addition & 1 deletion conda_envs/pytensor_ml.yml
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ channels:

dependencies:
- python>=3.12
- pytensor>=3.3.0,<3.4.0
- pytensor>=3.3.3,<3.4.0
- numpy
- safetensors
- scikit-learn
Expand Down
9 changes: 9 additions & 0 deletions docs/source/api/optim.rst
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@ Update rules
rprop
adagrad
adadelta
lbfgs

Transforms
----------
Expand Down Expand Up @@ -130,3 +131,11 @@ Low-level update functions
rprop_updates
adagrad_updates
adadelta_updates
lbfgs_updates

.. currentmodule:: pytensor_ml.optim.lbfgs

.. autosummary::
:toctree: generated/

LBFGSDirection
19 changes: 19 additions & 0 deletions docs/source/references.bib
Original file line number Diff line number Diff line change
Expand Up @@ -42,3 +42,22 @@ @inproceedings{glorot2010init
booktitle = {International Conference on Artificial Intelligence and Statistics},
year = {2010},
}

@book{nocedal2006numerical,
title = {Numerical Optimization},
author = {Nocedal, Jorge and Wright, Stephen J.},
edition = {2},
publisher = {Springer},
address = {New York},
year = {2006},
}

@article{liu1989lbfgs,
title = {On the Limited Memory {BFGS} Method for Large Scale Optimization},
author = {Liu, Dong C. and Nocedal, Jorge},
journal = {Mathematical Programming},
volume = {45},
number = {1--3},
pages = {503--528},
year = {1989},
}
4 changes: 2 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,7 @@ keywords = [
]

dependencies = [
"pytensor>=3.2.3,<4.0.0",
"pytensor>=3.3.3,<4.0.0",
"numpy",
]

Expand Down Expand Up @@ -163,7 +163,7 @@ platforms = ["osx-arm64", "linux-64", "win-64"]
# the two lists have to move together.
[tool.pixi.feature.docs.dependencies]
python = ">=3.12"
pytensor = ">=3.3.0,<3.4.0"
pytensor = ">=3.3.3,<3.4.0"
numpy = "*"
safetensors = "*"
# The gallery extension renders notebook thumbnails with matplotlib.
Expand Down
1 change: 1 addition & 0 deletions pytensor_ml/dispatch/mlx/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,4 +2,5 @@
# marker op that gets a kernel, mirroring the layout under pytensor_ml/layers.
import pytensor_ml.dispatch.mlx.attention
import pytensor_ml.dispatch.mlx.conv
import pytensor_ml.dispatch.mlx.lbfgs
import pytensor_ml.dispatch.mlx.pooling
51 changes: 51 additions & 0 deletions pytensor_ml/dispatch/mlx/lbfgs.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
import mlx.core as mx

from pytensor.link.mlx.dispatch import mlx_funcify

from pytensor_ml.optim.lbfgs import LBFGSDirection


@mlx_funcify.register(LBFGSDirection)
def mlx_funcify_LBFGSDirection(op, node=None, **kwargs):
"""Run the two-loop recursion as a Python loop over ``mx`` ops, since mlx has no scan."""
n, m = op.n_parameters, op.memory_size

def rows(stacks, slot):
# `count` is traced under mx.compile, so the slot is an mx scalar and the row is gathered rather
# than indexed from Python.
return [mx.take(stack, slot, axis=0) for stack in stacks]

def dot(left, right):
# Vector matmul is the fastest dot mlx has from 0.32.2 (ml-explore/mlx#3580); before that it ran
# one threadgroup and was slower than a fused reduction by two orders of magnitude.
return sum(a.reshape(-1) @ b.reshape(-1) for a, b in zip(left, right))

def direction(count, gamma, rho, *tensors):
gradients = tensors[:n]
S = tensors[n : 2 * n]
Y = tensors[2 * n :]

# Every row is gathered once, in ring order (oldest first), and reused by both loops.
order = [(count + offset) % m for offset in range(m)]
s_rows = [rows(S, slot) for slot in order]
y_rows = [rows(Y, slot) for slot in order]
curvatures = [mx.take(rho, slot) for slot in order]

vector = list(gradients)
alphas = [None] * m
for position in reversed(range(m)):
alphas[position] = curvatures[position] * dot(s_rows[position], vector)
vector = [
v - alphas[position].astype(v.dtype) * y_p
for v, y_p in zip(vector, y_rows[position])
]
vector = [gamma.astype(v.dtype) * v for v in vector]
for position in range(m):
beta = curvatures[position] * dot(y_rows[position], vector)
vector = [
v + (alphas[position] - beta).astype(v.dtype) * s_p
for v, s_p in zip(vector, s_rows[position])
]
return vector[0] if n == 1 else tuple(vector)

return direction
4 changes: 4 additions & 0 deletions pytensor_ml/optim/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
adam,
adamax,
adamw,
lbfgs,
nadam,
rmsprop,
rprop,
Expand Down Expand Up @@ -41,6 +42,7 @@
adam_updates,
adamax_updates,
adamw_updates,
lbfgs_updates,
nadam_updates,
rmsprop_updates,
rprop_updates,
Expand Down Expand Up @@ -96,6 +98,8 @@
"get_gradients",
"join_schedules",
"large_step",
"lbfgs",
"lbfgs_updates",
"linear_onecycle_schedule",
"linear_schedule",
"nadam",
Expand Down
51 changes: 51 additions & 0 deletions pytensor_ml/optim/alias.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
adam_updates,
adamax_updates,
adamw_updates,
lbfgs_updates,
nadam_updates,
rmsprop_updates,
rprop_updates,
Expand Down Expand Up @@ -357,6 +358,56 @@ def rule(
return rule


def lbfgs(
learning_rate: LearningRate = 1.0,
memory_size: int = 10,
scale_init_precond: bool = True,
*,
namespace: str = "lbfgs",
) -> Transform:
"""
L-BFGS optimizer. See :func:`~pytensor_ml.optim.rules.lbfgs_updates` for the update rule.

``learning_rate`` accepts a float, a scalar shared variable, any scalar graph, or a schedule, and
``namespace`` prefixes the state this rule allocates; see :func:`sgd`.

Examples
--------
A quasi-Newton direction from a memory of recent parameter and gradient differences, taken at a
fixed fraction with no line search. It reads the change between consecutive gradients as curvature,
so the loss has to be the same function from one step to the next: full batch, no dropout. For the
same reason it takes the loss's own gradients: put a clip after it in a chain, never ahead of it.

.. code-block:: python

import numpy as np

from pytensor_ml.layers import Input, Linear
from pytensor_ml.loss import SquaredError, supervised_loss
from pytensor_ml.optim import compile_train, lbfgs

X = Input("X", shape=(None, 4))
loss, target = supervised_loss(Linear("fc", n_in=4, n_out=1)(X), SquaredError())

step = compile_train(loss, lbfgs(learning_rate=0.5, memory_size=10))
loss_value = step(np.zeros((8, 4)), np.zeros((8, 1)))
"""

def rule(
loss_gradients_or_updates: LossGradientsOrUpdates, parameters: Sequence[Parameter]
) -> Updates:
return lbfgs_updates(
loss_gradients_or_updates,
parameters,
learning_rate=learning_rate,
memory_size=memory_size,
scale_init_precond=scale_init_precond,
namespace=namespace,
)

return rule


def rmsprop(
learning_rate: LearningRate = 1e-2,
rho: float = 0.9,
Expand Down
25 changes: 22 additions & 3 deletions pytensor_ml/optim/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,8 @@ class Gradients(Updates):

What :func:`to_updates` produces from a loss, and what everything ahead of the first rule in a chain
sees. A clip placed here bounds the gradient itself, so a spike never reaches the moment estimates.
A rule that reads curvature from consecutive gradients, such as
:func:`~pytensor_ml.optim.alias.lbfgs`, needs them unclipped, so its clip goes after it.
"""


Expand Down Expand Up @@ -481,9 +483,11 @@ def _unreachable_parameter_names(
]


def state_for(parameter: Parameter, slot: str, fill_value: float = 0.0) -> Parameter:
def state_for(
parameter: Parameter, slot: str, fill_value: float = 0.0, history_size: int | None = None
) -> Parameter:
"""
Return the optimizer-state shared variable shaped and typed like ``parameter``.
Return the optimizer-state shared variable typed like ``parameter``, or a stack of them.

The variable is named ``"{parameter.name}/{slot}"`` and carries the parameter's layer, so a checkpoint
numbers it where it numbers the parameter. The name is never used to *find* the variable at runtime --
Expand All @@ -501,6 +505,9 @@ def state_for(parameter: Parameter, slot: str, fill_value: float = 0.0) -> Param
A short role tag for the slot, e.g. ``"adam/first_moment"`` or ``"trace/velocity"``.
fill_value : float
Constant to initialize the state with. Default 0.0.
history_size : int, optional
Number of past values to stack along a new leading axis, so the state is shaped
``(history_size, *parameter.shape)``. Omitted, the state has the parameter's own shape.

Returns
-------
Expand All @@ -527,9 +534,21 @@ def state_for(parameter: Parameter, slot: str, fill_value: float = 0.0) -> Param
f"Cannot allocate optimizer state {slot!r} for an unnamed parameter. Stateful optimizers rely on "
"parameter names to identify their state at serialization boundaries; give the parameter a name."
)
if history_size is not None and history_size < 1:
raise ValueError(f"history_size must be at least 1, got {history_size}.")

value = parameter.get_value(borrow=True)
state = pytensor.shared(np.full_like(value, fill_value), name=f"{parameter.name}/{slot}")
shape = value.shape if history_size is None else (history_size, *value.shape)
static_shape = (
parameter.type.shape if history_size is None else (history_size, *parameter.type.shape)
)
# The declared dtype rather than the value's: after a step on mlx the value is a device array
# whose dtype numpy cannot read.
state = pytensor.shared(
np.full(shape, fill_value, dtype=parameter.type.dtype),
name=f"{parameter.name}/{slot}",
shape=static_shape,
)
# Keeps `Linear_1_W` and `Linear_1_W/adam/first_moment` numbered onto the same layer.
state.layer_name = getattr(parameter, "layer_name", None)
return state
Expand Down
Loading
Loading