diff --git a/.github/workflows/run_tests.yml b/.github/workflows/run_tests.yml index 829e952..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.30,<0.32"; fi + if [[ $INSTALL_MLX == "1" ]]; then pip install "mlx>=0.32.2,<0.33"; fi if [[ $INSTALL_TORCH == "1" ]]; then pip install torch --index-url https://download.pytorch.org/whl/cpu; fi env: INSTALL_JAX: ${{ matrix.install-jax }} diff --git a/conda_envs/environment-docs.yml b/conda_envs/environment-docs.yml index e550874..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.0,<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 ea6726a..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.2.3,<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 36717ef..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.0,<3.4.0 + - pytensor>=3.3.3,<3.4.0 - numpy - safetensors - scikit-learn diff --git a/docs/source/api/optim.rst b/docs/source/api/optim.rst index 632a92c..b3fb21d 100644 --- a/docs/source/api/optim.rst +++ b/docs/source/api/optim.rst @@ -26,6 +26,7 @@ Update rules rprop adagrad adadelta + lbfgs Transforms ---------- @@ -130,3 +131,11 @@ Low-level update functions rprop_updates adagrad_updates adadelta_updates + lbfgs_updates + +.. currentmodule:: pytensor_ml.optim.lbfgs + +.. autosummary:: + :toctree: generated/ + + LBFGSDirection 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/pyproject.toml b/pyproject.toml index 73fc934..7927f1b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -35,7 +35,7 @@ keywords = [ ] dependencies = [ - "pytensor>=3.2.3,<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.0,<3.4.0" +pytensor = ">=3.3.3,<3.4.0" numpy = "*" safetensors = "*" # The gallery extension renders notebook thumbnails with matplotlib. 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..77383a5 --- /dev/null +++ b/pytensor_ml/dispatch/mlx/lbfgs.py @@ -0,0 +1,51 @@ +import mlx.core as mx + +from pytensor.link.mlx.dispatch import mlx_funcify + +from pytensor_ml.optim.lbfgs import LBFGSDirection + + +@mlx_funcify.register(LBFGSDirection) +def mlx_funcify_LBFGSDirection(op, node=None, **kwargs): + """Run the two-loop recursion as a Python loop over ``mx`` ops, since mlx has no scan.""" + n, m = op.n_parameters, op.memory_size + + def rows(stacks, slot): + # `count` is traced under mx.compile, so the slot is an mx scalar and the row is gathered rather + # than indexed from Python. + return [mx.take(stack, slot, axis=0) for stack in stacks] + + def dot(left, right): + # Vector matmul is the fastest dot mlx has from 0.32.2 (ml-explore/mlx#3580); before that it ran + # one threadgroup and was slower than a fused reduction by two orders of magnitude. + return sum(a.reshape(-1) @ b.reshape(-1) for a, b in zip(left, right)) + + def direction(count, gamma, rho, *tensors): + gradients = tensors[:n] + S = tensors[n : 2 * n] + Y = tensors[2 * n :] + + # Every row is gathered once, in ring order (oldest first), and reused by both loops. + order = [(count + offset) % m for offset in range(m)] + s_rows = [rows(S, slot) for slot in order] + y_rows = [rows(Y, slot) for slot in order] + curvatures = [mx.take(rho, slot) for slot in order] + + vector = list(gradients) + alphas = [None] * m + for position in reversed(range(m)): + alphas[position] = curvatures[position] * dot(s_rows[position], vector) + vector = [ + v - alphas[position].astype(v.dtype) * y_p + for v, y_p in zip(vector, y_rows[position]) + ] + vector = [gamma.astype(v.dtype) * v for v in vector] + for position in range(m): + beta = curvatures[position] * dot(y_rows[position], vector) + vector = [ + v + (alphas[position] - beta).astype(v.dtype) * s_p + for v, s_p in zip(vector, s_rows[position]) + ] + return vector[0] if n == 1 else tuple(vector) + + return direction diff --git a/pytensor_ml/optim/__init__.py b/pytensor_ml/optim/__init__.py index a86eeb0..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, @@ -41,6 +42,7 @@ adam_updates, adamax_updates, adamw_updates, + lbfgs_updates, nadam_updates, rmsprop_updates, rprop_updates, @@ -96,6 +98,8 @@ "get_gradients", "join_schedules", "large_step", + "lbfgs", + "lbfgs_updates", "linear_onecycle_schedule", "linear_schedule", "nadam", diff --git a/pytensor_ml/optim/alias.py b/pytensor_ml/optim/alias.py index 8a860ea..335495b 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,56 @@ def rule( return rule +def lbfgs( + learning_rate: LearningRate = 1.0, + memory_size: int = 10, + scale_init_precond: bool = True, + *, + namespace: str = "lbfgs", +) -> Transform: + """ + L-BFGS optimizer. See :func:`~pytensor_ml.optim.rules.lbfgs_updates` for the update rule. + + ``learning_rate`` accepts a float, a scalar shared variable, any scalar graph, or a schedule, and + ``namespace`` prefixes the state this rule allocates; see :func:`sgd`. + + Examples + -------- + A quasi-Newton direction from a memory of recent parameter and gradient differences, taken at a + fixed fraction with no line search. It reads the change between consecutive gradients as curvature, + so the loss has to be the same function from one step to the next: full batch, no dropout. For the + same reason it takes the loss's own gradients: put a clip after it in a chain, never ahead of it. + + .. code-block:: python + + import numpy as np + + from pytensor_ml.layers import Input, Linear + from pytensor_ml.loss import SquaredError, supervised_loss + from pytensor_ml.optim import compile_train, lbfgs + + X = Input("X", shape=(None, 4)) + loss, target = supervised_loss(Linear("fc", n_in=4, n_out=1)(X), SquaredError()) + + step = compile_train(loss, lbfgs(learning_rate=0.5, memory_size=10)) + loss_value = step(np.zeros((8, 4)), np.zeros((8, 1))) + """ + + def rule( + loss_gradients_or_updates: LossGradientsOrUpdates, parameters: Sequence[Parameter] + ) -> Updates: + return lbfgs_updates( + loss_gradients_or_updates, + parameters, + learning_rate=learning_rate, + memory_size=memory_size, + scale_init_precond=scale_init_precond, + namespace=namespace, + ) + + return rule + + def rmsprop( learning_rate: LearningRate = 1e-2, rho: float = 0.9, diff --git a/pytensor_ml/optim/base.py b/pytensor_ml/optim/base.py index 3d2c963..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. """ @@ -481,9 +483,11 @@ def _unreachable_parameter_names( ] -def state_for(parameter: Parameter, slot: str, fill_value: float = 0.0) -> Parameter: +def state_for( + parameter: Parameter, slot: str, fill_value: float = 0.0, history_size: int | None = None +) -> Parameter: """ - Return the optimizer-state shared variable shaped and typed like ``parameter``. + Return the optimizer-state shared variable typed like ``parameter``, or a stack of them. The variable is named ``"{parameter.name}/{slot}"`` and carries the parameter's layer, so a checkpoint numbers it where it numbers the parameter. The name is never used to *find* the variable at runtime -- @@ -501,6 +505,9 @@ def state_for(parameter: Parameter, slot: str, fill_value: float = 0.0) -> Param A short role tag for the slot, e.g. ``"adam/first_moment"`` or ``"trace/velocity"``. fill_value : float Constant to initialize the state with. Default 0.0. + history_size : int, optional + Number of past values to stack along a new leading axis, so the state is shaped + ``(history_size, *parameter.shape)``. Omitted, the state has the parameter's own shape. Returns ------- @@ -527,9 +534,21 @@ def state_for(parameter: Parameter, slot: str, fill_value: float = 0.0) -> Param f"Cannot allocate optimizer state {slot!r} for an unnamed parameter. Stateful optimizers rely on " "parameter names to identify their state at serialization boundaries; give the parameter a name." ) + if history_size is not None and history_size < 1: + raise ValueError(f"history_size must be at least 1, got {history_size}.") value = parameter.get_value(borrow=True) - state = pytensor.shared(np.full_like(value, fill_value), name=f"{parameter.name}/{slot}") + shape = value.shape if history_size is None else (history_size, *value.shape) + static_shape = ( + parameter.type.shape if history_size is None else (history_size, *parameter.type.shape) + ) + # The declared dtype rather than the value's: after a step on mlx the value is a device array + # whose dtype numpy cannot read. + state = pytensor.shared( + np.full(shape, fill_value, dtype=parameter.type.dtype), + name=f"{parameter.name}/{slot}", + shape=static_shape, + ) # Keeps `Linear_1_W` and `Linear_1_W/adam/first_moment` numbered onto the same layer. state.layer_name = getattr(parameter, "layer_name", None) return state diff --git a/pytensor_ml/optim/lbfgs.py b/pytensor_ml/optim/lbfgs.py new file mode 100644 index 0000000..2d49b78 --- /dev/null +++ b/pytensor_ml/optim/lbfgs.py @@ -0,0 +1,179 @@ +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.scalar import upcast +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, 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. ``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`. 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. The + slot count of ``rho`` and the stacks has to be static: + + .. code-block:: python + + import pytensor + import pytensor.tensor as pt + + from pytensor_ml.optim.lbfgs import LBFGSDirection + + g = pt.vector("g") + 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) + + 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, rho, *raw = inputs + tensors = [pt.as_tensor_variable(tensor) for tensor in raw] + # 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, curvature_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, rho, *tensors = inputs + if len(tensors) != 3 * n: + raise ValueError( + 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)}." + ) + # 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, with a static length of " + f"memory_size={m}, but got {rho.type}." + ) + 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 + + def right_product(slot, *vector): + s = [stack[slot] for stack in S] + y = [stack[slot] for stack in Y] + 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( + right_product, + sequences=[order], + outputs_info=[*gradients, None], + go_backwards=True, + return_updates=False, + ) + 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 = 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. + 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 _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: + """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 != 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, with the slot count static, but got " + f"{stack.type} for a gradient of type {gradient.type}." + ) + + +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)]) diff --git a/pytensor_ml/optim/rules.py b/pytensor_ml/optim/rules.py index c6a6b60..bb922b6 100644 --- a/pytensor_ml/optim/rules.py +++ b/pytensor_ml/optim/rules.py @@ -1,10 +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 ( @@ -16,9 +19,11 @@ 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 +905,187 @@ 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: LearningRate = 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. + 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 + 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, 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) + 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] + 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 + ] + # 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. + 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) + # 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 + # 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) + ] + # 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) + + updates: Updates = Steps(incoming) + if scale_init_precond: + # 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( + 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 + + directions = LBFGSDirection(n_parameters=len(parameters), memory_size=memory_size)( + new_pairs_written, + identity_scale, + new_curvatures, + *gradients, + *new_value_memory, + *new_gradient_memory, + return_list=True, + ) + + 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] + 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 new file mode 100644 index 0000000..3ce5f86 --- /dev/null +++ b/tests/dispatch/mlx/test_lbfgs.py @@ -0,0 +1,137 @@ +import numpy as np +import pytensor +import pytensor.tensor as pt +import pytest + +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, mlx_mode +from tests.optim.lbfgs_reference import dense_inverse_hessian, ring_stacks + +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, 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) + ] + + 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, rho, *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 = 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 + ) + + +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, np.zeros(3), g_in, S_in, Y_in) + + 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, [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) + + 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 + # (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/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_composition.py b/tests/optim/test_composition.py index d58032a..d38734a 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,37 @@ 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_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.""" diff --git a/tests/optim/test_lbfgs.py b/tests/optim/test_lbfgs.py new file mode 100644 index 0000000..e912566 --- /dev/null +++ b/tests/optim/test_lbfgs.py @@ -0,0 +1,163 @@ +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 +from tests.optim.lbfgs_reference import dense_inverse_hessian, ring_stacks + +floatX = pytensor.config.floatX +RTOL = 1e-6 if floatX == "float64" else 1e-4 + + +@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 + # 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) + 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) + 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, 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, rho, *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)]) + want = dense_inverse_hessian(gamma, pairs, size) @ gradient + 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, 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") + + +def test_a_scalar_parameter_has_vector_stacks(): + g = pt.scalar("g", 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) + + # 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(): + # 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, 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(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(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(3), np.ones(2), np.ones((4, 2)), np.ones((3, 2))), + "memory_size=3", + ), + ( + {"n_parameters": 1, "memory_size": 3}, + (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(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", + ), + ( + {"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", + "no_memory", + "extra_tensor", + "wrong_slots", + "wrong_rank", + "wrong_dtype", + "wrong_rho_length", + "dynamic_slots", + ], +) +def test_malformed_inputs_are_refused_at_build_time(props, tensors, message): + with pytest.raises(ValueError, match=message): + LBFGSDirection(**props)(0, 1.0, *tensors) diff --git a/tests/optim/test_rules.py b/tests/optim/test_rules.py index 9ff5c00..3303be5 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, @@ -21,6 +22,8 @@ adamw_updates, compile_train, cosine_schedule, + lbfgs, + lbfgs_updates, nadam, nadam_updates, rmsprop, @@ -32,6 +35,7 @@ ) from pytensor_ml.optim import alias as alias_module from pytensor_ml.pytensorf import function +from tests.optim.lbfgs_reference import dense_inverse_hessian floatX = pytensor.config.floatX @@ -62,6 +66,7 @@ def trainable(value, name=None, **kwargs): nadam(learning_rate=1e-2), adamax(learning_rate=1e-2), rprop(learning_rate=1e-2), + lbfgs(), ], ids=[ "sgd", @@ -79,6 +84,7 @@ def trainable(value, name=None, **kwargs): "nadam", "adamax", "rprop", + "lbfgs", ], ) def test_rule_reduces_loss(run_training, rule): @@ -97,8 +103,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 @@ -260,6 +267,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 +558,234 @@ 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) + + +@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. 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]) + 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 + 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 + + 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_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() + 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. 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) + + 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_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_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 + 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_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_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"): + 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 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.