diff --git a/deepmd/pt_expt/model/make_model.py b/deepmd/pt_expt/model/make_model.py index c61a5edca0..6d2dedfb80 100644 --- a/deepmd/pt_expt/model/make_model.py +++ b/deepmd/pt_expt/model/make_model.py @@ -1,4 +1,5 @@ # SPDX-License-Identifier: LGPL-3.0-or-later +import logging import math import types from typing import ( @@ -24,6 +25,9 @@ compact_nodes, expand_node_values, ) +from deepmd.pt.utils.auto_batch_size import ( + AutoBatchSize, +) from deepmd.pt_expt.common import ( auto_wrapped_class, torch_module, @@ -32,6 +36,11 @@ fused_energy_force_enabled, fused_operators_enabled, ) +from deepmd.pt_expt.utils.env import ( + DP_HESSIAN_HVP_BATCH, + DP_HESSIAN_HVP_BATCH_CAP, + DP_HESSIAN_HVP_MEMORY_FRACTION, +) from deepmd.pt_expt.utils.graph_builder import ( build_neighbor_graph_for_method, build_ragged_neighbor_graph, @@ -376,6 +385,228 @@ def __call__(self, coord_flat: torch.Tensor) -> torch.Tensor: return atom_out.sum(dim=0).reshape(-1)[self.ci] +log = logging.getLogger(__name__) + + +def _hessian_graph_batched_hvp( + model: Any, + kk: str, + ci: int, + nloc: int, + coord_flat: torch.Tensor, + atype: torch.Tensor, + box: torch.Tensor | None, + method: str, + pair_excl: Any, + rcut: float, + fparam: torch.Tensor | None, + aparam: torch.Tensor | None, + spin: torch.Tensor | None, + charge_spin: torch.Tensor | None, + batch: int, + create_graph: bool, + max_rows: int | None = None, +) -> torch.Tensor: + """Hessian of one reduced-output component via batched Hessian-vector products. + + Exactly equivalent to ``torch.autograd.functional.hessian`` on + :class:`_WrapperForwardEnergyGraph`, but evaluates ``batch`` rows per + second-order backward instead of one. The structure is replicated ``batch`` + times along the frame axis the carry-all graph already has; frames are + independent, so the energy of the replicated system is a sum of independent + terms and its Hessian is block diagonal. One forward and one first-order + backward build a graph that every chunk of seed vectors then reuses. + + No ``vmap`` is involved, so custom autograd Functions that lack a batching + rule keep working, and no approximation is made: in float64 the result + matches the one-row-at-a-time path to 2.8e-14 relative RMS on the + DPA-4.0.1-Pro-MPtrj checkpoint (1e-16 on the smaller example model), against + a 6.4e-16 asymmetry in the unbatched path's own answer. In float32 the two + differ by ~3e-6 relative RMS, a few times the unbatched path's own ~5e-7 + asymmetry -- the longer summation chain, not a different computation. + + ``max_rows`` stops after that many leading rows. Only the memory probe + uses it: pricing one Hessian-vector product must not pay for the whole + Hessian first. + + Returns the ``(nloc * 3, nloc * 3)`` Hessian for output component ``ci`` + (or its first ``max_rows`` rows). + """ + ndof = nloc * 3 + nb = max(1, min(batch, ndof)) + if create_graph and coord_flat.requires_grad: + # Mirror ``torch.autograd.functional._grad_preprocess``: with + # ``create_graph`` the caller means to differentiate the Hessian + # itself, so the path back to the coordinates has to survive. + # Detaching would silently cut it and leave only the parameter path. + x = coord_flat.reshape(1, ndof).expand(nb, ndof).contiguous() + else: + x = coord_flat.detach().reshape(1, ndof).expand(nb, ndof).contiguous() + x = x.requires_grad_(True) + atype_b = atype.reshape(1, nloc).expand(nb, nloc).contiguous() + box_b = ( + box.reshape(1, -1).expand(nb, box.numel()).contiguous() + if box is not None + else None + ) + graph = build_neighbor_graph_for_method( + method, x.reshape(nb, nloc, 3), atype_b, box_b, rcut, pair_excl + ) + atomic_ret = model.atomic_model.forward_common_atomic_graph( + graph, + atype_b.reshape(-1), + fparam=fparam.expand(nb, -1).contiguous() if fparam is not None else None, + aparam=aparam.repeat(nb, 1) if aparam is not None else None, + spin=spin.repeat(nb, 1) if spin is not None else None, + charge_spin=charge_spin.expand(nb, -1).contiguous() + if charge_spin is not None + else None, + ) + # flat (nb * nloc, *def) -> one scalar per replica, summed: the replicas are + # independent, so d/dx_b only sees replica b. + total = atomic_ret[kk].reshape(nb, nloc, -1)[..., ci].sum() + wanted = ndof if max_rows is None else min(max_rows, ndof) + if not total.requires_grad: + # The reduced output does not depend on the coordinates at all, so both + # derivatives are zero. autograd refuses to differentiate an output + # that carries no graph, so this has to be caught before the first + # ``grad`` rather than after it. + return x.new_zeros(wanted, ndof) + + (grad,) = torch.autograd.grad( + total, x, create_graph=True, allow_unused=True, materialize_grads=True + ) + + if not grad.requires_grad: + # The output is linear in the coordinates: the first derivative exists + # but is constant, so every second derivative is zero. + # ``functional.hessian`` materialises that zero block under its default + # ``strict=False``; differentiating a constant ``grad`` again would + # instead raise, turning a legitimate model into a crash. + return x.new_zeros(wanted, ndof) + + rows: list[torch.Tensor] = [] + for start in range(0, wanted, nb): + stop = min(start + nb, wanted) + # One seed block per chunk rather than one ``ndof x ndof`` identity: + # the identity is the very allocation the batching exists to avoid, and + # a fresh block per chunk keeps ``create_graph`` from retaining a buffer + # that a later chunk would overwrite. + seeds = torch.zeros(nb, ndof, dtype=x.dtype, device=x.device) + seeds[ + torch.arange(stop - start, dtype=torch.int64, device=x.device), + torch.arange(start, stop, dtype=torch.int64, device=x.device), + ] = 1 + (hvp,) = torch.autograd.grad( + grad, + x, + grad_outputs=seeds, + retain_graph=True, + create_graph=create_graph, + allow_unused=True, + materialize_grads=True, + ) + rows.append(hvp) + # Padding rows carry a zero seed and sit at the end of the final chunk, so + # one slice drops both them and anything past ``max_rows``. + return torch.cat(rows)[:wanted] + + +def _hvp_replica_cost(device: torch.device, probe: Any) -> int: + """Peak memory, in bytes, that one Hessian-vector product costs.""" + torch.cuda.synchronize(device) + torch.cuda.reset_peak_memory_stats(device) + base = torch.cuda.memory_allocated(device) + probe() + torch.cuda.synchronize(device) + return max(torch.cuda.max_memory_allocated(device) - base, 1) + + +def _auto_hvp_batch(device: torch.device, probe: Any) -> int: + """Pick a batch from what one Hessian-vector product costs and what is free. + + Peak memory is linear in the batch, so the batch that fits is the one worth + taking: past it the run dies, and below it the device is idle. Pricing a + single product costs one row out of ``3 * nloc`` and needs no per-model + constants, which a fitted memory model would. + + The measured cost covers a whole product, while each batch step beyond the + first adds only its marginal share, so this reads high and the batch comes + out conservative -- the direction to err, since + :func:`_hessian_graph_row_block` can recover from a batch that turns out too + large but nothing recovers the time lost to one that was too small. + """ + if device.type != "cuda": + # Without allocator introspection there is nothing to size against, and + # without a recoverable out-of-memory error nothing to catch if the + # guess is wrong. Stay on the one-row-at-a-time path. + return 1 + cost = _hvp_replica_cost(device, probe) + free, _total = torch.cuda.mem_get_info(device) + # Blocks the caching allocator holds but is not using are free to us even + # though the driver counts them as taken. + reusable = torch.cuda.memory_reserved(device) - torch.cuda.memory_allocated(device) + budget = (free + reusable) * DP_HESSIAN_HVP_MEMORY_FRACTION + return max(1, min(int(budget // cost), DP_HESSIAN_HVP_BATCH_CAP)) + + +def _hessian_graph_row_block( + batch: int, + wrapper: Any, + coord_flat: torch.Tensor, + create_graph: bool, + **kwargs: Any, +) -> tuple[torch.Tensor, int]: + """One output component's Hessian, halving the batch if memory runs out. + + Batching changes how the Hessian is computed, never what it is, so a batch + that does not fit can simply be retried smaller instead of ending the run. + Reaching 1 hands over to the original one-row-at-a-time path rather than to + this helper with a batch of one, so the fallback bottoms out in exactly the + code a user who set 1 would have taken. + + The out-of-memory test is the repository's ``is_oom_error``: the allocator + failure does not always arrive as ``torch.OutOfMemoryError`` -- AOTInductor + rewraps it in a plain ``RuntimeError`` -- and a catch keyed on the + exception type alone lets the wrapped form end the run. + + Returns the Hessian and the batch that fit, so the caller starts the next + component there instead of re-climbing the ladder from the top: each failed + attempt is a full forward plus a first-order backward, spent in exactly the + regime the fallback exists for. + """ + while batch > 1: + try: + return ( + _hessian_graph_batched_hvp( + coord_flat=coord_flat, + batch=batch, + create_graph=create_graph, + **kwargs, + ), + batch, + ) + except Exception as e: + if not AutoBatchSize(silent=True).is_oom_error(e): + raise + # Round up, so 1 is reached only from 2: rounding down drops 3 + # straight to 1 and skips the batch of two that may well fit. + batch = (batch + 1) // 2 + log.warning( + "Hessian-vector product batch did not fit in memory; retrying " + "with %d. Set DP_HESSIAN_HVP_BATCH to choose it yourself.", + batch, + ) + return ( + torch.autograd.functional.hessian( + wrapper, + coord_flat, + create_graph=create_graph, + ), + 1, + ) + + def _cal_hessian_ext_graph( model: Any, kk: str, @@ -413,6 +644,19 @@ def _cal_hessian_ext_graph( if charge_spin is not None and charge_spin.ndim == 1 else charge_spin ) + # Priced per frame: within one frame every output component shares the + # neighbour count, but each frame has its own, so a batch priced on one + # frame does not price another. A price is an estimate, though, and an + # out-of-memory error is a measurement: once a batch has been refused, + # later frames of the same size are priced no higher than the batch that + # survived, rather than walking down again from a price the device has + # already refused -- and once a single row is all that survived, they are + # not priced at all. The size is the real-atom count, the only one known + # before the graph is built; frames of another size keep their own price, + # so one frame's batch is still not spent on another. + hvp_batch = DP_HESSIAN_HVP_BATCH + auto_batch = hvp_batch is None + ceilings: dict[int, int] = {} hessians = [] for ii in range(nf): node_index = torch.nonzero(atype[ii] >= 0, as_tuple=False).reshape(-1) @@ -427,6 +671,55 @@ def _cal_hessian_ext_graph( if charge_spin_by_frame is not None: frame_index = 0 if charge_spin_by_frame.shape[0] == 1 else ii charge_spin_frame = charge_spin_by_frame[frame_index : frame_index + 1] + hvp_kwargs = { + "model": model, + "kk": kk, + "nloc": n_real, + "atype": atype_frame, + "box": box[ii : ii + 1] if box is not None else None, + "method": method, + "pair_excl": pair_excl, + "rcut": rcut, + "fparam": fparam[ii : ii + 1] if fparam is not None else None, + "aparam": aparam_frame, + "spin": spin_frame, + "charge_spin": charge_spin_frame, + } + ceiling = ceilings.get(n_real) + if auto_batch and n_real and ceiling == 1: + # One row is all that fit at this size; no price can change that. + hvp_batch = 1 + elif auto_batch and n_real: + # Which component the probe differentiates does not change the + # graph, so the pricing probe fixes ci=0. + try: + hvp_batch = _auto_hvp_batch( + coord.device, + lambda: _hessian_graph_batched_hvp( + coord_flat=coord_flat, + batch=1, + create_graph=create_graph, + max_rows=1, + ci=0, + **hvp_kwargs, + ), + ) + except Exception as e: + # Pricing one product is itself a product, so it can be the + # allocation that does not fit. Letting that escape would + # end the run before the one-row-at-a-time path -- which + # might well have fit -- was ever tried. + if not AutoBatchSize(silent=True).is_oom_error(e): + raise + log.warning( + "Ran out of memory measuring the Hessian-vector " + "product; falling back to one row at a time. Set " + "DP_HESSIAN_HVP_BATCH to choose the batch yourself." + ) + hvp_batch = ceilings[n_real] = 1 + if ceiling is not None: + hvp_batch = min(hvp_batch, ceiling) + log.debug("Hessian-vector products batched %d rows at a time", hvp_batch) for ci in range(vsize): wrapper = _WrapperForwardEnergyGraph( model=model, @@ -443,11 +736,17 @@ def _cal_hessian_ext_graph( spin=spin_frame, charge_spin=charge_spin_frame, ) - hess = torch.autograd.functional.hessian( - wrapper, - coord_flat, + batch = hvp_batch if hvp_batch is not None else 1 + hess, hvp_batch = _hessian_graph_row_block( + batch=batch, + wrapper=wrapper, + coord_flat=coord_flat, create_graph=create_graph, + ci=ci, + **hvp_kwargs, ) # (n_real*3, n_real*3) + if hvp_batch < batch: + ceilings[n_real] = hvp_batch if n_real != nloc: component_index = ( node_index[:, None] * 3 diff --git a/deepmd/pt_expt/utils/env.py b/deepmd/pt_expt/utils/env.py index 0f4d38ba84..fbbc962c86 100644 --- a/deepmd/pt_expt/utils/env.py +++ b/deepmd/pt_expt/utils/env.py @@ -30,6 +30,40 @@ SAMPLER_RECORD = os.environ.get("SAMPLER_RECORD", False) DP_DTYPE_PROMOTION_STRICT = os.environ.get("DP_DTYPE_PROMOTION_STRICT", "0") == "1" +# Number of Hessian rows evaluated per second-order backward pass. The Hessian +# is built from Hessian-vector products; batching them over replicated frames +# trades memory for far fewer kernel launches. 1 keeps the one-row-at-a-time +# path and reproduces the pre-batching behaviour exactly. +# +# Left unset, the batch is chosen per frame: one Hessian-vector product is run to +# measure what a replica costs, and the batch is what the free memory affords, +# clamped to [1, DP_HESSIAN_HVP_BATCH_CAP]. That is not a tuning preference but +# a correctness matter, because peak memory is linear in this value while the +# speedup is not. Measured on one H20 with DPA-4.0.1-Pro-MPtrj in eval mode, +# float32, TF32 off: +# +# peak(MiB) = 782 + [4.50 + 3.27*(B-1)]*edges + [10.9 + 7.96*(B-1)]*natoms +# +# so at fcc-solid density a 96 GiB card holds ~381 atoms at B=1 but only ~63 at +# B=8, while the speedup falls from 5.06x at 72 edges to 1.24x at 5832 edges -- +# batching recovers kernel-launch overhead, which stops mattering once a single +# Hessian-vector product already saturates the device. A fixed large value +# therefore costs the size ceiling and buys nothing on the systems that need it. +# +# An explicit value is honoured as given, including above the cap. The +# out-of-memory fallback still applies to it: halving the batch changes how the +# Hessian is computed, never what it is, so a run that would have died is +# finished instead, with a warning naming the batch actually used. +_hessian_hvp_batch = os.environ.get("DP_HESSIAN_HVP_BATCH") +DP_HESSIAN_HVP_BATCH: int | None = ( + int(_hessian_hvp_batch) if _hessian_hvp_batch is not None else None +) +# Ceiling for the automatic choice. Past this the speedup has flattened on every +# system measured, so more batch would only cost memory. +DP_HESSIAN_HVP_BATCH_CAP = 8 +# Share of the free memory the automatic choice plans for. The rest absorbs the +# gap between one replica's measured cost and the marginal cost of the next. +DP_HESSIAN_HVP_MEMORY_FRACTION = 0.5 try: # only linux ncpus = len(os.sched_getaffinity(0)) diff --git a/doc/env.md b/doc/env.md index ba08490a06..586c610c69 100644 --- a/doc/env.md +++ b/doc/env.md @@ -70,6 +70,39 @@ available device memory, leaving a 10% margin. This policy applies to both Other GPU allocators and backends grow batches until an out-of-memory error. ::: +:::{envvar} DP_HESSIAN_HVP_BATCH + +**Default**: automatically sized on CUDA devices; `1` elsewhere + +{{ pytorch_icon }} Number of Hessian rows evaluated per second-order backward +pass when a PyTorch Exportable model computes a Hessian on the neighbour-graph +route. The Hessian is assembled from Hessian-vector products, +and batching them over replicated frames trades peak memory for far fewer +kernel launches. `1` evaluates one row at a time, which is how the Hessian was +computed before batching existed. + +Automatic sizing runs one Hessian-vector product per frame to measure what a +replica costs there -- neighbour counts differ between frames, so one +measurement does not price them all -- then takes the batch the free memory +affords, capped at 8. If a batch -- or the measurement itself -- then runs out +of memory, the batch that survived caps every later frame with as many real atoms +in the same call, so a device that is fuller than the measurement suggested is +not run out of memory again frame after frame. Peak memory +is linear in the batch while the speedup is not: batching recovers +kernel-launch overhead, which stops mattering once a single Hessian-vector +product already saturates the device. On one H20 with DPA-4, batching is worth +about 5x on a system of 72 neighbour pairs but only 1.2x at 5832, where it also +costs six times the memory -- so a large fixed batch lowers the system size that +fits and buys little on the systems that need it. + +Setting this variable disables automatic sizing and uses the value given; +out-of-memory errors can still reduce the batch, halving it (rounding up, so 2 +is always tried before 1), warning which +batch was used, and keeping the surviving batch for the rest of the call. The +result does not depend on the batch: it changes how the Hessian is computed, +not what it is. +::: + :::{envvar} DP_BACKEND **Default**: `tensorflow` diff --git a/source/tests/pt_expt/model/test_hessian_hvp_batch.py b/source/tests/pt_expt/model/test_hessian_hvp_batch.py new file mode 100644 index 0000000000..ec495918f2 --- /dev/null +++ b/source/tests/pt_expt/model/test_hessian_hvp_batch.py @@ -0,0 +1,784 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""``DP_HESSIAN_HVP_BATCH`` must not change the Hessian, only how it is computed. + +The graph route assembles the Hessian from Hessian-vector products and can +evaluate several rows per second-order backward by replicating the structure +along the frame axis. Frames are independent, so the replicated energy is a +sum of independent terms and its Hessian is block diagonal -- the batched +result is exact, and these tests pin that down in float64. + +``test_dpa2_graph_lower`` already reaches the batched helper, but only because +the batch it happens to get exceeds 1: a batch of 1 would turn that coverage +into coverage of the unbatched path alone, silently. Here the batch is the +object under test, and a counter asserts which branch ran so the coverage +cannot drift away again. + +``DP_HESSIAN_HVP_BATCH`` unset means "choose per frame from free memory", so +the choice itself and the out-of-memory fallback are tested too. +""" + +import pytest +import torch + +import deepmd.pt_expt.model.make_model as mm +from deepmd.pt.utils import ( + env, +) +from deepmd.pt_expt.descriptor.dpa1 import ( + DescrptDPA1, +) +from deepmd.pt_expt.fitting import ( + InvarFitting, +) +from deepmd.pt_expt.model import ( + EnergyModel, +) +from deepmd.pt_expt.model.graph_lower import ( + model_uses_graph_lower, +) + +from ...seed import ( + GLOBAL_SEED, +) + +NATOMS = 5 +NDOF = 3 * NATOMS # 15: odd, so most batch sizes leave a partial final chunk +RCUT = 4.0 +RCUT_SMTH = 0.5 +SEL = 20 # mixed-type single-int sel +NT = 2 + +# 15 % 2 == 1 and 15 % 8 == 7 exercise the zero-padded final chunk; 15 % 3 == 0 +# divides evenly; 16 exceeds NDOF and must be clamped back to a single chunk. +BATCHES = [2, 3, 4, 8, 16] +PADS = {b for b in BATCHES if NDOF % min(b, NDOF)} + + +@pytest.fixture +def route_counts(monkeypatch): + """Count which Hessian implementation each forward actually took.""" + counts = {"batched": 0, "graph": 0, "dense": 0} + originals = { + "batched": mm._hessian_graph_batched_hvp, + "graph": mm._cal_hessian_ext_graph, + "dense": mm._cal_hessian_ext, + } + names = { + "batched": "_hessian_graph_batched_hvp", + "graph": "_cal_hessian_ext_graph", + "dense": "_cal_hessian_ext", + } + + def make(key): + original = originals[key] + + def counted(*args, **kwargs): + counts[key] += 1 + return original(*args, **kwargs) + + return counted + + for key, name in names.items(): + monkeypatch.setattr(mm, name, make(key)) + return counts + + +class TestHessianHvpBatch: + """Batched and unbatched Hessian-vector products must agree exactly.""" + + def setup_method(self) -> None: + self.device = env.DEVICE + generator = torch.Generator(device=self.device).manual_seed(GLOBAL_SEED) + cell = torch.rand( + [3, 3], dtype=torch.float64, device=self.device, generator=generator + ) + cell = (cell + cell.T) + 5.0 * torch.eye( + 3, device=self.device, dtype=torch.float64 + ) + self.box = cell.reshape(1, 9) + coord = torch.rand( + [NATOMS, 3], + dtype=torch.float64, + device=self.device, + generator=generator, + ) + self.coord = (coord @ cell).unsqueeze(0) + self.atype = torch.tensor( + [[0, 0, 0, 1, 1]], dtype=torch.int64, device=self.device + ) + + def _make_model(self, graph: bool = True) -> EnergyModel: + ds = DescrptDPA1( + RCUT, + RCUT_SMTH, + SEL, + NT, + neuron=[3, 6], + axis_neuron=2, + attn=4, + attn_layer=0, + attn_dotr=True, + attn_mask=False, + # Smooth attention keeps sel-padding in the dense softmax + # denominator, which the carry-all graph omits; exact graph-vs-dense + # parity needs it off. + smooth_type_embedding=False, + activation_function="tanh", + set_davg_zero=False, + type_one_side=True, + precision="float64", + seed=GLOBAL_SEED, + ).to(self.device) + ft = InvarFitting( + "energy", + NT, + ds.get_dim_out(), + 1, + mixed_types=ds.mixed_types(), + precision="float64", + seed=GLOBAL_SEED, + ).to(self.device) + model = EnergyModel(ds, ft, type_map=["foo", "bar"]).to(self.device) + model.eval() + if not graph: + model.atomic_model.descriptor.disable_graph_lower() + model.enable_hessian() + assert model_uses_graph_lower(model) is graph + return model + + def _hessian(self, model: EnergyModel) -> torch.Tensor: + out = model.forward( + self.coord.clone().requires_grad_(True), self.atype, box=self.box + ) + return out["hessian"].reshape(NDOF, NDOF) + + def test_batch_one_takes_the_unbatched_path( + self, route_counts, monkeypatch + ) -> None: + """1 (and 0) must reach the original one-row-at-a-time implementation.""" + model = self._make_model() + for batch in (0, 1): + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", batch) + self._hessian(model) + assert route_counts["graph"] == 2 + assert route_counts["batched"] == 0, "batch<=1 must not take the batched helper" + assert route_counts["dense"] == 0 + + @pytest.mark.parametrize("batch", BATCHES) + def test_batched_matches_unbatched(self, batch, route_counts, monkeypatch) -> None: + """Every batch size must reproduce the unbatched Hessian to float64 precision.""" + model = self._make_model() + + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 1) + reference = self._hessian(model) + assert route_counts["batched"] == 0 + + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", batch) + batched = self._hessian(model) + assert route_counts["batched"] == 1, "the batched helper did not run" + assert route_counts["dense"] == 0 + + scale = reference.abs().max() + torch.testing.assert_close( + batched, reference, rtol=0.0, atol=float(1e-12 * scale) + ) + + @pytest.mark.parametrize("batch", sorted(PADS)) + def test_padded_final_chunk_is_discarded(self, batch, monkeypatch) -> None: + """A batch size that does not divide 3*nloc still yields exactly 3*nloc rows. + + The final chunk is zero-padded to keep the retained graph's shape; those + rows must be dropped, not returned. + """ + assert NDOF % batch, f"{batch} divides {NDOF}; it exercises no padding" + model = self._make_model() + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", batch) + hessian = self._hessian(model) + assert hessian.shape == (NDOF, NDOF) + # A dropped-row bug shows up as an all-zero row, and a padding leak as a + # zero row in the middle; neither is possible for a real Hessian here. + assert (hessian.abs().sum(dim=1) > 0).all() + + @pytest.mark.skipif( + not torch.cuda.is_available(), reason="the automatic choice needs CUDA" + ) + def test_probe_prices_one_product_not_the_whole_hessian(self, monkeypatch) -> None: + """The automatic choice must not pay for a Hessian to decide the batch. + + ``max_rows`` is what keeps the probe to a single Hessian-vector product; + ignoring it would still give the right batch, just after doing all the + work the batch was supposed to speed up. + """ + calls = [] + real = mm._hessian_graph_batched_hvp + + def record(*args, **kwargs): + out = real(*args, **kwargs) + calls.append((kwargs.get("max_rows"), out.shape[0])) + return out + + monkeypatch.setattr(mm, "_hessian_graph_batched_hvp", record) + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", None) + model = self._make_model() + self._hessian(model) + + assert calls, "the automatic choice never ran the probe" + probe_max_rows, probe_rows = calls[0] + assert probe_max_rows == 1 + assert probe_rows == 1, "the probe computed more than one row" + + def test_the_batch_is_repriced_for_each_frame(self, monkeypatch) -> None: + """One price does not fit all frames: neighbour counts differ per frame. + + Pricing once on the first frame spends a sparse frame's batch on a + dense one (recovered only by the OOM ladder) and a dense frame's batch + on a sparse one (never recovered: the device just idles). The probe + costs a single Hessian-vector product, so it runs once per frame; the + Hessians accumulated for earlier frames also shrink the free memory a + later price sees, which pricing once up front cannot know. + """ + prices = [] + real = mm._auto_hvp_batch + + def record(device, probe): + batch = real(device, probe) + prices.append(batch) + return batch + + monkeypatch.setattr(mm, "_auto_hvp_batch", record) + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", None) + model = self._make_model() + coord = torch.cat([self.coord, self.coord * 1.01], dim=0) + atype = torch.cat([self.atype, self.atype], dim=0) + box = torch.cat([self.box, self.box], dim=0) + out = model.forward(coord.clone().requires_grad_(True), atype, box=box) + batched = out["hessian"].reshape(2, NDOF, NDOF) + + assert len(prices) == 2, f"priced {len(prices)} times for 2 frames" + + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 1) + out = model.forward(coord.clone().requires_grad_(True), atype, box=box) + reference = out["hessian"].reshape(2, NDOF, NDOF) + scale = reference.abs().max() + torch.testing.assert_close( + batched, reference, rtol=0.0, atol=float(1e-12 * scale) + ) + assert len(prices) == 2, "a fixed batch must not price at all" + + def test_create_graph_keeps_the_path_to_the_coordinates(self, monkeypatch) -> None: + """``create_graph`` must leave the Hessian differentiable in the input. + + ``torch.autograd.functional.hessian`` keeps the input in the graph when + ``create_graph`` is set; detaching instead severs everything upstream of + the coordinates while leaving the parameter path intact, so the loss of + signal is silent. + """ + model = self._make_model() + model.train() + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 4) + upstream = self.coord.clone().requires_grad_(True) + # a non-leaf coordinate, which is what makes the severed path visible + out = model.forward(upstream * 1.0, self.atype, box=self.box) + hessian = out["hessian"].reshape(NDOF, NDOF) + assert hessian.requires_grad, "create_graph produced a detached Hessian" + (back,) = torch.autograd.grad( + hessian.sum(), upstream, retain_graph=True, allow_unused=True + ) + assert back is not None, "the Hessian no longer depends on the coordinates" + assert torch.isfinite(back).all() + + @pytest.mark.parametrize("dependence", ["linear", "constant"]) + def test_a_curvature_free_energy_gives_a_zero_hessian_not_a_crash( + self, dependence, monkeypatch + ) -> None: + """Constant *and* linear coordinate dependence must yield zeros. + + ``functional.hessian`` materialises the zero block under its default + ``strict=False``. The two shapes fail differently and need separate + guards: a linear output reaches the second derivative with a constant + first derivative, while a constant output carries no graph at all and + is refused by the *first* ``autograd.grad`` -- before any guard placed + after it can run. + """ + + class _CurvatureFreeAtomicModel: + """Energy with no second derivative in the coordinates.""" + + def __init__(self, dependence: str) -> None: + self.dependence = dependence + + def forward_common_atomic_graph(self, graph, atype_flat, **kwargs): + nb, nloc = graph.shape[0], graph.shape[1] + if self.dependence == "linear": + energy = (3.0 * graph).sum(-1) + else: # no dependence on the coordinates whatsoever + energy = torch.full( + (nb, nloc), 7.0, dtype=graph.dtype, device=graph.device + ) + return {"energy": energy.reshape(nb * nloc, 1)} + + class _Model: + atomic_model = _CurvatureFreeAtomicModel(dependence) + + monkeypatch.setattr( + mm, + "build_neighbor_graph_for_method", + lambda method, pos, atype, box, rcut, pair_excl: pos, + ) + nloc = NATOMS + coord_flat = self.coord.reshape(-1).clone() + hessian = mm._hessian_graph_batched_hvp( + model=_Model(), + kk="energy", + ci=0, + nloc=nloc, + coord_flat=coord_flat, + atype=self.atype, + box=self.box, + method="graph", + pair_excl=None, + rcut=RCUT, + fparam=None, + aparam=None, + spin=None, + charge_spin=None, + batch=4, + create_graph=False, + ) + assert hessian.shape == (NDOF, NDOF) + assert torch.count_nonzero(hessian) == 0, ( + f"a {dependence} energy has no curvature" + ) + + def test_the_probe_does_not_build_a_full_identity(self, monkeypatch) -> None: + """Pricing one product must not allocate the ``ndof x ndof`` identity. + + The identity is the allocation batching exists to avoid; building it to + decide the batch can itself be what runs the device out of memory, and + it inflates the very cost the probe is measuring. + """ + seen = [] + real_eye = torch.eye + + def record_eye(n, *args, **kwargs): + seen.append(n) + return real_eye(n, *args, **kwargs) + + monkeypatch.setattr(torch, "eye", record_eye) + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 4) + model = self._make_model() + self._hessian(model) + assert NDOF not in seen, ( + f"an identity of size {NDOF} was built; sizes seen: {seen}" + ) + + @pytest.mark.parametrize("wrapped", [False, True]) + def test_an_out_of_memory_probe_falls_back_instead_of_escaping( + self, wrapped, monkeypatch + ) -> None: + """Pricing a product is a product, so it can be the thing that does not fit. + + Letting that escape ends the run before the one-row-at-a-time path -- + which may well have fit -- is ever tried. The OOM need not arrive as + ``torch.OutOfMemoryError``; the wrapped form must fall back too. + """ + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", None) + + def always_oom(device, probe): + if wrapped: + raise RuntimeError("CUDA out of memory. Tried to allocate 2.00 GiB") + raise torch.OutOfMemoryError("probe did not fit") + + monkeypatch.setattr(mm, "_auto_hvp_batch", always_oom) + model = self._make_model() + hessian = self._hessian(model) + assert hessian.shape == (NDOF, NDOF) + assert torch.isfinite(hessian).all() + + def test_dense_route_is_untouched(self, route_counts, monkeypatch) -> None: + """The batch size must not reach, or change, the dense Hessian route.""" + dense_model = self._make_model(graph=False) + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 8) + dense = self._hessian(dense_model) + assert route_counts["dense"] == 1 + assert route_counts["graph"] == 0 + assert route_counts["batched"] == 0 + + graph_model = self._make_model(graph=True) + graph_batched = self._hessian(graph_model) + assert route_counts["batched"] == 1 + + # An independent implementation agreeing to float64 precision is a + # stronger statement than batched-vs-unbatched alone: both graph paths + # share the same wrapper, the dense route does not. + scale = dense.abs().max() + torch.testing.assert_close( + graph_batched, dense, rtol=0.0, atol=float(1e-9 * scale) + ) + + def test_oom_halves_the_batch_and_keeps_the_answer( + self, route_counts, monkeypatch + ) -> None: + model = self._make_model() + + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 1) + reference = self._hessian(model) + + attempted = [] + real = mm._hessian_graph_batched_hvp + + def only_small_batches_fit(*args, **kwargs): + attempted.append(kwargs["batch"]) + if kwargs["batch"] > 2: + raise torch.OutOfMemoryError("simulated") + return real(*args, **kwargs) + + monkeypatch.setattr(mm, "_hessian_graph_batched_hvp", only_small_batches_fit) + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 8) + recovered = self._hessian(model) + + assert attempted == [8, 4, 2], attempted + scale = reference.abs().max() + torch.testing.assert_close( + recovered, reference, rtol=0.0, atol=float(1e-12 * scale) + ) + + @pytest.mark.parametrize( + ("start", "fits", "expected"), + [ + (7, 2, [7, 4, 2]), # rounding down would go 7, 3 and then give up + (3, 2, [3, 2]), # rounding down would drop 3 straight to 1 + (5, 3, [5, 3]), # 5 halves to 3, not 2 + ], + ) + def test_halving_rounds_up_so_two_is_tried_before_one( + self, start, fits, expected, monkeypatch + ) -> None: + """An odd batch that does not fit must not skip the batch below it. + + Rounding down sends 3 straight to the one-row-at-a-time path, so a + batch of 2 that would have fit is never tried -- measured on DPA-4, + the call then takes about twice as long as it does at 2. + """ + model = self._make_model() + + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 1) + reference = self._hessian(model) + + attempted = [] + real = mm._hessian_graph_batched_hvp + + def refuse_above_fits(*args, **kwargs): + attempted.append(kwargs["batch"]) + if kwargs["batch"] > fits: + raise torch.OutOfMemoryError("simulated") + return real(*args, **kwargs) + + monkeypatch.setattr(mm, "_hessian_graph_batched_hvp", refuse_above_fits) + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", start) + recovered = self._hessian(model) + + assert attempted == expected, attempted + scale = reference.abs().max() + torch.testing.assert_close( + recovered, reference, rtol=0.0, atol=float(1e-12 * scale) + ) + + @pytest.mark.parametrize("wrapped", ["message", "cause", "aoti"]) + def test_a_wrapped_out_of_memory_still_falls_back( + self, wrapped, monkeypatch + ) -> None: + """An OOM need not arrive as ``torch.OutOfMemoryError``. + + AOTInductor rewraps the allocator failure in a plain ``RuntimeError``: + sometimes the original text survives in the message, sometimes only in + the ``__cause__`` chain, and sometimes both are stripped behind its + ``run_func_`` signature. A catch keyed on the exception type alone + lets all three forms end the run; the fallback must recognise what the + repository's ``is_oom_error`` recognises. + """ + model = self._make_model() + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 1) + reference = self._hessian(model) + + attempted = [] + real = mm._hessian_graph_batched_hvp + + def oom_above_two(*args, **kwargs): + attempted.append(kwargs["batch"]) + if kwargs["batch"] > 2: + if wrapped == "message": + raise RuntimeError("CUDA out of memory. Tried to allocate 2.00 GiB") + if wrapped == "cause": + raise RuntimeError( + "the forward call failed" + ) from torch.cuda.OutOfMemoryError("CUDA out of memory.") + raise RuntimeError( + "run_func_(...) API call failed at " + "/tmp/model_container_runner.cpp:123" + ) + return real(*args, **kwargs) + + monkeypatch.setattr(mm, "_hessian_graph_batched_hvp", oom_above_two) + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 8) + recovered = self._hessian(model) + + assert attempted == [8, 4, 2], attempted + scale = reference.abs().max() + torch.testing.assert_close( + recovered, reference, rtol=0.0, atol=float(1e-12 * scale) + ) + + def test_an_unrelated_runtime_error_is_not_retried(self, monkeypatch) -> None: + """Catching wider than ``OutOfMemoryError`` must not swallow real errors.""" + model = self._make_model() + + def bogus(*args, **kwargs): + raise RuntimeError("some unrelated failure") + + monkeypatch.setattr(mm, "_hessian_graph_batched_hvp", bogus) + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 8) + with pytest.raises(RuntimeError, match="some unrelated failure"): + self._hessian(model) + + def test_oom_all_the_way_down_lands_on_the_unbatched_path( + self, route_counts, monkeypatch + ) -> None: + """When nothing fits, the fallback bottoms out in the original path.""" + model = self._make_model() + + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 1) + reference = self._hessian(model) + batched_before = route_counts["batched"] + + def nothing_fits(*args, **kwargs): + raise torch.OutOfMemoryError("simulated") + + monkeypatch.setattr(mm, "_hessian_graph_batched_hvp", nothing_fits) + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 8) + recovered = self._hessian(model) + + assert route_counts["batched"] == batched_before, ( + "the counter wraps the real helper, which was replaced" + ) + scale = reference.abs().max() + torch.testing.assert_close( + recovered, reference, rtol=0.0, atol=float(1e-12 * scale) + ) + + def test_the_surviving_batch_is_kept_for_the_next_frame(self, monkeypatch) -> None: + """A batch the ladder discovers must be reused, not re-discovered. + + When the halved batch stayed local to the helper, every frame restarted + the ladder from the top: two frames with only batches up to 2 fitting + attempt [8, 4, 2, 8, 4, 2], and each failed attempt pays a full forward + plus a first-order backward in exactly the regime the fallback exists + for, printing a warning at every step. + """ + model = self._make_model() + coord = torch.cat([self.coord, self.coord * 1.01], dim=0) + atype = torch.cat([self.atype, self.atype], dim=0) + box = torch.cat([self.box, self.box], dim=0) + + def hessian2(batch: int) -> torch.Tensor: + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", batch) + out = model.forward(coord.clone().requires_grad_(True), atype, box=box) + return out["hessian"].reshape(2, NDOF, NDOF) + + reference = hessian2(1) + + attempted = [] + real = mm._hessian_graph_batched_hvp + + def oom_above_two(*args, **kwargs): + attempted.append(kwargs["batch"]) + if kwargs["batch"] > 2: + raise torch.OutOfMemoryError("simulated") + return real(*args, **kwargs) + + monkeypatch.setattr(mm, "_hessian_graph_batched_hvp", oom_above_two) + recovered = hessian2(8) + + assert attempted == [8, 4, 2, 2], attempted + scale = reference.abs().max() + torch.testing.assert_close( + recovered, reference, rtol=0.0, atol=float(1e-12 * scale) + ) + + @pytest.mark.parametrize( + ("prices", "fits", "virtual", "expected"), + [ + ([8, 8], {5: 2}, False, [8, 4, 2, 2]), # overpriced twice: one descent + ([8, 1], {5: 2}, False, [8, 4, 2]), # a lower price still wins + ([2, 8], {5: 8}, False, [2, 8]), # nothing refused, nothing capped + ([8, 8], {5: 2, 4: 8}, True, [8, 4, 2, 8]), # another size, own price + ], + ) + def test_a_refused_batch_is_not_priced_again( + self, prices, fits, virtual, expected, monkeypatch + ) -> None: + """In automatic mode a batch the device refused caps later prices. + + Each frame is priced afresh, and a price is an estimate that can sit + above what the device actually holds. Taking every fresh price as + given sends each frame back down the ladder, paying the failed + attempts and printing the warnings again. Once the ladder has cut a + batch, later frames with as many real atoms are priced no higher than + the batch that survived; the cap only ever lowers a price, a call + that never ran out of memory caps nothing, and a frame with another + real-atom count keeps its own price instead of idling at a batch + learned on a different size. + """ + model = self._make_model() + coord = torch.cat([self.coord, self.coord * 1.01], dim=0) + atype = torch.cat([self.atype, self.atype], dim=0) + if virtual: + atype[1, -1] = -1 # 4 real atoms in the second frame + box = torch.cat([self.box, self.box], dim=0) + + def hessian2() -> torch.Tensor: + out = model.forward(coord.clone().requires_grad_(True), atype, box=box) + return out["hessian"].reshape(2, NDOF, NDOF) + + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 1) + reference = hessian2() + + # The price stands in for the probe so that the test runs on any + # device: on CUDA the probe would price from the real free memory. + quoted = iter(prices) + attempted = [] + real = mm._hessian_graph_batched_hvp + + def refuse_above_fits(*args, **kwargs): + attempted.append(kwargs["batch"]) + if kwargs["batch"] > fits[kwargs["nloc"]]: + raise torch.OutOfMemoryError("simulated") + return real(*args, **kwargs) + + monkeypatch.setattr(mm, "_auto_hvp_batch", lambda device, probe: next(quoted)) + monkeypatch.setattr(mm, "_hessian_graph_batched_hvp", refuse_above_fits) + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", None) + recovered = hessian2() + + assert next(quoted, None) is None, "each frame must still be priced" + assert attempted == expected, attempted + scale = reference.abs().max() + torch.testing.assert_close( + recovered, reference, rtol=0.0, atol=float(1e-12 * scale) + ) + + @pytest.mark.parametrize( + ("refused", "expected"), + [ + ("probe", [None]), # the measurement itself did not fit + ("ladder", [8]), # every batch above one was refused + ], + ) + def test_nothing_is_priced_after_only_one_row_fit( + self, refused, expected, monkeypatch + ) -> None: + """Once only a single row has fit, later frames skip the probe. + + A price can only be capped down to the batch that survived, so after + that batch reaches one the probe can no longer change the answer; it + would spend a forward and a first-order backward per frame for + nothing, and in this state that forward is what ran the device out of + memory -- each time printing the warning again. Three frames: the + first refuses, the other two must neither be priced nor batched. + """ + model = self._make_model() + coord = torch.cat([self.coord, self.coord * 1.01, self.coord * 0.99], dim=0) + atype = torch.cat([self.atype] * 3, dim=0) + box = torch.cat([self.box] * 3, dim=0) + + def hessian3() -> torch.Tensor: + out = model.forward(coord.clone().requires_grad_(True), atype, box=box) + return out["hessian"].reshape(3, NDOF, NDOF) + + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 1) + reference = hessian3() + + priced = [] + + def price(device, probe): + if refused == "probe": + priced.append(None) + raise torch.OutOfMemoryError("simulated") + priced.append(8) + return 8 + + def refuse_batches(*args, **kwargs): + raise torch.OutOfMemoryError("simulated") + + monkeypatch.setattr(mm, "_auto_hvp_batch", price) + monkeypatch.setattr(mm, "_hessian_graph_batched_hvp", refuse_batches) + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", None) + recovered = hessian3() + + assert priced == expected, priced + scale = reference.abs().max() + torch.testing.assert_close( + recovered, reference, rtol=0.0, atol=float(1e-12 * scale) + ) + + def test_hessian_is_symmetric(self, monkeypatch) -> None: + """Batching changes the summation order; it must not break symmetry.""" + model = self._make_model() + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", 8) + hessian = self._hessian(model) + scale = hessian.abs().max() + torch.testing.assert_close( + hessian, hessian.T, rtol=0.0, atol=float(1e-12 * scale) + ) + + +class TestHvpBatchPolicy: + """The automatic batch: bounded, monotone in free memory, and recoverable.""" + + MIB = 1024 * 1024 + + def _probe(self, nbytes: int): + """A stand-in Hessian-vector product that costs a known amount.""" + + def probe(): + return torch.empty(nbytes, dtype=torch.uint8, device=env.DEVICE) + + return probe + + @pytest.mark.skipif( + not torch.cuda.is_available(), reason="the automatic batch needs CUDA" + ) + @pytest.mark.parametrize("free_mib", [8, 64, 256, 1024, 4096, 16384, 65536]) + def test_auto_batch_is_bounded(self, free_mib, monkeypatch) -> None: + """Whatever the free memory, the batch stays within [1, cap].""" + monkeypatch.setattr( + torch.cuda, "mem_get_info", lambda *a, **k: (free_mib * self.MIB, 0) + ) + batch = mm._auto_hvp_batch(env.DEVICE, self._probe(32 * self.MIB)) + assert 1 <= batch <= mm.DP_HESSIAN_HVP_BATCH_CAP + + @pytest.mark.skipif( + not torch.cuda.is_available(), reason="the automatic batch needs CUDA" + ) + def test_auto_batch_rises_with_free_memory(self, monkeypatch) -> None: + """More memory must never buy a smaller batch.""" + chosen = [] + for free_mib in (8, 32, 128, 512, 2048, 8192): + monkeypatch.setattr( + torch.cuda, + "mem_get_info", + lambda *a, _f=free_mib, **k: (_f * self.MIB, 0), + ) + chosen.append(mm._auto_hvp_batch(env.DEVICE, self._probe(32 * self.MIB))) + assert chosen == sorted(chosen), chosen + # The ends must actually differ, or monotonicity is vacuous. + assert chosen[0] == 1 + assert chosen[-1] == mm.DP_HESSIAN_HVP_BATCH_CAP + + def test_auto_batch_is_one_without_cuda(self) -> None: + """No allocator introspection and no recoverable OOM: stay unbatched.""" + + def explode(): # pragma: no cover - must never be called + raise AssertionError("the probe must not run on a non-CUDA device") + + assert mm._auto_hvp_batch(torch.device("cpu"), explode) == 1