From f8f3ad37c98e1e0a8de292b2a93118e61a81b886 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 17:02:29 +0800 Subject: [PATCH 01/31] Let state_for stack a history axis --- pytensor_ml/optim/base.py | 19 ++++++++++++++++--- tests/optim/test_training.py | 18 ++++++++++++++++++ 2 files changed, 34 insertions(+), 3 deletions(-) diff --git a/pytensor_ml/optim/base.py b/pytensor_ml/optim/base.py index 3d2c963..b6125f1 100644 --- a/pytensor_ml/optim/base.py +++ b/pytensor_ml/optim/base.py @@ -481,9 +481,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 -- @@ -501,6 +503,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 ------- @@ -527,9 +532,17 @@ 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) + state = pytensor.shared( + np.full(shape, fill_value, dtype=value.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 diff --git a/tests/optim/test_training.py b/tests/optim/test_training.py index 4261f1c..fae1607 100644 --- a/tests/optim/test_training.py +++ b/tests/optim/test_training.py @@ -218,6 +218,24 @@ def test_state_for_requires_named_parameter(): state_for(anonymous, "adam/first_moment") +def test_state_for_history_stacks_a_leading_axis(): + # The stack carries the parameter's static shape as well as its value's, so a write of the parameter + # itself into one slot type-checks; a `(?,)`-typed buffer would refuse a `(3,)`-typed parameter. + parameter = trainable(np.ones(3, dtype=config.floatX), name="w") + + stack = state_for(parameter, "lbfgs/value_differences", history_size=4) + + assert stack.type.shape == (4, 3) + assert stack.get_value().shape == (4, 3) + assert stack.get_value().dtype == parameter.get_value().dtype + + +def test_state_for_rejects_an_empty_history(): + parameter = trainable(np.ones(3, dtype=config.floatX), name="w") + with pytest.raises(ValueError, match="history_size must be at least 1"): + state_for(parameter, "lbfgs/value_differences", history_size=0) + + def test_compile_train_rejects_duplicate_parameter_names(): # Two parameters sharing a name give their optimizer state colliding names; compile_train refuses to # build a training step whose checkpointed state cannot be told apart. From 2d8e5b74bdd0f7e6b5133f7f0a32ca471f478c9e Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 17:52:55 +0800 Subject: [PATCH 02/31] Add an LBFGSDirection op with a scan inner graph --- docs/source/references.bib | 19 +++++ pytensor_ml/optim/lbfgs.py | 168 +++++++++++++++++++++++++++++++++++++ tests/optim/test_lbfgs.py | 122 +++++++++++++++++++++++++++ 3 files changed, 309 insertions(+) create mode 100644 pytensor_ml/optim/lbfgs.py create mode 100644 tests/optim/test_lbfgs.py diff --git a/docs/source/references.bib b/docs/source/references.bib index f5fe3e2..3ba15b7 100644 --- a/docs/source/references.bib +++ b/docs/source/references.bib @@ -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}, +} diff --git a/pytensor_ml/optim/lbfgs.py b/pytensor_ml/optim/lbfgs.py new file mode 100644 index 0000000..2b3525a --- /dev/null +++ b/pytensor_ml/optim/lbfgs.py @@ -0,0 +1,168 @@ +from collections.abc import Sequence + +import pytensor +import pytensor.tensor as pt + +from pytensor.compile.builders import SymbolicOp +from pytensor.graph.basic import Variable +from pytensor.tensor import TensorVariable + + +class LBFGSDirection(SymbolicOp): + r""" + Multiply a gradient by the L-BFGS inverse-Hessian approximation that a ring-buffered memory defines. + + Inputs are ``count, gamma, g_1..g_n, S_1..S_n, Y_1..Y_n`` and outputs are ``d_1..d_n = H g``, one per + parameter. ``S_p`` and ``Y_p`` are ``(memory_size, *shape)`` stacks of past parameter differences + :math:`s` and gradient differences :math:`y` for parameter ``p``, written as a ring: slot + ``(count - 1) % memory_size`` holds the newest pair and ``count`` is the number of pairs written so + far. A slot that holds nothing yet is all zeros and contributes nothing to the recursion, and a + writer retires a slot the same way. The op applies whatever pairs it is given: admitting only pairs + with :math:`y^\top s > 0`, which keeps the approximation positive definite, is the writer's job. + + The product is the two-loop recursion, algorithm 7.4 of :cite:t:`nocedal2006numerical`, with + :math:`\rho_i = 1 / (y_i^\top s_i)`. Each dot product sums over every parameter, so the memory of a + model with several parameters is treated as one vector and never copied into one. Starting from + :math:`\gamma I`, + + .. math:: + + q &\leftarrow g \\ + \alpha_i &= \rho_i s_i^\top q, \quad q \leftarrow q - \alpha_i y_i \quad \text{newest to oldest} \\ + r &\leftarrow \gamma q \\ + \beta_i &= \rho_i y_i^\top r, \quad r \leftarrow r + (\alpha_i - \beta_i) s_i \quad \text{oldest to newest} + + and :math:`d = r`. The loops are ``scan``s over the ring order, so the inner graph runs on any backend + with a scan dispatch, and a backend without one registers its own implementation of this op. + + Parameters + ---------- + n_parameters : int + How many parameters the gradient and memory are split across. + memory_size : int + Number of slots in each memory stack. + + Examples + -------- + Compile the direction for one vector parameter and a memory of four slots, with one pair written: + + .. code-block:: python + + import pytensor + import pytensor.tensor as pt + + from pytensor_ml.optim.lbfgs import LBFGSDirection + + g = pt.vector("g") + S = pt.matrix("S") + Y = pt.matrix("Y") + d = LBFGSDirection(n_parameters=1, memory_size=4)(1, 1.0, g, S, Y) + direction = pytensor.function([g, S, Y], d) + + References + ---------- + The limited-memory update is from :cite:t:`liu1989lbfgs`. + """ + + __props__ = ("n_parameters", "memory_size") + n_parameters: int + memory_size: int + + def __init__(self, input_types=None, **kwargs): + super().__init__(input_types, **kwargs) + if self.n_parameters < 1: + raise ValueError(f"n_parameters must be at least 1, got {self.n_parameters}.") + if self.memory_size < 1: + raise ValueError(f"memory_size must be at least 1, got {self.memory_size}.") + + @staticmethod + def filter_inputs(*inputs: Variable | float | int) -> tuple[Variable, ...]: + count, gamma, *raw = inputs + tensors = [pt.as_tensor_variable(tensor) for tensor in raw] + return ( + pt.as_tensor_variable(count).astype("int64"), + pt.as_tensor_variable(gamma).astype(tensors[0].dtype), + *tensors, + ) + + def build_inner_graph(self, *inputs: TensorVariable) -> list[Variable]: + n, m = self.n_parameters, self.memory_size + count, gamma, *tensors = inputs + if len(tensors) != 3 * n: + raise ValueError( + f"LBFGSDirection with n_parameters={n} takes {3 * n} tensors after count and gamma, a " + f"gradient and two memory stacks per parameter, but got {len(tensors)}." + ) + gradients = tensors[:n] + S = tensors[n : 2 * n] + Y = tensors[2 * n :] + for index, (gradient, s, y) in enumerate(zip(gradients, S, Y)): + for stack in (s, y): + _require_stack_of(stack, gradient, m, index) + + order = (count + pt.arange(m)) % m + curvatures = _curvatures(S, Y, m) + + def right_product(slot, *vector): + s = [stack[slot] for stack in S] + y = [stack[slot] for stack in Y] + alpha = curvatures[slot] * _dot(s, vector) + return [v - alpha * y_p for v, y_p in zip(vector, y)] + [alpha] + + *q, alphas = pytensor.scan( + right_product, + sequences=[order], + outputs_info=[*gradients, None], + go_backwards=True, + return_updates=False, + ) + r = [gamma * v[-1] for v in q] + + def left_product(slot, alpha, *vector): + s = [stack[slot] for stack in S] + y = [stack[slot] for stack in Y] + beta = curvatures[slot] * _dot(y, vector) + return [v + (alpha - beta) * s_p for v, s_p in zip(vector, s)] + + # The backward loop reports its alphas newest first and the forward loop reads them oldest first. + r = pytensor.scan( + left_product, + sequences=[order, alphas[::-1]], + outputs_info=r, + return_updates=False, + ) + if n == 1: + r = [r] + return [v[-1] for v in r] + + +def _require_stack_of( + stack: TensorVariable, gradient: TensorVariable, memory_size: int, index: int +) -> None: + """Raise unless ``stack`` is ``memory_size`` slots of ``gradient``'s shape and dtype.""" + slots = stack.type.shape[0] if stack.type.ndim else None + if ( + stack.type.ndim != gradient.type.ndim + 1 + or stack.type.dtype != gradient.type.dtype + or (slots is not None and slots != memory_size) + ): + raise ValueError( + f"The memory stacks of parameter {index} must be shaped (memory_size={memory_size}, " + f"*gradient.shape) at the gradient's dtype, but got {stack.type} for a gradient of type " + f"{gradient.type}." + ) + + +def _dot(left: Sequence[TensorVariable], right: Sequence[TensorVariable]) -> TensorVariable: + return pt.sum([pt.sum(a * b) for a, b in zip(left, right)]) + + +def _curvatures( + S: Sequence[TensorVariable], Y: Sequence[TensorVariable], memory_size: int +) -> TensorVariable: + """Return ``1 / (y_i . s_i)`` per slot, and zero for an empty slot rather than a division by zero.""" + products = pt.sum( + [pt.sum((s * y).reshape((memory_size, -1)), axis=1) for s, y in zip(S, Y)], axis=0 + ) + empty = pt.eq(products, 0.0) + return pt.switch(empty, 0.0, 1.0 / pt.switch(empty, 1.0, products)) diff --git a/tests/optim/test_lbfgs.py b/tests/optim/test_lbfgs.py new file mode 100644 index 0000000..db7252c --- /dev/null +++ b/tests/optim/test_lbfgs.py @@ -0,0 +1,122 @@ +import numpy as np +import pytensor +import pytensor.tensor as pt +import pytest + +from pytensor_ml.optim.lbfgs import LBFGSDirection +from pytensor_ml.pytensorf import function + +floatX = pytensor.config.floatX +RTOL = 1e-6 if floatX == "float64" else 1e-4 + + +def two_loop_direction(gamma, gradient, pairs): + """Nocedal and Wright 7.4 on one flat vector, over ``(s, y)`` pairs given oldest first.""" + q = gradient.astype(np.float64) + alphas = [] + for s, y in reversed(pairs): + rho = 1.0 / (y @ s) + alphas.append(rho * (s @ q)) + q = q - alphas[-1] * y + r = gamma * q + for (s, y), alpha in zip(pairs, reversed(alphas)): + rho = 1.0 / (y @ s) + beta = rho * (y @ r) + r = r + (alpha - beta) * s + return r + + +def ring_stacks(pairs, memory_size, count, shapes): + """Lay chronological flat pairs into per-parameter ring stacks, newest at ``(count - 1) % memory_size``.""" + S = [np.zeros((memory_size, *shape), dtype=floatX) for shape in shapes] + Y = [np.zeros((memory_size, *shape), dtype=floatX) for shape in shapes] + splits = np.cumsum([int(np.prod(shape)) for shape in shapes])[:-1] + for age, (s, y) in enumerate(reversed(pairs)): + slot = (count - 1 - age) % memory_size + for stack, piece in zip(S, np.split(s, splits)): + stack[slot] = piece.reshape(stack.shape[1:]) + for stack, piece in zip(Y, np.split(y, splits)): + stack[slot] = piece.reshape(stack.shape[1:]) + return S, Y + + +@pytest.mark.parametrize("n_pairs, count", [(2, 2), (4, 6)], ids=["not_yet_wrapped", "wrapped"]) +def test_direction_matches_the_two_loop_recursion_over_a_ring(n_pairs, count): + # The reference sees a flat vector and a chronological list, so it shares no ring or reshape + # arithmetic with the op. Two parameters of different rank exercise the cross-parameter dot products. + # Before the ring wraps its empty slots lead the order; after, the newest pair sits mid-ring. + rng = np.random.default_rng(0) + shapes = [(3, 2), (4,)] + size = sum(int(np.prod(shape)) for shape in shapes) + memory_size, gamma = 4, 0.7 + gradient = rng.normal(size=size).astype(floatX) + pairs = [] + for _ in range(n_pairs): + s = rng.normal(size=size).astype(floatX) + pairs.append((s, rng.normal(size=size).astype(floatX) + 0.5 * s)) # keeps y . s > 0 + S, Y = ring_stacks(pairs, memory_size, count, shapes) + + op = LBFGSDirection(n_parameters=2, memory_size=memory_size) + gradients = [pt.tensor(f"g{i}", shape=shape) for i, shape in enumerate(shapes)] + S_in = [pt.tensor(f"S{i}", shape=(memory_size, *shape)) for i, shape in enumerate(shapes)] + Y_in = [pt.tensor(f"Y{i}", shape=(memory_size, *shape)) for i, shape in enumerate(shapes)] + direction = function( + [*gradients, *S_in, *Y_in], op(count, gamma, *gradients, *S_in, *Y_in, return_list=True) + ) + + splits = np.cumsum([int(np.prod(shape)) for shape in shapes])[:-1] + gradient_pieces = [ + piece.reshape(shape) for piece, shape in zip(np.split(gradient, splits), shapes) + ] + got = np.concatenate([d.ravel() for d in direction(*gradient_pieces, *S, *Y)]) + np.testing.assert_allclose(got, two_loop_direction(gamma, gradient, pairs), rtol=RTOL) + + +def test_an_empty_memory_scales_the_gradient(): + # Built at float32 whatever floatX is, so the Python-float gamma has a narrower dtype to upcast. + rng = np.random.default_rng(1) + g = rng.normal(size=5).astype("float32") + S = np.zeros((3, 5), dtype="float32") + Y = np.zeros((3, 5), dtype="float32") + + d = LBFGSDirection(n_parameters=1, memory_size=3)(0, 0.25, g, S, Y) + + assert d.dtype == "float32" + np.testing.assert_allclose(d.eval(), 0.25 * g, rtol=1e-6) + + +@pytest.mark.parametrize( + "props, tensors, message", + [ + ({"n_parameters": 0, "memory_size": 3}, (), "n_parameters must be at least 1"), + ( + {"n_parameters": 1, "memory_size": 0}, + (np.ones(2), np.ones((0, 2)), np.ones((0, 2))), + "memory_size must be at least 1", + ), + ( + {"n_parameters": 1, "memory_size": 3}, + (np.ones(2), np.ones((3, 2)), np.ones((3, 2)), np.ones((3, 2))), + "takes 3 tensors", + ), + ( + {"n_parameters": 1, "memory_size": 3}, + (np.ones(2), np.ones((4, 2)), np.ones((3, 2))), + "memory_size=3", + ), + ( + {"n_parameters": 1, "memory_size": 3}, + (np.ones(2), np.ones((3, 2, 1)), np.ones((3, 2))), + "memory_size=3", + ), + ( + {"n_parameters": 1, "memory_size": 3}, + (np.ones(2, dtype="float32"), np.ones((3, 2)), np.ones((3, 2))), + "dtype", + ), + ], + ids=["no_parameters", "no_memory", "extra_tensor", "wrong_slots", "wrong_rank", "wrong_dtype"], +) +def test_malformed_inputs_are_refused_at_build_time(props, tensors, message): + with pytest.raises(ValueError, match=message): + LBFGSDirection(**props)(0, 1.0, *tensors) From e4ebd9fe25e60bbb16d264ee0d66abed155d9b6f Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 18:12:05 +0800 Subject: [PATCH 03/31] Take the recursion's dot products through BLAS --- pytensor_ml/optim/lbfgs.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/pytensor_ml/optim/lbfgs.py b/pytensor_ml/optim/lbfgs.py index 2b3525a..a462d6c 100644 --- a/pytensor_ml/optim/lbfgs.py +++ b/pytensor_ml/optim/lbfgs.py @@ -154,7 +154,8 @@ def _require_stack_of( def _dot(left: Sequence[TensorVariable], right: Sequence[TensorVariable]) -> TensorVariable: - return pt.sum([pt.sum(a * b) for a, b in zip(left, right)]) + # A dot of two raveled rows reaches BLAS under numba, where a fused multiply-and-sum does not. + return pt.sum([pt.dot(a.ravel(), b.ravel()) for a, b in zip(left, right)]) def _curvatures( From e4893cf333245dea96460b8f490ab2f612de5393 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 18:12:06 +0800 Subject: [PATCH 04/31] Dispatch LBFGSDirection to mlx as a loop over ring rows --- pytensor_ml/dispatch/mlx/__init__.py | 1 + pytensor_ml/dispatch/mlx/lbfgs.py | 50 +++++++++++++++++++++++ tests/dispatch/mlx/test_lbfgs.py | 60 ++++++++++++++++++++++++++++ 3 files changed, 111 insertions(+) create mode 100644 pytensor_ml/dispatch/mlx/lbfgs.py create mode 100644 tests/dispatch/mlx/test_lbfgs.py diff --git a/pytensor_ml/dispatch/mlx/__init__.py b/pytensor_ml/dispatch/mlx/__init__.py index b559c53..113a295 100644 --- a/pytensor_ml/dispatch/mlx/__init__.py +++ b/pytensor_ml/dispatch/mlx/__init__.py @@ -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 diff --git a/pytensor_ml/dispatch/mlx/lbfgs.py b/pytensor_ml/dispatch/mlx/lbfgs.py new file mode 100644 index 0000000..3f04a48 --- /dev/null +++ b/pytensor_ml/dispatch/mlx/lbfgs.py @@ -0,0 +1,50 @@ +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): + return sum(mx.sum(a * b) for a, b in zip(left, right)) + + def direction(count, gamma, *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 = [] + for s, y in zip(s_rows, y_rows): + product = dot(s, y) + curvatures.append( + mx.where(product == 0, 0.0, 1.0 / mx.where(product == 0, 1.0, product)) + ) + + 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] * y_p for v, y_p in zip(vector, y_rows[position])] + vector = [gamma * v for v in vector] + for position in range(m): + beta = curvatures[position] * dot(y_rows[position], vector) + vector = [ + v + (alphas[position] - beta) * s_p for v, s_p in zip(vector, s_rows[position]) + ] + return vector[0] if n == 1 else tuple(vector) + + return direction diff --git a/tests/dispatch/mlx/test_lbfgs.py b/tests/dispatch/mlx/test_lbfgs.py new file mode 100644 index 0000000..2490241 --- /dev/null +++ b/tests/dispatch/mlx/test_lbfgs.py @@ -0,0 +1,60 @@ +import numpy as np +import pytensor +import pytensor.tensor as pt +import pytest + +pytest.importorskip("mlx.core") + +from pytensor_ml.optim.lbfgs import LBFGSDirection +from tests.dispatch.mlx.test_basic import compare_mlx_and_py +from tests.optim.test_lbfgs import ring_stacks, two_loop_direction + +floatX = pytensor.config.floatX + + +@pytest.mark.parametrize("n_pairs, count", [(2, 2), (4, 6)], ids=["not_yet_wrapped", "wrapped"]) +def test_direction_matches_py(n_pairs, count): + rng = np.random.default_rng(sum(map(ord, "MLX LBFGS"))) + shapes = [(3, 2), (4,)] + size = sum(int(np.prod(shape)) for shape in shapes) + memory_size, gamma = 4, 0.7 + gradient = rng.normal(size=size).astype(floatX) + pairs = [] + for _ in range(n_pairs): + s = rng.normal(size=size).astype(floatX) + pairs.append((s, rng.normal(size=size).astype(floatX) + 0.5 * s)) + S, Y = ring_stacks(pairs, memory_size, count, shapes) + splits = np.cumsum([int(np.prod(shape)) for shape in shapes])[:-1] + gradient_pieces = [ + piece.reshape(shape) for piece, shape in zip(np.split(gradient, splits), shapes) + ] + + gradients = [pt.tensor(f"g{i}", shape=shape) for i, shape in enumerate(shapes)] + S_in = [pt.tensor(f"S{i}", shape=(memory_size, *shape)) for i, shape in enumerate(shapes)] + Y_in = [pt.tensor(f"Y{i}", shape=(memory_size, *shape)) for i, shape in enumerate(shapes)] + op = LBFGSDirection(n_parameters=2, memory_size=memory_size) + outputs = op(count, gamma, *gradients, *S_in, *Y_in, return_list=True) + + _, got = compare_mlx_and_py( + [*gradients, *S_in, *Y_in], + outputs, + [*gradient_pieces, *S, *Y], + assert_fn=lambda got, want: np.testing.assert_allclose(got, want, rtol=1e-4), + ) + want = two_loop_direction(gamma, gradient, pairs) + np.testing.assert_allclose( + np.concatenate([np.asarray(d).ravel() for d in got]), want, rtol=1e-4 + ) + + +def test_a_single_parameter_returns_one_array(): + g = np.arange(5, dtype=floatX) + S = np.zeros((3, 5), dtype=floatX) + Y = np.zeros((3, 5), dtype=floatX) + g_in = pt.tensor("g", shape=(5,)) + S_in = pt.tensor("S", shape=(3, 5)) + Y_in = pt.tensor("Y", shape=(3, 5)) + + d = LBFGSDirection(n_parameters=1, memory_size=3)(0, 0.5, g_in, S_in, Y_in) + + compare_mlx_and_py([g_in, S_in, Y_in], d, [g, S, Y]) From 63c5d6f49f8417648e58cfd8ad6fcc53df6142ff Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 19:28:54 +0800 Subject: [PATCH 05/31] Test the direction op against the dense BFGS matrix --- tests/dispatch/mlx/test_lbfgs.py | 4 +-- tests/optim/test_lbfgs.py | 53 +++++++++++++++++++++----------- 2 files changed, 37 insertions(+), 20 deletions(-) diff --git a/tests/dispatch/mlx/test_lbfgs.py b/tests/dispatch/mlx/test_lbfgs.py index 2490241..4527db4 100644 --- a/tests/dispatch/mlx/test_lbfgs.py +++ b/tests/dispatch/mlx/test_lbfgs.py @@ -7,7 +7,7 @@ from pytensor_ml.optim.lbfgs import LBFGSDirection from tests.dispatch.mlx.test_basic import compare_mlx_and_py -from tests.optim.test_lbfgs import ring_stacks, two_loop_direction +from tests.optim.test_lbfgs import dense_inverse_hessian, ring_stacks floatX = pytensor.config.floatX @@ -41,7 +41,7 @@ def test_direction_matches_py(n_pairs, count): [*gradient_pieces, *S, *Y], assert_fn=lambda got, want: np.testing.assert_allclose(got, want, rtol=1e-4), ) - want = two_loop_direction(gamma, gradient, pairs) + want = dense_inverse_hessian(gamma, pairs, size) @ gradient np.testing.assert_allclose( np.concatenate([np.asarray(d).ravel() for d in got]), want, rtol=1e-4 ) diff --git a/tests/optim/test_lbfgs.py b/tests/optim/test_lbfgs.py index db7252c..72be32c 100644 --- a/tests/optim/test_lbfgs.py +++ b/tests/optim/test_lbfgs.py @@ -10,20 +10,17 @@ RTOL = 1e-6 if floatX == "float64" else 1e-4 -def two_loop_direction(gamma, gradient, pairs): - """Nocedal and Wright 7.4 on one flat vector, over ``(s, y)`` pairs given oldest first.""" - q = gradient.astype(np.float64) - alphas = [] - for s, y in reversed(pairs): +def dense_inverse_hessian(gamma, pairs, size): + """The matrix the two-loop recursion multiplies by, built from its definition: BFGS updates from + ``gamma I`` over ``(s, y)`` pairs oldest first, ``H <- V^T H V + rho s s^T`` with ``V = I - rho y s^T`` + (Nocedal and Wright, equation 7.16).""" + H = gamma * np.eye(size) + for s, y in pairs: + s, y = s.astype(np.float64), y.astype(np.float64) rho = 1.0 / (y @ s) - alphas.append(rho * (s @ q)) - q = q - alphas[-1] * y - r = gamma * q - for (s, y), alpha in zip(pairs, reversed(alphas)): - rho = 1.0 / (y @ s) - beta = rho * (y @ r) - r = r + (alpha - beta) * s - return r + V = np.eye(size) - rho * np.outer(y, s) + H = V.T @ H @ V + rho * np.outer(s, s) + return H def ring_stacks(pairs, memory_size, count, shapes): @@ -42,9 +39,10 @@ def ring_stacks(pairs, memory_size, count, shapes): @pytest.mark.parametrize("n_pairs, count", [(2, 2), (4, 6)], ids=["not_yet_wrapped", "wrapped"]) def test_direction_matches_the_two_loop_recursion_over_a_ring(n_pairs, count): - # The reference sees a flat vector and a chronological list, so it shares no ring or reshape - # arithmetic with the op. Two parameters of different rank exercise the cross-parameter dot products. - # Before the ring wraps its empty slots lead the order; after, the newest pair sits mid-ring. + # The reference is the dense matrix the recursion is an algorithm for, built from the textbook update + # on a flat vector, so it shares neither the loop nor the ring or reshape arithmetic with the op. Two + # parameters of different rank exercise the cross-parameter dot products. Before the ring wraps its + # empty slots lead the order; after, the newest pair sits mid-ring. rng = np.random.default_rng(0) shapes = [(3, 2), (4,)] size = sum(int(np.prod(shape)) for shape in shapes) @@ -53,7 +51,8 @@ def test_direction_matches_the_two_loop_recursion_over_a_ring(n_pairs, count): pairs = [] for _ in range(n_pairs): s = rng.normal(size=size).astype(floatX) - pairs.append((s, rng.normal(size=size).astype(floatX) + 0.5 * s)) # keeps y . s > 0 + noise = rng.normal(size=size).astype(floatX) + pairs.append((s, noise - (noise @ s) / (s @ s) * s + 0.5 * s)) # y . s = 0.5 s . s > 0 S, Y = ring_stacks(pairs, memory_size, count, shapes) op = LBFGSDirection(n_parameters=2, memory_size=memory_size) @@ -69,7 +68,25 @@ def test_direction_matches_the_two_loop_recursion_over_a_ring(n_pairs, count): piece.reshape(shape) for piece, shape in zip(np.split(gradient, splits), shapes) ] got = np.concatenate([d.ravel() for d in direction(*gradient_pieces, *S, *Y)]) - np.testing.assert_allclose(got, two_loop_direction(gamma, gradient, pairs), rtol=RTOL) + want = dense_inverse_hessian(gamma, pairs, size) @ gradient + np.testing.assert_allclose(got, want, rtol=RTOL) + + +def test_a_scalar_parameter_has_vector_stacks(): + g = pt.scalar("g", dtype=floatX) + S = pt.vector("S", dtype=floatX) + Y = pt.vector("Y", dtype=floatX) + + d = LBFGSDirection(n_parameters=1, memory_size=3)(1, 1.0, g, S, Y) + + # One pair (s, y) with y = 2 s: H y = s, so H maps g onto g / 2. + np.testing.assert_allclose( + d.eval( + {g: 4.0, S: np.array([0, 0, 1.5], dtype=floatX), Y: np.array([0, 0, 3.0], dtype=floatX)} + ), + 2.0, + rtol=RTOL, + ) def test_an_empty_memory_scales_the_gradient(): From 959b00489aa77a724b9f52228e2fdf56ecdec3fe Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 19:29:03 +0800 Subject: [PATCH 06/31] Keep each parameter's dtype through the recursion --- pytensor_ml/optim/lbfgs.py | 6 +++--- tests/optim/test_lbfgs.py | 15 +++++++++++++++ 2 files changed, 18 insertions(+), 3 deletions(-) diff --git a/pytensor_ml/optim/lbfgs.py b/pytensor_ml/optim/lbfgs.py index a462d6c..a0db27a 100644 --- a/pytensor_ml/optim/lbfgs.py +++ b/pytensor_ml/optim/lbfgs.py @@ -107,7 +107,7 @@ def right_product(slot, *vector): s = [stack[slot] for stack in S] y = [stack[slot] for stack in Y] alpha = curvatures[slot] * _dot(s, vector) - return [v - alpha * y_p for v, y_p in zip(vector, y)] + [alpha] + return [v - alpha.astype(v.dtype) * y_p for v, y_p in zip(vector, y)] + [alpha] *q, alphas = pytensor.scan( right_product, @@ -116,13 +116,13 @@ def right_product(slot, *vector): go_backwards=True, return_updates=False, ) - r = [gamma * v[-1] for v in q] + r = [gamma.astype(v.dtype) * v[-1] for v in q] def left_product(slot, alpha, *vector): s = [stack[slot] for stack in S] y = [stack[slot] for stack in Y] beta = curvatures[slot] * _dot(y, vector) - return [v + (alpha - beta) * s_p for v, s_p in zip(vector, s)] + return [v + (alpha - beta).astype(v.dtype) * s_p for v, s_p in zip(vector, s)] # The backward loop reports its alphas newest first and the forward loop reads them oldest first. r = pytensor.scan( diff --git a/tests/optim/test_lbfgs.py b/tests/optim/test_lbfgs.py index 72be32c..36772b4 100644 --- a/tests/optim/test_lbfgs.py +++ b/tests/optim/test_lbfgs.py @@ -72,6 +72,21 @@ def test_direction_matches_the_two_loop_recursion_over_a_ring(n_pairs, count): np.testing.assert_allclose(got, want, rtol=RTOL) +def test_parameters_of_different_dtypes_keep_their_own(): + # The cross-parameter dot products upcast to the widest dtype; each carried vector has to be cast + # back or the scan refuses the narrower parameter's recurrence. + g_wide = pt.tensor("g_wide", shape=(3,), dtype="float64") + g_narrow = pt.tensor("g_narrow", shape=(2,), dtype="float32") + S_wide, Y_wide = (pt.tensor(name, shape=(2, 3), dtype="float64") for name in "SY") + S_narrow, Y_narrow = (pt.tensor(name, shape=(2, 2), dtype="float32") for name in ("s", "y")) + + wide, narrow = LBFGSDirection(n_parameters=2, memory_size=2)( + 1, 0.5, g_wide, g_narrow, S_wide, S_narrow, Y_wide, Y_narrow, return_list=True + ) + + assert (wide.dtype, narrow.dtype) == ("float64", "float32") + + def test_a_scalar_parameter_has_vector_stacks(): g = pt.scalar("g", dtype=floatX) S = pt.vector("S", dtype=floatX) From 75af3370b0ffc1b4baa7a0a98351ac7a975ad5f6 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 19:29:15 +0800 Subject: [PATCH 07/31] Compute the curvatures with the dot the rule's guard uses --- pytensor_ml/optim/lbfgs.py | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/pytensor_ml/optim/lbfgs.py b/pytensor_ml/optim/lbfgs.py index a0db27a..7ee3c35 100644 --- a/pytensor_ml/optim/lbfgs.py +++ b/pytensor_ml/optim/lbfgs.py @@ -106,7 +106,7 @@ def build_inner_graph(self, *inputs: TensorVariable) -> list[Variable]: def right_product(slot, *vector): s = [stack[slot] for stack in S] y = [stack[slot] for stack in Y] - alpha = curvatures[slot] * _dot(s, vector) + alpha = curvatures[slot] * flat_dot(s, vector) return [v - alpha.astype(v.dtype) * y_p for v, y_p in zip(vector, y)] + [alpha] *q, alphas = pytensor.scan( @@ -121,7 +121,7 @@ def right_product(slot, *vector): def left_product(slot, alpha, *vector): s = [stack[slot] for stack in S] y = [stack[slot] for stack in Y] - beta = curvatures[slot] * _dot(y, vector) + beta = curvatures[slot] * flat_dot(y, vector) return [v + (alpha - beta).astype(v.dtype) * s_p for v, s_p in zip(vector, s)] # The backward loop reports its alphas newest first and the forward loop reads them oldest first. @@ -153,8 +153,8 @@ def _require_stack_of( ) -def _dot(left: Sequence[TensorVariable], right: Sequence[TensorVariable]) -> TensorVariable: - # A dot of two raveled rows reaches BLAS under numba, where a fused multiply-and-sum does not. +def flat_dot(left: Sequence[TensorVariable], right: Sequence[TensorVariable]) -> TensorVariable: + """Dot product of two lists of tensors read as one flat vector each, through BLAS under numba.""" return pt.sum([pt.dot(a.ravel(), b.ravel()) for a, b in zip(left, right)]) @@ -162,8 +162,10 @@ def _curvatures( S: Sequence[TensorVariable], Y: Sequence[TensorVariable], memory_size: int ) -> TensorVariable: """Return ``1 / (y_i . s_i)`` per slot, and zero for an empty slot rather than a division by zero.""" - products = pt.sum( - [pt.sum((s * y).reshape((memory_size, -1)), axis=1) for s, y in zip(S, Y)], axis=0 + # The same dot the writer's admission test uses, so a pair it admitted never rounds to a negative + # curvature here. + products = pt.stack( + [flat_dot([s[slot] for s in S], [y[slot] for y in Y]) for slot in range(memory_size)] ) empty = pt.eq(products, 0.0) return pt.switch(empty, 0.0, 1.0 / pt.switch(empty, 1.0, products)) From 5e865426c31c90c8f8dd84aad1dfb397ce4a2441 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 19:29:17 +0800 Subject: [PATCH 08/31] Allocate optimizer state at the parameter's declared dtype --- pytensor_ml/optim/base.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/pytensor_ml/optim/base.py b/pytensor_ml/optim/base.py index b6125f1..034b14d 100644 --- a/pytensor_ml/optim/base.py +++ b/pytensor_ml/optim/base.py @@ -538,8 +538,10 @@ def state_for( value = parameter.get_value(borrow=True) 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=value.dtype), + np.full(shape, fill_value, dtype=parameter.type.dtype), name=f"{parameter.name}/{slot}", shape=static_shape, ) From fceed2aad9e4087b1eef206379dbec34cea5ac91 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 19:30:54 +0800 Subject: [PATCH 09/31] Add lbfgs_updates with a ring memory and curvature guard --- docs/source/api/optim.rst | 1 + pytensor_ml/optim/__init__.py | 2 + pytensor_ml/optim/base.py | 4 +- pytensor_ml/optim/rules.py | 153 +++++++++++++++++++++++++++++++ tests/dispatch/mlx/test_lbfgs.py | 31 +++++++ tests/optim/test_rules.py | 137 +++++++++++++++++++++++++++ 6 files changed, 327 insertions(+), 1 deletion(-) diff --git a/docs/source/api/optim.rst b/docs/source/api/optim.rst index 632a92c..6c90791 100644 --- a/docs/source/api/optim.rst +++ b/docs/source/api/optim.rst @@ -130,3 +130,4 @@ Low-level update functions rprop_updates adagrad_updates adadelta_updates + lbfgs_updates diff --git a/pytensor_ml/optim/__init__.py b/pytensor_ml/optim/__init__.py index a86eeb0..daf14e9 100644 --- a/pytensor_ml/optim/__init__.py +++ b/pytensor_ml/optim/__init__.py @@ -41,6 +41,7 @@ adam_updates, adamax_updates, adamw_updates, + lbfgs_updates, nadam_updates, rmsprop_updates, rprop_updates, @@ -96,6 +97,7 @@ "get_gradients", "join_schedules", "large_step", + "lbfgs_updates", "linear_onecycle_schedule", "linear_schedule", "nadam", diff --git a/pytensor_ml/optim/base.py b/pytensor_ml/optim/base.py index 034b14d..2febacc 100644 --- a/pytensor_ml/optim/base.py +++ b/pytensor_ml/optim/base.py @@ -537,7 +537,9 @@ def state_for( value = parameter.get_value(borrow=True) 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) + 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( diff --git a/pytensor_ml/optim/rules.py b/pytensor_ml/optim/rules.py index c6a6b60..bdb778f 100644 --- a/pytensor_ml/optim/rules.py +++ b/pytensor_ml/optim/rules.py @@ -1,5 +1,6 @@ from collections.abc import Callable, Sequence +import numpy as np import pytensor.tensor as pt from pytensor import config @@ -11,14 +12,17 @@ LearningRate, LossGradientsOrUpdates, Parameter, + Rate, Steps, Updates, gradients_to_descend, rate_on, read_rate, + scalar_state, state_for, to_floatx, ) +from pytensor_ml.optim.lbfgs import LBFGSDirection, flat_dot from pytensor_ml.params import step_counter @@ -900,3 +904,152 @@ def rprop_updates( updates[parameter] = parameter - pt.sign(effective_gradient) * new_step_size return updates + + +def lbfgs_updates( + loss_gradients_or_updates: LossGradientsOrUpdates, + parameters: Sequence[Parameter], + learning_rate: Rate = 1.0, + memory_size: int = 10, + scale_init_precond: bool = True, + namespace: str = "lbfgs", +) -> Updates: + r""" + L-BFGS: descend along the gradient multiplied by a limited-memory inverse-Hessian approximation. + + The approximation is built from the last ``memory_size`` accepted pairs of parameter differences + :math:`s = p_{k+1} - p_k` and gradient differences :math:`y = g_{k+1} - g_k`, applied to the gradient + by the two-loop recursion of :class:`~pytensor_ml.optim.lbfgs.LBFGSDirection` starting from + :math:`\gamma I`, with :math:`\gamma = s^\top y / y^\top y` for the newest pair. A pair enters the + memory only when :math:`y^\top s > \epsilon\, y^\top y`, which keeps the approximation positive + definite, so a step through a non-convex region leaves the memory as it was. Before any pair is + accepted :math:`\gamma = \min(1, 1 / \|g\|)`, which keeps the first step inside the unit ball. The + step is :math:`p \leftarrow p - \eta H g`. + + The direction is well scaled once the memory holds a pair, so :math:`\eta = 1` is the natural rate + and a line search the natural way to back off from it. Consecutive gradients have to be measured on + the same objective for their difference to be curvature, so the rule assumes a deterministic, + full-batch loss. + + Parameters + ---------- + loss_gradients_or_updates : TensorVariable, sequence of TensorVariable, or Updates + Scalar loss to differentiate, precomputed gradients, or the updates dict an earlier transform in + a chain produced. + parameters : sequence of shared tensor variable + Parameters to update. + learning_rate : float or shared tensor variable + Step size :math:`\eta`. Default 1.0. + memory_size : int + Number of pairs the memory holds. Default 10. + scale_init_precond : bool + Start the recursion from :math:`\gamma I` as above. When False it starts from the identity, and + the first step is the raw gradient. Default True. + + namespace : str + Prefix for every state slot this rule allocates, so two rules in one graph keep separate state + rather than reusing each other's. Default is the rule's own name. + + Returns + ------- + updates : Updates + Mapping from each parameter and its memory buffers to their next values. + + Examples + -------- + Compile the step yourself rather than going through :func:`~pytensor_ml.optim.train.compile_train`. + The rule returns the updates dict directly, at a fixed rate and with no line search: + + .. 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 lbfgs_updates + from pytensor_ml.pytensorf import collect_trainable_params, function + + X = Input("X", shape=(None, 4)) + loss, target = supervised_loss(Linear("fc", n_in=4, n_out=1)(X), SquaredError()) + + updates = lbfgs_updates(loss, collect_trainable_params(loss), learning_rate=0.5) + step = function([X, target], loss, updates=updates) + loss_value = step(np.zeros((8, 4)), np.zeros((8, 1))) + """ + if memory_size < 1: + raise ValueError(f"memory_size must be at least 1, got {memory_size}.") + + incoming, gradients = gradients_to_descend(loss_gradients_or_updates, parameters, namespace) + learning_rate = to_floatx(learning_rate) + + step_count = step_counter(f"{namespace}/step_count") + pairs_written = scalar_state(f"{namespace}/pairs_written", dtype="int64") + previous_values = [state_for(p, f"{namespace}/previous_value") for p in parameters] + previous_gradients = [state_for(p, f"{namespace}/previous_gradient") for p in parameters] + value_memory = [ + state_for(p, f"{namespace}/value_differences", history_size=memory_size) for p in parameters + ] + gradient_memory = [ + state_for(p, f"{namespace}/gradient_differences", history_size=memory_size) + for p in parameters + ] + + # The buffers hold zeros before the first step, so the differences read off them are meaningless + # until a previous point exists; the guard below never lets those into the memory. + value_differences = [p - previous for p, previous in zip(parameters, previous_values)] + gradient_differences = [g - previous for g, previous in zip(gradients, previous_gradients)] + curvature = flat_dot(gradient_differences, value_differences) + gradient_change = flat_dot(gradient_differences, gradient_differences) + epsilon = max(np.finfo(gradient.dtype).eps for gradient in gradients) + accept = (step_count > 0) & (curvature > epsilon * gradient_change) + + # Rejection rewrites the slot with itself, so the write stays in place and unconditional; only the + # count decides whether the slot is now part of the memory. The slot is a one-element index vector + # rather than a scalar because mlx cannot trace a scalar index (pymc-devs/pytensor#2422). + slot = (pairs_written % memory_size)[None] + new_value_memory = [ + pt.set_subtensor(memory[slot], pt.switch(accept, s[None], memory[slot])) + for memory, s in zip(value_memory, value_differences) + ] + new_gradient_memory = [ + pt.set_subtensor(memory[slot], pt.switch(accept, y[None], memory[slot])) + for memory, y in zip(gradient_memory, gradient_differences) + ] + new_pairs_written = pairs_written + accept.astype(pairs_written.dtype) + + if scale_init_precond: + has_pair = new_pairs_written > 0 + newest = ((new_pairs_written - 1) % memory_size)[None] + newest_s = [memory[newest] for memory in new_value_memory] + newest_y = [memory[newest] for memory in new_gradient_memory] + newest_curvature = flat_dot(newest_s, newest_y) + newest_change = pt.switch(has_pair, flat_dot(newest_y, newest_y), 1.0) + gradient_norm = pt.sqrt(flat_dot(gradients, gradients)) + identity_scale = pt.switch( + has_pair, + newest_curvature / newest_change, + pt.minimum(1.0, 1.0 / pt.switch(gradient_norm > 0, gradient_norm, 1.0)), + ) + else: + identity_scale = 1.0 + + directions = LBFGSDirection(n_parameters=len(parameters), memory_size=memory_size)( + new_pairs_written, + identity_scale, + *gradients, + *new_value_memory, + *new_gradient_memory, + return_list=True, + ) + + updates: Updates = Steps(incoming) + updates[step_count] = step_count + 1 + updates[pairs_written] = new_pairs_written + for index, parameter in enumerate(parameters): + updates[previous_values[index]] = parameter + updates[previous_gradients[index]] = gradients[index] + updates[value_memory[index]] = new_value_memory[index] + updates[gradient_memory[index]] = new_gradient_memory[index] + updates[parameter] = parameter - learning_rate * directions[index] + + return updates diff --git a/tests/dispatch/mlx/test_lbfgs.py b/tests/dispatch/mlx/test_lbfgs.py index 4527db4..ecb7a90 100644 --- a/tests/dispatch/mlx/test_lbfgs.py +++ b/tests/dispatch/mlx/test_lbfgs.py @@ -5,7 +5,13 @@ pytest.importorskip("mlx.core") +from pytensor.compile.mode import Mode +from pytensor.link.mlx.linker import MLXLinker + +from pytensor_ml.optim import lbfgs_updates from pytensor_ml.optim.lbfgs import LBFGSDirection +from pytensor_ml.params import trainable +from pytensor_ml.pytensorf import function from tests.dispatch.mlx.test_basic import compare_mlx_and_py from tests.optim.test_lbfgs import dense_inverse_hessian, ring_stacks @@ -58,3 +64,28 @@ def test_a_single_parameter_returns_one_array(): d = LBFGSDirection(n_parameters=1, memory_size=3)(0, 0.5, g_in, S_in, Y_in) compare_mlx_and_py([g_in, S_in, Y_in], d, [g, S, Y]) + + +@pytest.mark.parametrize("use_compile", [True, False], ids=["compiled", "eager"]) +def test_the_rule_reaches_the_minimum_of_a_quadratic(use_compile): + # The rule reads and writes its ring with a traced slot, which mlx traces only as advanced indexing + # (pymc-devs/pytensor#2422); this is the end-to-end check that the whole step compiles and runs. + A = np.array([[3.0, 0.5], [0.5, 1.0]]) + b = np.array([1.0, -2.0]) + u = trainable(np.array([5.0], dtype=floatX), name="u") + v = trainable(np.array([-3.0], dtype=floatX), name="v") + x = pt.concatenate([u, v]) + loss = 0.5 * x @ pt.constant(A, dtype=floatX) @ x - pt.constant(b, dtype=floatX) @ x + mode = Mode(linker=MLXLinker(use_compile=use_compile), optimizer="fast_run") + step = function( + [], loss, updates=lbfgs_updates(loss, [u, v], learning_rate=1.0, memory_size=2), mode=mode + ) + + for _ in range(12): + step() + + np.testing.assert_allclose( + np.concatenate([np.asarray(u.get_value()), np.asarray(v.get_value())]), + np.linalg.solve(A, b), + rtol=1e-4, + ) diff --git a/tests/optim/test_rules.py b/tests/optim/test_rules.py index 9ff5c00..f3a8107 100644 --- a/tests/optim/test_rules.py +++ b/tests/optim/test_rules.py @@ -21,6 +21,7 @@ adamw_updates, compile_train, cosine_schedule, + lbfgs_updates, nadam, nadam_updates, rmsprop, @@ -32,6 +33,7 @@ ) from pytensor_ml.optim import alias as alias_module from pytensor_ml.pytensorf import function +from tests.optim.test_lbfgs import dense_inverse_hessian floatX = pytensor.config.floatX @@ -260,6 +262,7 @@ def test_two_rules_of_one_kind_keep_separate_state_when_named(): (rmsprop_updates, "rmsprop"), (adadelta_updates, "adadelta"), (rprop_updates, "rprop"), + (lbfgs_updates, "lbfgs"), ], ids=lambda value: value if isinstance(value, str) else "", ) @@ -550,6 +553,140 @@ def test_rprop_shrinks_and_skips_on_sign_flip(): np.testing.assert_allclose(p.get_value(), [-lr + lr * eta_minus]) +def test_lbfgs_satisfies_the_secant_condition_on_the_newest_pair(): + """The inverse-Hessian estimate maps the newest gradient difference onto the parameter difference + that produced it, ``H y = s``, whatever the initial scaling. A zero gradient holds the parameters + still while the move before it becomes the newest pair with ``y = -g``, so feeding ``-g`` next has + to move them by ``-lr * s``. Two parameters, so the memory is split across tensors.""" + g_u, g_v = pt.vector("g_u"), pt.vector("g_v") + u = trainable(np.zeros(2), name="u") + v = trainable(np.zeros(1), name="v") + lr = 0.3 + fn = function([g_u, g_v], [u, v], updates=lbfgs_updates([g_u, g_v], [u, v], learning_rate=lr)) + g = [np.array([1.0, -2.0], dtype=floatX), np.array([0.5], dtype=floatX)] + + fn(*g) + fn(*[0.5 * gp for gp in g]) + before_move = [x.get_value().copy() for x in (u, v)] + fn(*[0.5 * gp for gp in g]) + after_move = [x.get_value().copy() for x in (u, v)] + fn(*[np.zeros_like(gp) for gp in g]) # no move; (after - before, -0.5 g) is now the newest pair + fn(*[-0.5 * gp for gp in g]) + + for x, x_after_move, x_before_move in zip((u, v), after_move, before_move): + np.testing.assert_allclose( + x.get_value(), x_after_move - lr * (x_after_move - x_before_move), rtol=RTOL + ) + + +def test_lbfgs_reaches_the_minimum_of_a_quadratic(): + # Two slots for two dimensions: once both hold pairs the estimate is close to the true inverse + # Hessian and unit steps close in on the minimizer, which is known in closed form. + A = np.array([[3.0, 0.5], [0.5, 1.0]]) + b = np.array([1.0, -2.0]) + u = trainable(np.array([5.0]), name="u") + v = trainable(np.array([-3.0]), name="v") + x = pt.concatenate([u, v]) + loss = 0.5 * x @ pt.constant(A, dtype=floatX) @ x - pt.constant(b, dtype=floatX) @ x + step = function([], loss, updates=lbfgs_updates(loss, [u, v], learning_rate=1.0, memory_size=2)) + + for _ in range(12): + step() + + np.testing.assert_allclose( + np.concatenate([u.get_value(), v.get_value()]), np.linalg.solve(A, b), rtol=1e-4 + ) + + +@pytest.mark.parametrize("gradient", [[3.0, -4.0], [0.3, -0.4]], ids=["long", "short"]) +def test_lbfgs_first_step_is_the_gradient_capped_to_the_unit_ball(gradient): + # A gradient of norm 5 is cut to unit length, one of norm 0.5 is left as it is. + p = trainable(np.zeros(2), name="w") + loss = (pt.constant(np.array(gradient), dtype=floatX) * p).sum() + step = function([], loss, updates=lbfgs_updates(loss, [p], learning_rate=1.0)) + + step() + + g = np.array(gradient) + np.testing.assert_allclose(p.get_value(), -min(1.0, 1.0 / np.linalg.norm(g)) * g, rtol=RTOL) + + +def test_lbfgs_step_matches_the_dense_update_through_a_ring_wrap(): + """On a strictly convex quadratic every pair is accepted, so the memory is the last ``memory_size`` + chronological pairs and each step is ``-lr * H g`` for the dense BFGS matrix built from them. Two + slots over six steps wrap the ring twice; a rule that overwrote the wrong slot or read the newest + pair off by one would drift from the dense reference from the third step on.""" + A = np.diag([1.0, 2.0, 3.0, 4.0, 5.0]) + 0.1 + A = A @ A.T + b = np.array([0.3, -1.0, 2.0, 0.5, -0.7]) + u = trainable(np.array([1.0, -2.0, 0.5]), name="u") + v = trainable(np.array([3.0, 1.0]), name="v") + x = pt.concatenate([u, v]) + loss = 0.5 * x @ pt.constant(A, dtype=floatX) @ x - pt.constant(b, dtype=floatX) @ x + memory_size, lr = 2, 0.5 + updates = lbfgs_updates(loss, [u, v], learning_rate=lr, memory_size=memory_size) + step = function([], pt.grad(loss, [u, v]), updates=updates) # gradient before the update + + pairs = [] + previous = None + for _ in range(6): + x_before = np.concatenate([u.get_value(), v.get_value()]) + g_before = np.concatenate([g.ravel() for g in step()]) + if previous is not None: + pairs.append((x_before - previous[0], g_before - previous[1])) + if pairs: + s, y = pairs[-1] + gamma = (s @ y) / (y @ y) + else: + gamma = min(1.0, 1.0 / np.linalg.norm(g_before)) + H = dense_inverse_hessian(gamma, pairs[-memory_size:], x_before.size) + np.testing.assert_allclose( + np.concatenate([u.get_value(), v.get_value()]), x_before - lr * H @ g_before, rtol=RTOL + ) + previous = (x_before, g_before) + + +def test_lbfgs_without_initial_scaling_starts_along_the_raw_gradient(): + p = trainable(np.array([3.0, -4.0]), name="w") + loss = (pt.constant(np.array([3.0, -4.0]), dtype=floatX) * p).sum() + step = function( + [], loss, updates=lbfgs_updates(loss, [p], learning_rate=0.1, scale_init_precond=False) + ) + + step() + + np.testing.assert_allclose(p.get_value(), [3.0, -4.0] - 0.1 * np.array([3.0, -4.0]), rtol=RTOL) + + +def test_lbfgs_rejects_a_pair_with_negative_curvature(): + """A step whose gradient change opposes the parameter change would make the inverse-Hessian estimate + indefinite, so the pair is left out of the memory, the ring index does not advance, and the next step + is the one an empty memory gives.""" + g = pt.vector("g") + p = trainable(np.zeros(2), name="w") + lr = 0.1 + updates = lbfgs_updates([g], [p], learning_rate=lr, memory_size=2) + memory = next(key for key in updates if key.name == "w/lbfgs/value_differences") + pairs_written = next(key for key in updates if key.name == "lbfgs/pairs_written") + fn = function([g], p, updates=updates) + + fn(np.array([1.0, 0.0], dtype=floatX)) # first step: no previous point, nothing to write + before = p.get_value().copy() + fn(np.array([2.0, 0.0], dtype=floatX)) # p moved along -g and g grew: y . s < 0, rejected + assert int(pairs_written.get_value()) == 0 + np.testing.assert_array_equal(memory.get_value(), 0.0) + np.testing.assert_allclose(p.get_value(), before - lr * 0.5 * np.array([2.0, 0.0]), rtol=RTOL) + fn(np.array([0.5, 0.0], dtype=floatX)) # g shrank along the move: y . s > 0, accepted + assert int(pairs_written.get_value()) == 1 + assert np.any(memory.get_value()[0] != 0.0) + + +def test_lbfgs_rejects_a_zero_memory_size(): + p = trainable(np.zeros(2), name="w") + with pytest.raises(ValueError, match="memory_size must be at least 1"): + lbfgs_updates((p**2).sum(), [p], memory_size=0) + + def test_amsgrad_caps_step_after_gradient_spike(): """AMSGrad divides by the running maximum of the second moment, so a large gradient permanently caps the denominator. Once gradients shrink it therefore takes a smaller step than plain Adam, whose decaying From 8b5a8dec4abf0a9a43e470f1fe64a839d12aeb32 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 19:31:04 +0800 Subject: [PATCH 10/31] Raise the mlx floor to 0.32.1 and dot through its vector matmul --- .github/workflows/run_tests.yml | 2 +- pytensor_ml/dispatch/mlx/lbfgs.py | 4 +++- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/.github/workflows/run_tests.yml b/.github/workflows/run_tests.yml index 829e952..3144111 100644 --- a/.github/workflows/run_tests.yml +++ b/.github/workflows/run_tests.yml @@ -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.1,<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 }} diff --git a/pytensor_ml/dispatch/mlx/lbfgs.py b/pytensor_ml/dispatch/mlx/lbfgs.py index 3f04a48..e1c505f 100644 --- a/pytensor_ml/dispatch/mlx/lbfgs.py +++ b/pytensor_ml/dispatch/mlx/lbfgs.py @@ -16,7 +16,9 @@ def rows(stacks, slot): return [mx.take(stack, slot, axis=0) for stack in stacks] def dot(left, right): - return sum(mx.sum(a * b) for a, b in zip(left, right)) + # Vector matmul is the fastest dot mlx has from 0.32.1 (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, *tensors): gradients = tensors[:n] From 13eba4590fda990529e3d51ca360482c8b7c947f Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 20:33:06 +0800 Subject: [PATCH 11/31] Set the mlx floor to 0.32.2 --- .github/workflows/run_tests.yml | 2 +- pytensor_ml/dispatch/mlx/lbfgs.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/run_tests.yml b/.github/workflows/run_tests.yml index 3144111..76e1d24 100644 --- a/.github/workflows/run_tests.yml +++ b/.github/workflows/run_tests.yml @@ -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.32.1,<0.33"; 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 }} diff --git a/pytensor_ml/dispatch/mlx/lbfgs.py b/pytensor_ml/dispatch/mlx/lbfgs.py index e1c505f..25ce3ac 100644 --- a/pytensor_ml/dispatch/mlx/lbfgs.py +++ b/pytensor_ml/dispatch/mlx/lbfgs.py @@ -16,7 +16,7 @@ def rows(stacks, slot): 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.1 (ml-explore/mlx#3580); before that it ran + # 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)) From 05d2ffb696d20d79e312d4d27e9c968647774997 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 20:55:47 +0800 Subject: [PATCH 12/31] Add the lbfgs alias --- docs/source/api/optim.rst | 1 + pytensor_ml/optim/__init__.py | 2 ++ pytensor_ml/optim/alias.py | 50 +++++++++++++++++++++++++++++++++++ pytensor_ml/optim/rules.py | 9 +++---- tests/optim/test_rules.py | 20 +++++++++++++- 5 files changed, 76 insertions(+), 6 deletions(-) diff --git a/docs/source/api/optim.rst b/docs/source/api/optim.rst index 6c90791..6de1b94 100644 --- a/docs/source/api/optim.rst +++ b/docs/source/api/optim.rst @@ -26,6 +26,7 @@ Update rules rprop adagrad adadelta + lbfgs Transforms ---------- diff --git a/pytensor_ml/optim/__init__.py b/pytensor_ml/optim/__init__.py index daf14e9..f064e85 100644 --- a/pytensor_ml/optim/__init__.py +++ b/pytensor_ml/optim/__init__.py @@ -4,6 +4,7 @@ adam, adamax, adamw, + lbfgs, nadam, rmsprop, rprop, @@ -97,6 +98,7 @@ "get_gradients", "join_schedules", "large_step", + "lbfgs", "lbfgs_updates", "linear_onecycle_schedule", "linear_schedule", diff --git a/pytensor_ml/optim/alias.py b/pytensor_ml/optim/alias.py index 8a860ea..a1c6208 100644 --- a/pytensor_ml/optim/alias.py +++ b/pytensor_ml/optim/alias.py @@ -14,6 +14,7 @@ adam_updates, adamax_updates, adamw_updates, + lbfgs_updates, nadam_updates, rmsprop_updates, rprop_updates, @@ -357,6 +358,55 @@ 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. 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. + + .. 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, diff --git a/pytensor_ml/optim/rules.py b/pytensor_ml/optim/rules.py index bdb778f..2a230eb 100644 --- a/pytensor_ml/optim/rules.py +++ b/pytensor_ml/optim/rules.py @@ -12,7 +12,6 @@ LearningRate, LossGradientsOrUpdates, Parameter, - Rate, Steps, Updates, gradients_to_descend, @@ -909,7 +908,7 @@ def rprop_updates( def lbfgs_updates( loss_gradients_or_updates: LossGradientsOrUpdates, parameters: Sequence[Parameter], - learning_rate: Rate = 1.0, + learning_rate: LearningRate = 1.0, memory_size: int = 10, scale_init_precond: bool = True, namespace: str = "lbfgs", @@ -958,7 +957,7 @@ def lbfgs_updates( Examples -------- Compile the step yourself rather than going through :func:`~pytensor_ml.optim.train.compile_train`. - The rule returns the updates dict directly, at a fixed rate and with no line search: + The rule returns the updates dict directly, with no line search: .. code-block:: python @@ -980,9 +979,9 @@ def lbfgs_updates( raise ValueError(f"memory_size must be at least 1, got {memory_size}.") incoming, gradients = gradients_to_descend(loss_gradients_or_updates, parameters, namespace) - learning_rate = to_floatx(learning_rate) - step_count = step_counter(f"{namespace}/step_count") + learning_rate = to_floatx(rate_on(learning_rate, step_count)) + pairs_written = scalar_state(f"{namespace}/pairs_written", dtype="int64") previous_values = [state_for(p, f"{namespace}/previous_value") for p in parameters] previous_gradients = [state_for(p, f"{namespace}/previous_gradient") for p in parameters] diff --git a/tests/optim/test_rules.py b/tests/optim/test_rules.py index f3a8107..904506e 100644 --- a/tests/optim/test_rules.py +++ b/tests/optim/test_rules.py @@ -21,6 +21,7 @@ adamw_updates, compile_train, cosine_schedule, + lbfgs, lbfgs_updates, nadam, nadam_updates, @@ -64,6 +65,7 @@ def trainable(value, name=None, **kwargs): nadam(learning_rate=1e-2), adamax(learning_rate=1e-2), rprop(learning_rate=1e-2), + lbfgs(learning_rate=1e-2), ], ids=[ "sgd", @@ -81,6 +83,7 @@ def trainable(value, name=None, **kwargs): "nadam", "adamax", "rprop", + "lbfgs", ], ) def test_rule_reduces_loss(run_training, rule): @@ -99,8 +102,9 @@ def test_rule_reduces_loss(run_training, rule): (rmsprop, "rmsprop_updates"), (adagrad, "adagrad_updates"), (adadelta, "adadelta_updates"), + (lbfgs, "lbfgs_updates"), ], - ids=["adam", "adamw", "nadam", "adamax", "rprop", "rmsprop", "adagrad", "adadelta"], + ids=["adam", "adamw", "nadam", "adamax", "rprop", "rmsprop", "adagrad", "adadelta", "lbfgs"], ) def test_alias_forwards_every_argument_to_the_matching_parameter(alias, updates_name, monkeypatch): # test_rule_reduces_loss cannot see a mis-forward: the loss still falls if beta1 and beta2 are @@ -646,6 +650,20 @@ def test_lbfgs_step_matches_the_dense_update_through_a_ring_wrap(): previous = (x_before, g_before) +def test_lbfgs_schedule_reads_the_rules_own_clock(): + # The rule keeps a step counter to tell the first step apart; a scheduled rate must read that same + # clock rather than allocate a second one measuring the same time. + parameter = trainable(np.array([1.0, -2.0]), name="w") + loss = (parameter**2).sum() + + step = compile_train(loss, lbfgs(cosine_schedule(0.1, 10), memory_size=2), inputs=[]) + + counters = [ + str(shared.name) for shared in step.get_shared() if str(shared.name).endswith("step_count") + ] + assert counters == ["lbfgs/step_count"] + + def test_lbfgs_without_initial_scaling_starts_along_the_raw_gradient(): p = trainable(np.array([3.0, -4.0]), name="w") loss = (pt.constant(np.array([3.0, -4.0]), dtype=floatX) * p).sum() From a20f465d2dc614e25296260594949e7aaeb3a8fa Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 22:18:00 +0800 Subject: [PATCH 13/31] Build literal count and gamma at their dtype instead of casting --- pytensor_ml/optim/lbfgs.py | 14 +++++++++----- 1 file changed, 9 insertions(+), 5 deletions(-) diff --git a/pytensor_ml/optim/lbfgs.py b/pytensor_ml/optim/lbfgs.py index 7ee3c35..dec1378 100644 --- a/pytensor_ml/optim/lbfgs.py +++ b/pytensor_ml/optim/lbfgs.py @@ -79,11 +79,7 @@ def __init__(self, input_types=None, **kwargs): def filter_inputs(*inputs: Variable | float | int) -> tuple[Variable, ...]: count, gamma, *raw = inputs tensors = [pt.as_tensor_variable(tensor) for tensor in raw] - return ( - pt.as_tensor_variable(count).astype("int64"), - pt.as_tensor_variable(gamma).astype(tensors[0].dtype), - *tensors, - ) + return (_scalar_at(count, "int64"), _scalar_at(gamma, tensors[0].dtype), *tensors) def build_inner_graph(self, *inputs: TensorVariable) -> list[Variable]: n, m = self.n_parameters, self.memory_size @@ -136,6 +132,14 @@ def left_product(slot, alpha, *vector): return [v[-1] for v in r] +def _scalar_at(value: Variable | float | int, dtype: str) -> TensorVariable: + """Return ``value`` as a scalar of ``dtype``, built at that dtype rather than cast to it when it is a + literal, so no ``Cast`` node enters the graph for a Python number.""" + if isinstance(value, Variable): + return pt.as_tensor_variable(value).astype(dtype) + return pt.constant(value, dtype=dtype) + + def _require_stack_of( stack: TensorVariable, gradient: TensorVariable, memory_size: int, index: int ) -> None: From 914fffdca2a5f07344e79a9c20656747fe67a163 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 22:49:01 +0800 Subject: [PATCH 14/31] Require pytensor 3.3.2 --- conda_envs/environment-docs.yml | 2 +- conda_envs/pytensor_ml-gpu_jax.yml | 2 +- conda_envs/pytensor_ml.yml | 2 +- pyproject.toml | 4 ++-- 4 files changed, 5 insertions(+), 5 deletions(-) diff --git a/conda_envs/environment-docs.yml b/conda_envs/environment-docs.yml index e550874..f92673b 100644 --- a/conda_envs/environment-docs.yml +++ b/conda_envs/environment-docs.yml @@ -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.2,<3.4.0 - numpy - safetensors # The gallery extension renders notebook thumbnails with matplotlib. diff --git a/conda_envs/pytensor_ml-gpu_jax.yml b/conda_envs/pytensor_ml-gpu_jax.yml index ea6726a..eb08d29 100644 --- a/conda_envs/pytensor_ml-gpu_jax.yml +++ b/conda_envs/pytensor_ml-gpu_jax.yml @@ -6,7 +6,7 @@ channels: dependencies: - python>=3.12 - - pytensor>=3.2.3,<4.0.0 + - pytensor>=3.3.2,<4.0.0 - numpy - scikit-learn diff --git a/conda_envs/pytensor_ml.yml b/conda_envs/pytensor_ml.yml index 36717ef..aa8c2b0 100644 --- a/conda_envs/pytensor_ml.yml +++ b/conda_envs/pytensor_ml.yml @@ -5,7 +5,7 @@ channels: dependencies: - python>=3.12 - - pytensor>=3.3.0,<3.4.0 + - pytensor>=3.3.2,<3.4.0 - numpy - safetensors - scikit-learn diff --git a/pyproject.toml b/pyproject.toml index 73fc934..c8904d4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -35,7 +35,7 @@ keywords = [ ] dependencies = [ - "pytensor>=3.2.3,<4.0.0", + "pytensor>=3.3.2,<4.0.0", "numpy", ] @@ -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.2,<3.4.0" numpy = "*" safetensors = "*" # The gallery extension renders notebook thumbnails with matplotlib. From 41274c254ec76a12162fcbc46e15d0e19c7dc03d Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 12:10:35 -0500 Subject: [PATCH 15/31] Cast the mlx recursion to each parameter's dtype --- pytensor_ml/dispatch/mlx/lbfgs.py | 10 +++++-- tests/dispatch/mlx/test_lbfgs.py | 48 ++++++++++++++++++++++++++++++- 2 files changed, 54 insertions(+), 4 deletions(-) diff --git a/pytensor_ml/dispatch/mlx/lbfgs.py b/pytensor_ml/dispatch/mlx/lbfgs.py index 25ce3ac..9cd26d5 100644 --- a/pytensor_ml/dispatch/mlx/lbfgs.py +++ b/pytensor_ml/dispatch/mlx/lbfgs.py @@ -40,12 +40,16 @@ def direction(count, gamma, *tensors): alphas = [None] * m for position in reversed(range(m)): alphas[position] = curvatures[position] * dot(s_rows[position], vector) - vector = [v - alphas[position] * y_p for v, y_p in zip(vector, y_rows[position])] - vector = [gamma * v for v in 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) * s_p for v, s_p in zip(vector, s_rows[position]) + 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) diff --git a/tests/dispatch/mlx/test_lbfgs.py b/tests/dispatch/mlx/test_lbfgs.py index ecb7a90..caeaba9 100644 --- a/tests/dispatch/mlx/test_lbfgs.py +++ b/tests/dispatch/mlx/test_lbfgs.py @@ -12,7 +12,7 @@ from pytensor_ml.optim.lbfgs import LBFGSDirection from pytensor_ml.params import trainable from pytensor_ml.pytensorf import function -from tests.dispatch.mlx.test_basic import compare_mlx_and_py +from tests.dispatch.mlx.test_basic import compare_mlx_and_py, mlx_mode from tests.optim.test_lbfgs import dense_inverse_hessian, ring_stacks floatX = pytensor.config.floatX @@ -66,6 +66,52 @@ def test_a_single_parameter_returns_one_array(): compare_mlx_and_py([g_in, S_in, Y_in], d, [g, S, Y]) +def test_parameters_of_different_dtypes_keep_their_own(): + # The cross-parameter dot products come back at the widest dtype, and mlx promotes every vector they + # scale, so the narrower parameter's direction has to be cast back to the dtype its output declares. + rng = np.random.default_rng(sum(map(ord, "mixed dtypes"))) + dtypes = ["float32", "float16"] + shapes = [(3,), (2,)] + size = sum(int(np.prod(shape)) for shape in shapes) + gamma = 0.5 + gradient = rng.normal(size=size) + s = rng.normal(size=size) + y = s + rng.normal(scale=0.1, size=size) + splits = np.cumsum([int(np.prod(shape)) for shape in shapes])[:-1] + + def per_parameter(flat): + return [piece.astype(dtype) for piece, dtype in zip(np.split(flat, splits), dtypes)] + + def one_pair_stack(flat): + return [np.stack([piece, np.zeros_like(piece)]) for piece in per_parameter(flat)] + + gradients = [ + pt.tensor(f"g{i}", shape=shape, dtype=dtype) + for i, (shape, dtype) in enumerate(zip(shapes, dtypes)) + ] + S_in = [ + pt.tensor(f"S{i}", shape=(2, *shape), dtype=dtype) + for i, (shape, dtype) in enumerate(zip(shapes, dtypes)) + ] + Y_in = [ + pt.tensor(f"Y{i}", shape=(2, *shape), dtype=dtype) + for i, (shape, dtype) in enumerate(zip(shapes, dtypes)) + ] + outputs = LBFGSDirection(n_parameters=2, memory_size=2)( + 1, gamma, *gradients, *S_in, *Y_in, return_list=True + ) + direction = pytensor.function([*gradients, *S_in, *Y_in], outputs, mode=mlx_mode) + + got = [ + np.asarray(d) + for d in direction(*per_parameter(gradient), *one_pair_stack(s), *one_pair_stack(y)) + ] + + assert [d.dtype for d in got] == dtypes + want = dense_inverse_hessian(gamma, [(s, y)], size) @ gradient + np.testing.assert_allclose(np.concatenate(got), want, rtol=1e-2) + + @pytest.mark.parametrize("use_compile", [True, False], ids=["compiled", "eager"]) def test_the_rule_reaches_the_minimum_of_a_quadratic(use_compile): # The rule reads and writes its ring with a traced slot, which mlx traces only as advanced indexing From 5177de84a0a525a7362005628f9070be641390f3 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 12:11:04 -0500 Subject: [PATCH 16/31] Cache each admitted pair's curvature in rule state --- pytensor_ml/dispatch/mlx/lbfgs.py | 9 +---- pytensor_ml/optim/lbfgs.py | 63 +++++++++++++++---------------- pytensor_ml/optim/rules.py | 17 +++++++++ tests/dispatch/mlx/test_lbfgs.py | 8 ++-- tests/optim/test_lbfgs.py | 58 +++++++++++++++++++++------- tests/optim/test_rules.py | 7 +++- 6 files changed, 105 insertions(+), 57 deletions(-) diff --git a/pytensor_ml/dispatch/mlx/lbfgs.py b/pytensor_ml/dispatch/mlx/lbfgs.py index 9cd26d5..77383a5 100644 --- a/pytensor_ml/dispatch/mlx/lbfgs.py +++ b/pytensor_ml/dispatch/mlx/lbfgs.py @@ -20,7 +20,7 @@ def dot(left, right): # 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, *tensors): + def direction(count, gamma, rho, *tensors): gradients = tensors[:n] S = tensors[n : 2 * n] Y = tensors[2 * n :] @@ -29,12 +29,7 @@ def direction(count, gamma, *tensors): 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 = [] - for s, y in zip(s_rows, y_rows): - product = dot(s, y) - curvatures.append( - mx.where(product == 0, 0.0, 1.0 / mx.where(product == 0, 1.0, product)) - ) + curvatures = [mx.take(rho, slot) for slot in order] vector = list(gradients) alphas = [None] * m diff --git a/pytensor_ml/optim/lbfgs.py b/pytensor_ml/optim/lbfgs.py index dec1378..8ae9e3f 100644 --- a/pytensor_ml/optim/lbfgs.py +++ b/pytensor_ml/optim/lbfgs.py @@ -5,6 +5,7 @@ from pytensor.compile.builders import SymbolicOp from pytensor.graph.basic import Variable +from pytensor.scalar import upcast from pytensor.tensor import TensorVariable @@ -12,18 +13,18 @@ class LBFGSDirection(SymbolicOp): r""" Multiply a gradient by the L-BFGS inverse-Hessian approximation that a ring-buffered memory defines. - Inputs are ``count, gamma, g_1..g_n, S_1..S_n, Y_1..Y_n`` and outputs are ``d_1..d_n = H g``, one per - parameter. ``S_p`` and ``Y_p`` are ``(memory_size, *shape)`` stacks of past parameter differences + Inputs are ``count, gamma, rho, g_1..g_n, S_1..S_n, Y_1..Y_n`` and outputs are ``d_1..d_n = H g``, one + per parameter. ``S_p`` and ``Y_p`` are ``(memory_size, *shape)`` stacks of past parameter differences :math:`s` and gradient differences :math:`y` for parameter ``p``, written as a ring: slot ``(count - 1) % memory_size`` holds the newest pair and ``count`` is the number of pairs written so - far. A slot that holds nothing yet is all zeros and contributes nothing to the recursion, and a - writer retires a slot the same way. The op applies whatever pairs it is given: admitting only pairs - with :math:`y^\top s > 0`, which keeps the approximation positive definite, is the writer's job. + far. ``rho`` is the ``(memory_size,)`` vector of each slot's curvature :math:`\rho_i = 1 / (y_i^\top + s_i)`. A slot that holds nothing yet has zero stacks and a zero ``rho`` and contributes nothing to the + recursion, and a writer retires a slot the same way. The op applies whatever pairs it is given: admitting only + pairs with :math:`y^\top s > 0`, which keeps the approximation positive definite, is the writer's job. - The product is the two-loop recursion, algorithm 7.4 of :cite:t:`nocedal2006numerical`, with - :math:`\rho_i = 1 / (y_i^\top s_i)`. Each dot product sums over every parameter, so the memory of a - model with several parameters is treated as one vector and never copied into one. Starting from - :math:`\gamma I`, + The product is the two-loop recursion, algorithm 7.4 of :cite:t:`nocedal2006numerical`. Each dot + product sums over every parameter, so the memory of a model with several parameters is treated as one + vector and never copied into one. Starting from :math:`\gamma I`, .. math:: @@ -54,10 +55,11 @@ class LBFGSDirection(SymbolicOp): from pytensor_ml.optim.lbfgs import LBFGSDirection g = pt.vector("g") + rho = pt.vector("rho") S = pt.matrix("S") Y = pt.matrix("Y") - d = LBFGSDirection(n_parameters=1, memory_size=4)(1, 1.0, g, S, Y) - direction = pytensor.function([g, S, Y], d) + d = LBFGSDirection(n_parameters=1, memory_size=4)(1, 1.0, rho, g, S, Y) + direction = pytensor.function([rho, g, S, Y], d) References ---------- @@ -77,17 +79,28 @@ def __init__(self, input_types=None, **kwargs): @staticmethod def filter_inputs(*inputs: Variable | float | int) -> tuple[Variable, ...]: - count, gamma, *raw = inputs + count, gamma, rho, *raw = inputs tensors = [pt.as_tensor_variable(tensor) for tensor in raw] - return (_scalar_at(count, "int64"), _scalar_at(gamma, tensors[0].dtype), *tensors) + # The curvatures are cross-parameter dot products, so they live at the widest parameter dtype. + curvature_dtype = upcast(*(tensor.dtype for tensor in tensors)) + return ( + _scalar_at(count, "int64"), + _scalar_at(gamma, tensors[0].dtype), + pt.as_tensor_variable(rho).astype(curvature_dtype), + *tensors, + ) def build_inner_graph(self, *inputs: TensorVariable) -> list[Variable]: n, m = self.n_parameters, self.memory_size - count, gamma, *tensors = inputs + count, gamma, rho, *tensors = inputs if len(tensors) != 3 * n: raise ValueError( - f"LBFGSDirection with n_parameters={n} takes {3 * n} tensors after count and gamma, a " - f"gradient and two memory stacks per parameter, but got {len(tensors)}." + f"LBFGSDirection with n_parameters={n} takes {3 * n} tensors after count, gamma and rho, " + f"a gradient and two memory stacks per parameter, but got {len(tensors)}." + ) + if rho.type.ndim != 1 or rho.type.shape[0] not in (None, m): + raise ValueError( + f"rho must be a vector of one curvature per slot (memory_size={m}), but got {rho.type}." ) gradients = tensors[:n] S = tensors[n : 2 * n] @@ -97,12 +110,11 @@ def build_inner_graph(self, *inputs: TensorVariable) -> list[Variable]: _require_stack_of(stack, gradient, m, index) order = (count + pt.arange(m)) % m - curvatures = _curvatures(S, Y, m) def right_product(slot, *vector): s = [stack[slot] for stack in S] y = [stack[slot] for stack in Y] - alpha = curvatures[slot] * flat_dot(s, vector) + alpha = rho[slot] * flat_dot(s, vector) return [v - alpha.astype(v.dtype) * y_p for v, y_p in zip(vector, y)] + [alpha] *q, alphas = pytensor.scan( @@ -117,7 +129,7 @@ def right_product(slot, *vector): def left_product(slot, alpha, *vector): s = [stack[slot] for stack in S] y = [stack[slot] for stack in Y] - beta = curvatures[slot] * flat_dot(y, vector) + beta = rho[slot] * flat_dot(y, vector) return [v + (alpha - beta).astype(v.dtype) * s_p for v, s_p in zip(vector, s)] # The backward loop reports its alphas newest first and the forward loop reads them oldest first. @@ -160,16 +172,3 @@ def _require_stack_of( def flat_dot(left: Sequence[TensorVariable], right: Sequence[TensorVariable]) -> TensorVariable: """Dot product of two lists of tensors read as one flat vector each, through BLAS under numba.""" return pt.sum([pt.dot(a.ravel(), b.ravel()) for a, b in zip(left, right)]) - - -def _curvatures( - S: Sequence[TensorVariable], Y: Sequence[TensorVariable], memory_size: int -) -> TensorVariable: - """Return ``1 / (y_i . s_i)`` per slot, and zero for an empty slot rather than a division by zero.""" - # The same dot the writer's admission test uses, so a pair it admitted never rounds to a negative - # curvature here. - products = pt.stack( - [flat_dot([s[slot] for s in S], [y[slot] for y in Y]) for slot in range(memory_size)] - ) - empty = pt.eq(products, 0.0) - return pt.switch(empty, 0.0, 1.0 / pt.switch(empty, 1.0, products)) diff --git a/pytensor_ml/optim/rules.py b/pytensor_ml/optim/rules.py index 2a230eb..f3b9d90 100644 --- a/pytensor_ml/optim/rules.py +++ b/pytensor_ml/optim/rules.py @@ -1,11 +1,13 @@ from collections.abc import Callable, Sequence import numpy as np +import pytensor import pytensor.tensor as pt from pytensor import config from pytensor.compile.sharedvalue import SharedVariable from pytensor.graph.basic import Variable +from pytensor.scalar import upcast from pytensor.tensor import TensorVariable from pytensor_ml.optim.base import ( @@ -992,6 +994,13 @@ def lbfgs_updates( state_for(p, f"{namespace}/gradient_differences", history_size=memory_size) for p in parameters ] + # One curvature per slot, at the dtype of the cross-parameter dot that measures it. A slot that holds + # no pair keeps a zero, which the recursion reads as a pair that contributes nothing. + curvatures = pytensor.shared( + np.zeros(memory_size, dtype=upcast(*(gradient.dtype for gradient in gradients))), + name=f"{namespace}/curvatures", + shape=(memory_size,), + ) # The buffers hold zeros before the first step, so the differences read off them are meaningless # until a previous point exists; the guard below never lets those into the memory. @@ -1014,6 +1023,12 @@ def lbfgs_updates( pt.set_subtensor(memory[slot], pt.switch(accept, y[None], memory[slot])) for memory, y in zip(gradient_memory, gradient_differences) ] + # Stored inverted, from the very dot the guard tested, so an admitted pair's curvature is positive by + # construction. The divisor is swapped out on rejection so a zero curvature never reaches it. + rho = pt.reciprocal(pt.switch(accept, curvature, 1.0)).astype(curvatures.dtype) + new_curvatures = pt.set_subtensor( + curvatures[slot], pt.switch(accept, rho[None], curvatures[slot]) + ) new_pairs_written = pairs_written + accept.astype(pairs_written.dtype) if scale_init_precond: @@ -1035,6 +1050,7 @@ def lbfgs_updates( directions = LBFGSDirection(n_parameters=len(parameters), memory_size=memory_size)( new_pairs_written, identity_scale, + new_curvatures, *gradients, *new_value_memory, *new_gradient_memory, @@ -1044,6 +1060,7 @@ def lbfgs_updates( updates: Updates = Steps(incoming) updates[step_count] = step_count + 1 updates[pairs_written] = new_pairs_written + updates[curvatures] = new_curvatures for index, parameter in enumerate(parameters): updates[previous_values[index]] = parameter updates[previous_gradients[index]] = gradients[index] diff --git a/tests/dispatch/mlx/test_lbfgs.py b/tests/dispatch/mlx/test_lbfgs.py index caeaba9..e32342e 100644 --- a/tests/dispatch/mlx/test_lbfgs.py +++ b/tests/dispatch/mlx/test_lbfgs.py @@ -29,7 +29,7 @@ def test_direction_matches_py(n_pairs, count): for _ in range(n_pairs): s = rng.normal(size=size).astype(floatX) pairs.append((s, rng.normal(size=size).astype(floatX) + 0.5 * s)) - S, Y = ring_stacks(pairs, memory_size, count, shapes) + S, Y, rho = ring_stacks(pairs, memory_size, count, shapes) splits = np.cumsum([int(np.prod(shape)) for shape in shapes])[:-1] gradient_pieces = [ piece.reshape(shape) for piece, shape in zip(np.split(gradient, splits), shapes) @@ -39,7 +39,7 @@ def test_direction_matches_py(n_pairs, count): S_in = [pt.tensor(f"S{i}", shape=(memory_size, *shape)) for i, shape in enumerate(shapes)] Y_in = [pt.tensor(f"Y{i}", shape=(memory_size, *shape)) for i, shape in enumerate(shapes)] op = LBFGSDirection(n_parameters=2, memory_size=memory_size) - outputs = op(count, gamma, *gradients, *S_in, *Y_in, return_list=True) + outputs = op(count, gamma, rho, *gradients, *S_in, *Y_in, return_list=True) _, got = compare_mlx_and_py( [*gradients, *S_in, *Y_in], @@ -61,7 +61,7 @@ def test_a_single_parameter_returns_one_array(): S_in = pt.tensor("S", shape=(3, 5)) Y_in = pt.tensor("Y", shape=(3, 5)) - d = LBFGSDirection(n_parameters=1, memory_size=3)(0, 0.5, g_in, S_in, Y_in) + d = LBFGSDirection(n_parameters=1, memory_size=3)(0, 0.5, np.zeros(3), g_in, S_in, Y_in) compare_mlx_and_py([g_in, S_in, Y_in], d, [g, S, Y]) @@ -98,7 +98,7 @@ def one_pair_stack(flat): for i, (shape, dtype) in enumerate(zip(shapes, dtypes)) ] outputs = LBFGSDirection(n_parameters=2, memory_size=2)( - 1, gamma, *gradients, *S_in, *Y_in, return_list=True + 1, gamma, [1 / (y @ s), 0.0], *gradients, *S_in, *Y_in, return_list=True ) direction = pytensor.function([*gradients, *S_in, *Y_in], outputs, mode=mlx_mode) diff --git a/tests/optim/test_lbfgs.py b/tests/optim/test_lbfgs.py index 36772b4..5d63ca4 100644 --- a/tests/optim/test_lbfgs.py +++ b/tests/optim/test_lbfgs.py @@ -24,9 +24,11 @@ def dense_inverse_hessian(gamma, pairs, size): def ring_stacks(pairs, memory_size, count, shapes): - """Lay chronological flat pairs into per-parameter ring stacks, newest at ``(count - 1) % memory_size``.""" + """Lay chronological flat pairs into per-parameter ring stacks, newest at ``(count - 1) % memory_size``, + and each slot's curvature ``1 / (y . s)`` into a vector beside them, zero where a slot is empty.""" S = [np.zeros((memory_size, *shape), dtype=floatX) for shape in shapes] Y = [np.zeros((memory_size, *shape), dtype=floatX) for shape in shapes] + rho = np.zeros(memory_size, dtype=floatX) splits = np.cumsum([int(np.prod(shape)) for shape in shapes])[:-1] for age, (s, y) in enumerate(reversed(pairs)): slot = (count - 1 - age) % memory_size @@ -34,7 +36,8 @@ def ring_stacks(pairs, memory_size, count, shapes): stack[slot] = piece.reshape(stack.shape[1:]) for stack, piece in zip(Y, np.split(y, splits)): stack[slot] = piece.reshape(stack.shape[1:]) - return S, Y + rho[slot] = 1.0 / (y.astype(np.float64) @ s.astype(np.float64)) + return S, Y, rho @pytest.mark.parametrize("n_pairs, count", [(2, 2), (4, 6)], ids=["not_yet_wrapped", "wrapped"]) @@ -53,14 +56,15 @@ def test_direction_matches_the_two_loop_recursion_over_a_ring(n_pairs, count): s = rng.normal(size=size).astype(floatX) noise = rng.normal(size=size).astype(floatX) pairs.append((s, noise - (noise @ s) / (s @ s) * s + 0.5 * s)) # y . s = 0.5 s . s > 0 - S, Y = ring_stacks(pairs, memory_size, count, shapes) + S, Y, rho = ring_stacks(pairs, memory_size, count, shapes) op = LBFGSDirection(n_parameters=2, memory_size=memory_size) gradients = [pt.tensor(f"g{i}", shape=shape) for i, shape in enumerate(shapes)] S_in = [pt.tensor(f"S{i}", shape=(memory_size, *shape)) for i, shape in enumerate(shapes)] Y_in = [pt.tensor(f"Y{i}", shape=(memory_size, *shape)) for i, shape in enumerate(shapes)] direction = function( - [*gradients, *S_in, *Y_in], op(count, gamma, *gradients, *S_in, *Y_in, return_list=True) + [*gradients, *S_in, *Y_in], + op(count, gamma, rho, *gradients, *S_in, *Y_in, return_list=True), ) splits = np.cumsum([int(np.prod(shape)) for shape in shapes])[:-1] @@ -81,7 +85,7 @@ def test_parameters_of_different_dtypes_keep_their_own(): S_narrow, Y_narrow = (pt.tensor(name, shape=(2, 2), dtype="float32") for name in ("s", "y")) wide, narrow = LBFGSDirection(n_parameters=2, memory_size=2)( - 1, 0.5, g_wide, g_narrow, S_wide, S_narrow, Y_wide, Y_narrow, return_list=True + 1, 0.5, np.zeros(2), g_wide, g_narrow, S_wide, S_narrow, Y_wide, Y_narrow, return_list=True ) assert (wide.dtype, narrow.dtype) == ("float64", "float32") @@ -92,7 +96,7 @@ def test_a_scalar_parameter_has_vector_stacks(): S = pt.vector("S", dtype=floatX) Y = pt.vector("Y", dtype=floatX) - d = LBFGSDirection(n_parameters=1, memory_size=3)(1, 1.0, g, S, Y) + d = LBFGSDirection(n_parameters=1, memory_size=3)(1, 1.0, [0.0, 0.0, 1 / (1.5 * 3.0)], g, S, Y) # One pair (s, y) with y = 2 s: H y = s, so H maps g onto g / 2. np.testing.assert_allclose( @@ -111,43 +115,71 @@ def test_an_empty_memory_scales_the_gradient(): S = np.zeros((3, 5), dtype="float32") Y = np.zeros((3, 5), dtype="float32") - d = LBFGSDirection(n_parameters=1, memory_size=3)(0, 0.25, g, S, Y) + d = LBFGSDirection(n_parameters=1, memory_size=3)(0, 0.25, np.zeros(3), g, S, Y) assert d.dtype == "float32" np.testing.assert_allclose(d.eval(), 0.25 * g, rtol=1e-6) +def test_a_zero_curvature_retires_its_slot_whatever_the_stacks_hold(): + # The op applies the curvatures it is given and never measures them from the stacks, so a slot whose + # curvature is zero drops out of the recursion even with a pair still written in it. + rng = np.random.default_rng(2) + g = rng.normal(size=4).astype(floatX) + s = rng.normal(size=4).astype(floatX) + y = (s + 0.5 * rng.normal(size=4)).astype(floatX) + S = np.stack([s, np.zeros_like(s)]) + Y = np.stack([y, np.zeros_like(y)]) + + d = LBFGSDirection(n_parameters=1, memory_size=2)(1, 0.5, np.zeros(2), g, S, Y) + + np.testing.assert_allclose(d.eval(), 0.5 * g, rtol=RTOL) + + @pytest.mark.parametrize( "props, tensors, message", [ ({"n_parameters": 0, "memory_size": 3}, (), "n_parameters must be at least 1"), ( {"n_parameters": 1, "memory_size": 0}, - (np.ones(2), np.ones((0, 2)), np.ones((0, 2))), + (np.ones(0), np.ones(2), np.ones((0, 2)), np.ones((0, 2))), "memory_size must be at least 1", ), ( {"n_parameters": 1, "memory_size": 3}, - (np.ones(2), np.ones((3, 2)), np.ones((3, 2)), np.ones((3, 2))), + (np.ones(3), np.ones(2), np.ones((3, 2)), np.ones((3, 2)), np.ones((3, 2))), "takes 3 tensors", ), ( {"n_parameters": 1, "memory_size": 3}, - (np.ones(2), np.ones((4, 2)), np.ones((3, 2))), + (np.ones(3), np.ones(2), np.ones((4, 2)), np.ones((3, 2))), "memory_size=3", ), ( {"n_parameters": 1, "memory_size": 3}, - (np.ones(2), np.ones((3, 2, 1)), np.ones((3, 2))), + (np.ones(3), np.ones(2), np.ones((3, 2, 1)), np.ones((3, 2))), "memory_size=3", ), ( {"n_parameters": 1, "memory_size": 3}, - (np.ones(2, dtype="float32"), np.ones((3, 2)), np.ones((3, 2))), + (np.ones(3), np.ones(2, dtype="float32"), np.ones((3, 2)), np.ones((3, 2))), "dtype", ), + ( + {"n_parameters": 1, "memory_size": 3}, + (np.ones(4), np.ones(2), np.ones((3, 2)), np.ones((3, 2))), + "one curvature per slot", + ), + ], + ids=[ + "no_parameters", + "no_memory", + "extra_tensor", + "wrong_slots", + "wrong_rank", + "wrong_dtype", + "wrong_rho_length", ], - ids=["no_parameters", "no_memory", "extra_tensor", "wrong_slots", "wrong_rank", "wrong_dtype"], ) def test_malformed_inputs_are_refused_at_build_time(props, tensors, message): with pytest.raises(ValueError, match=message): diff --git a/tests/optim/test_rules.py b/tests/optim/test_rules.py index 904506e..b9e8205 100644 --- a/tests/optim/test_rules.py +++ b/tests/optim/test_rules.py @@ -679,12 +679,14 @@ def test_lbfgs_without_initial_scaling_starts_along_the_raw_gradient(): def test_lbfgs_rejects_a_pair_with_negative_curvature(): """A step whose gradient change opposes the parameter change would make the inverse-Hessian estimate indefinite, so the pair is left out of the memory, the ring index does not advance, and the next step - is the one an empty memory gives.""" + is the one an empty memory gives. An accepted pair's curvature is stored beside it as ``1 / (y . s)``.""" g = pt.vector("g") p = trainable(np.zeros(2), name="w") lr = 0.1 updates = lbfgs_updates([g], [p], learning_rate=lr, memory_size=2) memory = next(key for key in updates if key.name == "w/lbfgs/value_differences") + gradient_memory = next(key for key in updates if key.name == "w/lbfgs/gradient_differences") + curvatures = next(key for key in updates if key.name == "lbfgs/curvatures") pairs_written = next(key for key in updates if key.name == "lbfgs/pairs_written") fn = function([g], p, updates=updates) @@ -693,10 +695,13 @@ def test_lbfgs_rejects_a_pair_with_negative_curvature(): fn(np.array([2.0, 0.0], dtype=floatX)) # p moved along -g and g grew: y . s < 0, rejected assert int(pairs_written.get_value()) == 0 np.testing.assert_array_equal(memory.get_value(), 0.0) + np.testing.assert_array_equal(curvatures.get_value(), 0.0) np.testing.assert_allclose(p.get_value(), before - lr * 0.5 * np.array([2.0, 0.0]), rtol=RTOL) fn(np.array([0.5, 0.0], dtype=floatX)) # g shrank along the move: y . s > 0, accepted assert int(pairs_written.get_value()) == 1 assert np.any(memory.get_value()[0] != 0.0) + s, y = memory.get_value()[0], gradient_memory.get_value()[0] + np.testing.assert_allclose(curvatures.get_value(), [1.0 / (y @ s), 0.0], rtol=RTOL) def test_lbfgs_rejects_a_zero_memory_size(): From 215e6ba87322c8aaf6b549eb5a1e35482dbd5685 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:08:34 -0500 Subject: [PATCH 17/31] Require pytensor 3.3.3 --- conda_envs/environment-docs.yml | 2 +- conda_envs/pytensor_ml-gpu_jax.yml | 2 +- conda_envs/pytensor_ml.yml | 2 +- pyproject.toml | 4 ++-- 4 files changed, 5 insertions(+), 5 deletions(-) diff --git a/conda_envs/environment-docs.yml b/conda_envs/environment-docs.yml index f92673b..a6b7516 100644 --- a/conda_envs/environment-docs.yml +++ b/conda_envs/environment-docs.yml @@ -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.2,<3.4.0 + - pytensor>=3.3.3,<3.4.0 - numpy - safetensors # The gallery extension renders notebook thumbnails with matplotlib. diff --git a/conda_envs/pytensor_ml-gpu_jax.yml b/conda_envs/pytensor_ml-gpu_jax.yml index eb08d29..5ed9891 100644 --- a/conda_envs/pytensor_ml-gpu_jax.yml +++ b/conda_envs/pytensor_ml-gpu_jax.yml @@ -6,7 +6,7 @@ channels: dependencies: - python>=3.12 - - pytensor>=3.3.2,<4.0.0 + - pytensor>=3.3.3,<4.0.0 - numpy - scikit-learn diff --git a/conda_envs/pytensor_ml.yml b/conda_envs/pytensor_ml.yml index aa8c2b0..e8924c7 100644 --- a/conda_envs/pytensor_ml.yml +++ b/conda_envs/pytensor_ml.yml @@ -5,7 +5,7 @@ channels: dependencies: - python>=3.12 - - pytensor>=3.3.2,<3.4.0 + - pytensor>=3.3.3,<3.4.0 - numpy - safetensors - scikit-learn diff --git a/pyproject.toml b/pyproject.toml index c8904d4..7927f1b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -35,7 +35,7 @@ keywords = [ ] dependencies = [ - "pytensor>=3.3.2,<4.0.0", + "pytensor>=3.3.3,<4.0.0", "numpy", ] @@ -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.2,<3.4.0" +pytensor = ">=3.3.3,<3.4.0" numpy = "*" safetensors = "*" # The gallery extension renders notebook thumbnails with matplotlib. From 08d62157a21a0b7d38ad4cafad3d9f8885e11ac7 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:08:50 -0500 Subject: [PATCH 18/31] Move the L-BFGS test references into their own module --- tests/dispatch/mlx/test_lbfgs.py | 2 +- tests/optim/lbfgs_reference.py | 34 ++++++++++++++++++++++++++++++++ tests/optim/test_lbfgs.py | 31 +---------------------------- tests/optim/test_rules.py | 2 +- 4 files changed, 37 insertions(+), 32 deletions(-) create mode 100644 tests/optim/lbfgs_reference.py diff --git a/tests/dispatch/mlx/test_lbfgs.py b/tests/dispatch/mlx/test_lbfgs.py index e32342e..3ce5f86 100644 --- a/tests/dispatch/mlx/test_lbfgs.py +++ b/tests/dispatch/mlx/test_lbfgs.py @@ -13,7 +13,7 @@ from pytensor_ml.params import trainable from pytensor_ml.pytensorf import function from tests.dispatch.mlx.test_basic import compare_mlx_and_py, mlx_mode -from tests.optim.test_lbfgs import dense_inverse_hessian, ring_stacks +from tests.optim.lbfgs_reference import dense_inverse_hessian, ring_stacks floatX = pytensor.config.floatX diff --git a/tests/optim/lbfgs_reference.py b/tests/optim/lbfgs_reference.py new file mode 100644 index 0000000..3a8d7cb --- /dev/null +++ b/tests/optim/lbfgs_reference.py @@ -0,0 +1,34 @@ +import numpy as np +import pytensor + +floatX = pytensor.config.floatX + + +def dense_inverse_hessian(gamma, pairs, size): + """The matrix the two-loop recursion multiplies by, built from its definition: BFGS updates from + ``gamma I`` over ``(s, y)`` pairs oldest first, ``H <- V^T H V + rho s s^T`` with ``V = I - rho y s^T`` + (Nocedal and Wright, equation 7.16).""" + H = gamma * np.eye(size) + for s, y in pairs: + s, y = s.astype(np.float64), y.astype(np.float64) + rho = 1.0 / (y @ s) + V = np.eye(size) - rho * np.outer(y, s) + H = V.T @ H @ V + rho * np.outer(s, s) + return H + + +def ring_stacks(pairs, memory_size, count, shapes): + """Lay chronological flat pairs into per-parameter ring stacks, newest at ``(count - 1) % memory_size``, + and each slot's curvature ``1 / (y . s)`` into a vector beside them, zero where a slot is empty.""" + S = [np.zeros((memory_size, *shape), dtype=floatX) for shape in shapes] + Y = [np.zeros((memory_size, *shape), dtype=floatX) for shape in shapes] + rho = np.zeros(memory_size, dtype=floatX) + splits = np.cumsum([int(np.prod(shape)) for shape in shapes])[:-1] + for age, (s, y) in enumerate(reversed(pairs)): + slot = (count - 1 - age) % memory_size + for stack, piece in zip(S, np.split(s, splits)): + stack[slot] = piece.reshape(stack.shape[1:]) + for stack, piece in zip(Y, np.split(y, splits)): + stack[slot] = piece.reshape(stack.shape[1:]) + rho[slot] = 1.0 / (y.astype(np.float64) @ s.astype(np.float64)) + return S, Y, rho diff --git a/tests/optim/test_lbfgs.py b/tests/optim/test_lbfgs.py index 5d63ca4..d41f089 100644 --- a/tests/optim/test_lbfgs.py +++ b/tests/optim/test_lbfgs.py @@ -5,41 +5,12 @@ from pytensor_ml.optim.lbfgs import LBFGSDirection from pytensor_ml.pytensorf import function +from tests.optim.lbfgs_reference import dense_inverse_hessian, ring_stacks floatX = pytensor.config.floatX RTOL = 1e-6 if floatX == "float64" else 1e-4 -def dense_inverse_hessian(gamma, pairs, size): - """The matrix the two-loop recursion multiplies by, built from its definition: BFGS updates from - ``gamma I`` over ``(s, y)`` pairs oldest first, ``H <- V^T H V + rho s s^T`` with ``V = I - rho y s^T`` - (Nocedal and Wright, equation 7.16).""" - H = gamma * np.eye(size) - for s, y in pairs: - s, y = s.astype(np.float64), y.astype(np.float64) - rho = 1.0 / (y @ s) - V = np.eye(size) - rho * np.outer(y, s) - H = V.T @ H @ V + rho * np.outer(s, s) - return H - - -def ring_stacks(pairs, memory_size, count, shapes): - """Lay chronological flat pairs into per-parameter ring stacks, newest at ``(count - 1) % memory_size``, - and each slot's curvature ``1 / (y . s)`` into a vector beside them, zero where a slot is empty.""" - S = [np.zeros((memory_size, *shape), dtype=floatX) for shape in shapes] - Y = [np.zeros((memory_size, *shape), dtype=floatX) for shape in shapes] - rho = np.zeros(memory_size, dtype=floatX) - splits = np.cumsum([int(np.prod(shape)) for shape in shapes])[:-1] - for age, (s, y) in enumerate(reversed(pairs)): - slot = (count - 1 - age) % memory_size - for stack, piece in zip(S, np.split(s, splits)): - stack[slot] = piece.reshape(stack.shape[1:]) - for stack, piece in zip(Y, np.split(y, splits)): - stack[slot] = piece.reshape(stack.shape[1:]) - rho[slot] = 1.0 / (y.astype(np.float64) @ s.astype(np.float64)) - return S, Y, rho - - @pytest.mark.parametrize("n_pairs, count", [(2, 2), (4, 6)], ids=["not_yet_wrapped", "wrapped"]) def test_direction_matches_the_two_loop_recursion_over_a_ring(n_pairs, count): # The reference is the dense matrix the recursion is an algorithm for, built from the textbook update diff --git a/tests/optim/test_rules.py b/tests/optim/test_rules.py index b9e8205..2f35d19 100644 --- a/tests/optim/test_rules.py +++ b/tests/optim/test_rules.py @@ -34,7 +34,7 @@ ) from pytensor_ml.optim import alias as alias_module from pytensor_ml.pytensorf import function -from tests.optim.test_lbfgs import dense_inverse_hessian +from tests.optim.lbfgs_reference import dense_inverse_hessian floatX = pytensor.config.floatX From c6d83387fc0fed0f8156bc1f9006e13fa3f716be Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:09:41 -0500 Subject: [PATCH 19/31] Require static slot counts in LBFGSDirection --- pytensor_ml/optim/lbfgs.py | 24 ++++++++++++++---------- tests/optim/test_lbfgs.py | 10 ++++++++-- 2 files changed, 22 insertions(+), 12 deletions(-) diff --git a/pytensor_ml/optim/lbfgs.py b/pytensor_ml/optim/lbfgs.py index 8ae9e3f..043197d 100644 --- a/pytensor_ml/optim/lbfgs.py +++ b/pytensor_ml/optim/lbfgs.py @@ -45,7 +45,8 @@ class LBFGSDirection(SymbolicOp): Examples -------- - Compile the direction for one vector parameter and a memory of four slots, with one pair written: + Compile the direction for one vector parameter and a memory of four slots, with one pair written. The + slot count of ``rho`` and the stacks has to be static: .. code-block:: python @@ -55,9 +56,9 @@ class LBFGSDirection(SymbolicOp): from pytensor_ml.optim.lbfgs import LBFGSDirection g = pt.vector("g") - rho = pt.vector("rho") - S = pt.matrix("S") - Y = pt.matrix("Y") + rho = pt.tensor("rho", shape=(4,)) + S = pt.tensor("S", shape=(4, None)) + Y = pt.tensor("Y", shape=(4, None)) d = LBFGSDirection(n_parameters=1, memory_size=4)(1, 1.0, rho, g, S, Y) direction = pytensor.function([rho, g, S, Y], d) @@ -98,9 +99,12 @@ def build_inner_graph(self, *inputs: TensorVariable) -> list[Variable]: f"LBFGSDirection with n_parameters={n} takes {3 * n} tensors after count, gamma and rho, " f"a gradient and two memory stacks per parameter, but got {len(tensors)}." ) - if rho.type.ndim != 1 or rho.type.shape[0] not in (None, m): + # The recursion walks exactly memory_size slots, so a slot count known only at runtime could + # silently drop rows or index past the end. + if rho.type.ndim != 1 or rho.type.shape[0] != m: raise ValueError( - f"rho must be a vector of one curvature per slot (memory_size={m}), but got {rho.type}." + f"rho must be a vector of one curvature per slot, with a static length of " + f"memory_size={m}, but got {rho.type}." ) gradients = tensors[:n] S = tensors[n : 2 * n] @@ -155,17 +159,17 @@ def _scalar_at(value: Variable | float | int, dtype: str) -> TensorVariable: def _require_stack_of( stack: TensorVariable, gradient: TensorVariable, memory_size: int, index: int ) -> None: - """Raise unless ``stack`` is ``memory_size`` slots of ``gradient``'s shape and dtype.""" + """Raise unless ``stack`` is a static ``memory_size`` slots of ``gradient``'s shape and dtype.""" slots = stack.type.shape[0] if stack.type.ndim else None if ( stack.type.ndim != gradient.type.ndim + 1 or stack.type.dtype != gradient.type.dtype - or (slots is not None and slots != memory_size) + or slots != memory_size ): raise ValueError( f"The memory stacks of parameter {index} must be shaped (memory_size={memory_size}, " - f"*gradient.shape) at the gradient's dtype, but got {stack.type} for a gradient of type " - f"{gradient.type}." + f"*gradient.shape) at the gradient's dtype, with the slot count static, but got " + f"{stack.type} for a gradient of type {gradient.type}." ) diff --git a/tests/optim/test_lbfgs.py b/tests/optim/test_lbfgs.py index d41f089..e912566 100644 --- a/tests/optim/test_lbfgs.py +++ b/tests/optim/test_lbfgs.py @@ -64,8 +64,8 @@ def test_parameters_of_different_dtypes_keep_their_own(): def test_a_scalar_parameter_has_vector_stacks(): g = pt.scalar("g", dtype=floatX) - S = pt.vector("S", dtype=floatX) - Y = pt.vector("Y", dtype=floatX) + S = pt.tensor("S", shape=(3,), dtype=floatX) + Y = pt.tensor("Y", shape=(3,), dtype=floatX) d = LBFGSDirection(n_parameters=1, memory_size=3)(1, 1.0, [0.0, 0.0, 1 / (1.5 * 3.0)], g, S, Y) @@ -141,6 +141,11 @@ def test_a_zero_curvature_retires_its_slot_whatever_the_stacks_hold(): (np.ones(4), np.ones(2), np.ones((3, 2)), np.ones((3, 2))), "one curvature per slot", ), + ( + {"n_parameters": 1, "memory_size": 3}, + (np.ones(3), np.ones(2), pt.tensor("S", shape=(None, 2)), np.ones((3, 2))), + "slot count static", + ), ], ids=[ "no_parameters", @@ -150,6 +155,7 @@ def test_a_zero_curvature_retires_its_slot_whatever_the_stacks_hold(): "wrong_rank", "wrong_dtype", "wrong_rho_length", + "dynamic_slots", ], ) def test_malformed_inputs_are_refused_at_build_time(props, tensors, message): From f443bc5a8371dbc882605508be3a61ec594f6b20 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:10:11 -0500 Subject: [PATCH 20/31] Build gamma at the curvature dtype --- pytensor_ml/optim/lbfgs.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/pytensor_ml/optim/lbfgs.py b/pytensor_ml/optim/lbfgs.py index 043197d..2d49b78 100644 --- a/pytensor_ml/optim/lbfgs.py +++ b/pytensor_ml/optim/lbfgs.py @@ -82,11 +82,12 @@ def __init__(self, input_types=None, **kwargs): def filter_inputs(*inputs: Variable | float | int) -> tuple[Variable, ...]: count, gamma, rho, *raw = inputs tensors = [pt.as_tensor_variable(tensor) for tensor in raw] - # The curvatures are cross-parameter dot products, so they live at the widest parameter dtype. + # The curvatures and gamma come from cross-parameter dot products, so they live at the widest + # parameter dtype. curvature_dtype = upcast(*(tensor.dtype for tensor in tensors)) return ( _scalar_at(count, "int64"), - _scalar_at(gamma, tensors[0].dtype), + _scalar_at(gamma, curvature_dtype), pt.as_tensor_variable(rho).astype(curvature_dtype), *tensors, ) From 7744434ff13f0384b059b2dfd15d5e2cb2bd64cc Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:10:28 -0500 Subject: [PATCH 21/31] Carry the L-BFGS identity scale in rule state --- pytensor_ml/optim/rules.py | 21 ++++++++++++--------- 1 file changed, 12 insertions(+), 9 deletions(-) diff --git a/pytensor_ml/optim/rules.py b/pytensor_ml/optim/rules.py index f3b9d90..9409bf6 100644 --- a/pytensor_ml/optim/rules.py +++ b/pytensor_ml/optim/rules.py @@ -1031,19 +1031,23 @@ def lbfgs_updates( ) new_pairs_written = pairs_written + accept.astype(pairs_written.dtype) + updates: Updates = Steps(incoming) if scale_init_precond: - has_pair = new_pairs_written > 0 - newest = ((new_pairs_written - 1) % memory_size)[None] - newest_s = [memory[newest] for memory in new_value_memory] - newest_y = [memory[newest] for memory in new_gradient_memory] - newest_curvature = flat_dot(newest_s, newest_y) - newest_change = pt.switch(has_pair, flat_dot(newest_y, newest_y), 1.0) + # The newest admitted pair's s . y / y . y, carried from the step that admitted it, since a + # rejected step leaves the newest pair in memory unchanged. + newest_pair_scale = scalar_state(f"{namespace}/identity_scale", dtype=curvatures.dtype) + new_newest_pair_scale = pt.switch( + accept, + curvature / pt.switch(accept, gradient_change, 1.0), + newest_pair_scale, + ).astype(newest_pair_scale.dtype) gradient_norm = pt.sqrt(flat_dot(gradients, gradients)) identity_scale = pt.switch( - has_pair, - newest_curvature / newest_change, + new_pairs_written > 0, + new_newest_pair_scale, pt.minimum(1.0, 1.0 / pt.switch(gradient_norm > 0, gradient_norm, 1.0)), ) + updates[newest_pair_scale] = new_newest_pair_scale else: identity_scale = 1.0 @@ -1057,7 +1061,6 @@ def lbfgs_updates( return_list=True, ) - updates: Updates = Steps(incoming) updates[step_count] = step_count + 1 updates[pairs_written] = new_pairs_written updates[curvatures] = new_curvatures From 2a29785a1a6978a0301917e4ecf805e81f2114af Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:10:46 -0500 Subject: [PATCH 22/31] Refuse L-BFGS pairs whose curvature overflows --- pytensor_ml/optim/rules.py | 9 ++++++++- tests/optim/test_rules.py | 14 ++++++++++++++ 2 files changed, 22 insertions(+), 1 deletion(-) diff --git a/pytensor_ml/optim/rules.py b/pytensor_ml/optim/rules.py index 9409bf6..8b0be8a 100644 --- a/pytensor_ml/optim/rules.py +++ b/pytensor_ml/optim/rules.py @@ -1009,7 +1009,14 @@ def lbfgs_updates( curvature = flat_dot(gradient_differences, value_differences) gradient_change = flat_dot(gradient_differences, gradient_differences) epsilon = max(np.finfo(gradient.dtype).eps for gradient in gradients) - accept = (step_count > 0) & (curvature > epsilon * gradient_change) + # Near a minimum y . y underflows to zero, so the relative test alone admits a pair whose inverse + # curvature or identity scale overflows; both have to be representable to enter the memory. + accept = ( + (step_count > 0) + & (curvature > epsilon * gradient_change) + & pt.isfinite(pt.reciprocal(curvature)) + & pt.isfinite(curvature / gradient_change) + ) # Rejection rewrites the slot with itself, so the write stays in place and unconditional; only the # count decides whether the slot is now part of the memory. The slot is a one-element index vector diff --git a/tests/optim/test_rules.py b/tests/optim/test_rules.py index 2f35d19..a5a80c6 100644 --- a/tests/optim/test_rules.py +++ b/tests/optim/test_rules.py @@ -704,6 +704,20 @@ def test_lbfgs_rejects_a_pair_with_negative_curvature(): np.testing.assert_allclose(curvatures.get_value(), [1.0 / (y @ s), 0.0], rtol=RTOL) +def test_lbfgs_stays_finite_after_it_converges(): + """Past the minimum the gradient changes underflow, so ``y . y`` reaches zero while ``y . s`` is + still a positive subnormal; a pair admitted then stores an infinite ``1 / (y . s)`` and the next + step is NaN. float32 at any floatX, where the underflow arrives within a few steps.""" + p = params.trainable(np.ones(2, dtype="float32"), name="w") + loss = (p**2).sum() + step = function([], loss, updates=lbfgs_updates(loss, [p])) + + for _ in range(10): + step() + + np.testing.assert_array_equal(p.get_value(), 0.0) + + def test_lbfgs_rejects_a_zero_memory_size(): p = trainable(np.zeros(2), name="w") with pytest.raises(ValueError, match="memory_size must be at least 1"): From 20e231da5e76a3f4d2860e08ec07fb04cc997068 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:11:03 -0500 Subject: [PATCH 23/31] Test L-BFGS on a one-slot ring --- tests/optim/test_rules.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/tests/optim/test_rules.py b/tests/optim/test_rules.py index a5a80c6..9bf74d1 100644 --- a/tests/optim/test_rules.py +++ b/tests/optim/test_rules.py @@ -615,11 +615,13 @@ def test_lbfgs_first_step_is_the_gradient_capped_to_the_unit_ball(gradient): np.testing.assert_allclose(p.get_value(), -min(1.0, 1.0 / np.linalg.norm(g)) * g, rtol=RTOL) -def test_lbfgs_step_matches_the_dense_update_through_a_ring_wrap(): +@pytest.mark.parametrize("memory_size", [1, 2], ids=["one_slot", "two_slots"]) +def test_lbfgs_step_matches_the_dense_update_through_a_ring_wrap(memory_size): """On a strictly convex quadratic every pair is accepted, so the memory is the last ``memory_size`` - chronological pairs and each step is ``-lr * H g`` for the dense BFGS matrix built from them. Two - slots over six steps wrap the ring twice; a rule that overwrote the wrong slot or read the newest - pair off by one would drift from the dense reference from the third step on.""" + chronological pairs and each step is ``-lr * H g`` for the dense BFGS matrix built from them. Over + six steps one slot is overwritten every step, where the newest pair is also the oldest, and two slots + wrap the ring twice; a rule that overwrote the wrong slot or read the newest pair off by one would + drift from the dense reference from the third step on.""" A = np.diag([1.0, 2.0, 3.0, 4.0, 5.0]) + 0.1 A = A @ A.T b = np.array([0.3, -1.0, 2.0, 0.5, -0.7]) @@ -627,7 +629,7 @@ def test_lbfgs_step_matches_the_dense_update_through_a_ring_wrap(): v = trainable(np.array([3.0, 1.0]), name="v") x = pt.concatenate([u, v]) loss = 0.5 * x @ pt.constant(A, dtype=floatX) @ x - pt.constant(b, dtype=floatX) @ x - memory_size, lr = 2, 0.5 + lr = 0.5 updates = lbfgs_updates(loss, [u, v], learning_rate=lr, memory_size=memory_size) step = function([], pt.grad(loss, [u, v]), updates=updates) # gradient before the update From 3e036f3a7c3838264f265fbb458111155287ede6 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:11:20 -0500 Subject: [PATCH 24/31] Test that L-BFGS rejects a negligible positive curvature --- tests/optim/test_rules.py | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/tests/optim/test_rules.py b/tests/optim/test_rules.py index 9bf74d1..4fec3a9 100644 --- a/tests/optim/test_rules.py +++ b/tests/optim/test_rules.py @@ -706,6 +706,22 @@ def test_lbfgs_rejects_a_pair_with_negative_curvature(): np.testing.assert_allclose(curvatures.get_value(), [1.0 / (y @ s), 0.0], rtol=RTOL) +def test_lbfgs_rejects_a_pair_whose_curvature_is_positive_but_negligible(): + """The guard asks for ``y . s > eps * y . y``, not only a positive sign: a pair whose curvature is + tiny next to its gradient change would put a near-singular ``1 / (y . s)`` into the memory.""" + g = pt.vector("g") + p = trainable(np.zeros(2), name="w") + updates = lbfgs_updates([g], [p], learning_rate=0.1, memory_size=2) + pairs_written = next(key for key in updates if key.name == "lbfgs/pairs_written") + fn = function([g], p, updates=updates) + + fn(np.array([1.0, 0.0], dtype=floatX)) # s = [-0.1, 0] on the next step + # y = [-0.5, 1e9]: y . s = 0.05 > 0, but eps * y . y is about 1e18 * eps, far above it + fn(np.array([0.5, 1e9], dtype=floatX)) + + assert int(pairs_written.get_value()) == 0 + + def test_lbfgs_stays_finite_after_it_converges(): """Past the minimum the gradient changes underflow, so ``y . y`` reaches zero while ``y . s`` is still a positive subnormal; a pair admitted then stores an infinite ``1 / (y . s)`` and the next From e5629c2c2d18dea4ccc900e28f228e55bbb97561 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:11:37 -0500 Subject: [PATCH 25/31] Test L-BFGS with parameters of mixed dtypes --- tests/optim/test_rules.py | 21 +++++++++++++++++++++ 1 file changed, 21 insertions(+) diff --git a/tests/optim/test_rules.py b/tests/optim/test_rules.py index 4fec3a9..6352951 100644 --- a/tests/optim/test_rules.py +++ b/tests/optim/test_rules.py @@ -736,6 +736,27 @@ def test_lbfgs_stays_finite_after_it_converges(): np.testing.assert_array_equal(p.get_value(), 0.0) +def test_lbfgs_parameters_of_different_dtypes_reach_the_minimum(): + # The curvatures and the identity scale are cross-parameter dots, so they are kept at the widest + # parameter dtype while each parameter keeps its own. + u = params.trainable(np.array([5.0], dtype="float64"), name="u") + v = params.trainable(np.array([-3.0], dtype="float32"), name="v") + loss = 0.5 * ((u - 1.0) ** 2).sum() + 2.0 * ((v + 2.0) ** 2).sum() + updates = lbfgs_updates(loss, [u, v], memory_size=2) + curvatures = next(key for key in updates if key.name == "lbfgs/curvatures") + step = function([], loss, updates=updates) + + for _ in range(12): + step() + + assert (u.get_value().dtype, v.get_value().dtype, curvatures.dtype) == ( + "float64", + "float32", + "float64", + ) + np.testing.assert_allclose([u.get_value()[0], v.get_value()[0]], [1.0, -2.0], rtol=1e-5) + + def test_lbfgs_rejects_a_zero_memory_size(): p = trainable(np.zeros(2), name="w") with pytest.raises(ValueError, match="memory_size must be at least 1"): From 883e7422214b0f929d30d743096963a7b5a62507 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:11:54 -0500 Subject: [PATCH 26/31] Test that L-BFGS resumes from a checkpoint --- tests/optim/test_rules.py | 23 +++++++++++++++++++++++ 1 file changed, 23 insertions(+) diff --git a/tests/optim/test_rules.py b/tests/optim/test_rules.py index 6352951..232527d 100644 --- a/tests/optim/test_rules.py +++ b/tests/optim/test_rules.py @@ -8,6 +8,7 @@ from pytensor.gradient import DisconnectedInputError, grad from pytensor_ml import params +from pytensor_ml.checkpoint import load_state, save_state from pytensor_ml.optim import ( adadelta, adadelta_updates, @@ -757,6 +758,28 @@ def test_lbfgs_parameters_of_different_dtypes_reach_the_minimum(): np.testing.assert_allclose([u.get_value()[0], v.get_value()[0]], [1.0, -2.0], rtol=1e-5) +def test_lbfgs_resumes_its_trajectory_from_a_checkpoint(tmp_path): + """Every piece of the rule's state is named and saved, so a run restored mid-way retraces the + steps it took the first time; a ring index or curvature left behind would desynchronize the memory + from its order.""" + A = np.array([[3.0, 0.5], [0.5, 1.0]]) + p = trainable(np.array([5.0, -3.0]), name="w") + loss = 0.5 * p @ pt.constant(A, dtype=floatX) @ p + updates = lbfgs_updates(loss, [p], learning_rate=0.5, memory_size=2) + state = list(updates) + step = function([], loss, updates=updates) + for _ in range(3): + step() + path = tmp_path / "lbfgs.safetensors" + save_state(state, path) + + first = [float(step()) for _ in range(4)] + load_state(state, path) + second = [float(step()) for _ in range(4)] + + np.testing.assert_array_equal(second, first) + + def test_lbfgs_rejects_a_zero_memory_size(): p = trainable(np.zeros(2), name="w") with pytest.raises(ValueError, match="memory_size must be at least 1"): From 218b6355bf02db6b09cdbbd0780487973ddd5a58 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:12:12 -0500 Subject: [PATCH 27/31] Run lbfgs at its default rate in the loss-reduction test --- tests/optim/test_rules.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/optim/test_rules.py b/tests/optim/test_rules.py index 232527d..3303be5 100644 --- a/tests/optim/test_rules.py +++ b/tests/optim/test_rules.py @@ -66,7 +66,7 @@ def trainable(value, name=None, **kwargs): nadam(learning_rate=1e-2), adamax(learning_rate=1e-2), rprop(learning_rate=1e-2), - lbfgs(learning_rate=1e-2), + lbfgs(), ], ids=[ "sgd", From e2fbdf3fa551c1795e10e2052a70c18233936627 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:12:29 -0500 Subject: [PATCH 28/31] Test lbfgs with its step clipped after it --- tests/optim/test_composition.py | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/tests/optim/test_composition.py b/tests/optim/test_composition.py index d58032a..9b9f30c 100644 --- a/tests/optim/test_composition.py +++ b/tests/optim/test_composition.py @@ -16,6 +16,7 @@ compile_train, cosine_schedule, large_step, + lbfgs, reduce_on_plateau, scalar_state, scale, @@ -42,6 +43,18 @@ def state_named(step, name): return next(variable for variable in step.get_shared() if variable.name == name) +def test_lbfgs_converges_with_its_step_clipped_after_it(): + """L-BFGS reads curvature from the parameter move and the raw gradients, so a clip after the rule + bounds the step without corrupting the pairs it stores, and the run still reaches the minimizer.""" + p, loss = quadratic_problem() + step = compile_train(loss, chain(lbfgs(), clip_by_global_norm(0.5))) + + for _ in range(10): + step(GOOD) + + np.testing.assert_allclose(p.get_value(), 0.0, atol=1e-6) + + def test_clipping_bounds_a_rules_step_end_to_end(): """The clipping transform is otherwise only exercised on a hand-built updates dict; here it has to survive a real rule, a real gradient, and compile_train's assembly.""" From a6fb4d1f9d7d232c779eedf0a13f7ce33f792662 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:12:47 -0500 Subject: [PATCH 29/31] Test that skip_if holds back the L-BFGS state --- tests/optim/test_composition.py | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) diff --git a/tests/optim/test_composition.py b/tests/optim/test_composition.py index 9b9f30c..d38734a 100644 --- a/tests/optim/test_composition.py +++ b/tests/optim/test_composition.py @@ -55,6 +55,25 @@ def test_lbfgs_converges_with_its_step_clipped_after_it(): np.testing.assert_allclose(p.get_value(), 0.0, atol=1e-6) +def test_skip_if_holds_back_all_of_the_lbfgs_state(): + """A skipped step must leave the memory exactly as it was, or the next applied step pairs a stale + previous gradient with a fresh parameter and stores a pair that is not a secant.""" + _, loss = quadratic_problem() + step = compile_train(loss, skip_if(lbfgs(), max_consecutive_skips=None)) + for _ in range(2): + step(GOOD) + lbfgs_state = [variable for variable in step.get_shared() if "lbfgs/" in str(variable.name)] + before = {variable.name: np.array(variable.get_value()) for variable in lbfgs_state} + + step(BAD) + + for variable in lbfgs_state: + np.testing.assert_array_equal(variable.get_value(), before[variable.name]) + step(GOOD) + pairs_written = state_named(step, "lbfgs/pairs_written") + assert int(pairs_written.get_value()) == int(before["lbfgs/pairs_written"]) + 1 + + def test_clipping_bounds_a_rules_step_end_to_end(): """The clipping transform is otherwise only exercised on a hand-built updates dict; here it has to survive a real rule, a real gradient, and compile_train's assembly.""" From 18ee0167931dee38b869bed112a6a7b9b31f2808 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:13:05 -0500 Subject: [PATCH 30/31] Document the fixed-step contract of lbfgs --- pytensor_ml/optim/alias.py | 5 +++-- pytensor_ml/optim/base.py | 2 ++ pytensor_ml/optim/rules.py | 16 +++++++++++++--- 3 files changed, 18 insertions(+), 5 deletions(-) diff --git a/pytensor_ml/optim/alias.py b/pytensor_ml/optim/alias.py index a1c6208..335495b 100644 --- a/pytensor_ml/optim/alias.py +++ b/pytensor_ml/optim/alias.py @@ -374,8 +374,9 @@ def lbfgs( Examples -------- A quasi-Newton direction from a memory of recent parameter and gradient differences, taken at a - fixed fraction. 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. + 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 diff --git a/pytensor_ml/optim/base.py b/pytensor_ml/optim/base.py index 2febacc..068d603 100644 --- a/pytensor_ml/optim/base.py +++ b/pytensor_ml/optim/base.py @@ -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. """ diff --git a/pytensor_ml/optim/rules.py b/pytensor_ml/optim/rules.py index 8b0be8a..bb922b6 100644 --- a/pytensor_ml/optim/rules.py +++ b/pytensor_ml/optim/rules.py @@ -927,11 +927,22 @@ def lbfgs_updates( accepted :math:`\gamma = \min(1, 1 / \|g\|)`, which keeps the first step inside the unit ball. The step is :math:`p \leftarrow p - \eta H g`. - The direction is well scaled once the memory holds a pair, so :math:`\eta = 1` is the natural rate - and a line search the natural way to back off from it. Consecutive gradients have to be measured on + The direction is well scaled once the memory holds a pair, so :math:`\eta = 1` is the natural rate. + The rule takes every step at that rate: there is no line search yet (pymc-devs/pytensor-ml#58), so + the rate is the only safeguard against a bad direction. Consecutive gradients have to be measured on the same objective for their difference to be curvature, so the rule assumes a deterministic, full-batch loss. + Three uses break that assumption. A gradient transform ahead of the rule in a chain, such as + :func:`~pytensor_ml.optim.clipping.clip_by_global_norm`, hands it gradients whose differences are not + curvature, and the rule can diverge. Clip after the rule instead: the parameter differences are read + off the parameters, so a clipped step still forms a valid pair. Wrapping the rule in + :func:`~pytensor_ml.optim.guards.skip_if` does not rescue a bad step either, because the loss is + deterministic and a skipped step is recomputed unchanged on the next call until the guard raises. A + step that a guard would skip calls for a smaller rate. Finally, a parameter written between steps, + with ``set_value`` for instance, forms a pair from a move the rule did not make, and that pair stays + in the memory for up to ``memory_size`` steps. + Parameters ---------- loss_gradients_or_updates : TensorVariable, sequence of TensorVariable, or Updates @@ -946,7 +957,6 @@ def lbfgs_updates( scale_init_precond : bool Start the recursion from :math:`\gamma I` as above. When False it starts from the identity, and the first step is the raw gradient. Default True. - namespace : str Prefix for every state slot this rule allocates, so two rules in one graph keep separate state rather than reusing each other's. Default is the rule's own name. From 0bf07b09f897cec0a6f996501dab5d5971403cc2 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Fri, 2 Oct 2026 13:13:23 -0500 Subject: [PATCH 31/31] List LBFGSDirection in the optim API docs --- docs/source/api/optim.rst | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/docs/source/api/optim.rst b/docs/source/api/optim.rst index 6de1b94..b3fb21d 100644 --- a/docs/source/api/optim.rst +++ b/docs/source/api/optim.rst @@ -132,3 +132,10 @@ Low-level update functions adagrad_updates adadelta_updates lbfgs_updates + +.. currentmodule:: pytensor_ml.optim.lbfgs + +.. autosummary:: + :toctree: generated/ + + LBFGSDirection