From a7d95a2a060af42e014fc4e0db69b6c9119cb46b Mon Sep 17 00:00:00 2001 From: Duo <50307526+iProzd@users.noreply.github.com> Date: Sun, 20 Sep 2026 13:34:29 -0400 Subject: [PATCH 01/15] perf(pt_expt): batch the Hessian-vector products on the graph route _cal_hessian_ext_graph called torch.autograd.functional.hessian, which with vectorize=False evaluates one Hessian row per second-order backward pass: for DPA-4 that is 3*nloc sequential passes, each far too small to keep the GPU busy. Replicate the structure along the frame axis the carry-all graph already has, so one forward and one first-order backward build a graph that every chunk of seed vectors reuses, and each second-order backward returns `batch` rows at once. Frames are independent, so the replicated energy is a sum of independent terms and its Hessian is block diagonal -- the result is exact, not an approximation, and no vmap is involved, so custom autograd Functions without a batching rule keep working. Measured on one H20, float64, TF32 off, DPA-4 from examples/water/dpa4: natoms B=1 (old) B=24 speedup max|H_B - H_1| 16 2.876 s 0.264 s 10.9x 8.3e-17 32 6.837 s 0.613 s 11.2x 1.2e-16 Peak memory grows with the batch: at 32 atoms 28.3 GiB (B=1) -> 34.6 GiB (B=24). DP_HESSIAN_HVP_BATCH tunes the trade-off; 1 restores the previous behaviour exactly. --- deepmd/pt_expt/model/make_model.py | 112 +++++++++++++++++++++++++++-- deepmd/pt_expt/utils/env.py | 5 ++ 2 files changed, 112 insertions(+), 5 deletions(-) diff --git a/deepmd/pt_expt/model/make_model.py b/deepmd/pt_expt/model/make_model.py index c61a5edca0..9f76eecc94 100644 --- a/deepmd/pt_expt/model/make_model.py +++ b/deepmd/pt_expt/model/make_model.py @@ -32,6 +32,9 @@ fused_energy_force_enabled, fused_operators_enabled, ) +from deepmd.pt_expt.utils.env import ( + DP_HESSIAN_HVP_BATCH, +) from deepmd.pt_expt.utils.graph_builder import ( build_neighbor_graph_for_method, build_ragged_neighbor_graph, @@ -376,6 +379,85 @@ def __call__(self, coord_flat: torch.Tensor) -> torch.Tensor: return atom_out.sum(dim=0).reshape(-1)[self.ci] +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, +) -> 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: on DPA-4 in float64 the + result matches the one-row-at-a-time path to ~1e-16, i.e. machine precision. + + Returns the ``(nloc * 3, nloc * 3)`` Hessian for output component ``ci``. + """ + ndof = nloc * 3 + nb = max(1, min(batch, ndof)) + 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() + (grad,) = torch.autograd.grad(total, x, create_graph=True) + + eye = torch.eye(ndof, dtype=x.dtype, device=x.device) + rows: list[torch.Tensor] = [] + for start in range(0, ndof, nb): + seeds = eye[start : start + nb] + if seeds.shape[0] < nb: # pad so the seed batch keeps the graph's shape + seeds = torch.cat([seeds, seeds.new_zeros(nb - seeds.shape[0], ndof)]) + (hvp,) = torch.autograd.grad( + grad, + x, + grad_outputs=seeds, + retain_graph=True, + create_graph=create_graph, + ) + rows.append(hvp[: min(nb, ndof - start)]) + return torch.cat(rows) + + def _cal_hessian_ext_graph( model: Any, kk: str, @@ -443,11 +525,31 @@ def _cal_hessian_ext_graph( spin=spin_frame, charge_spin=charge_spin_frame, ) - hess = torch.autograd.functional.hessian( - wrapper, - coord_flat, - create_graph=create_graph, - ) # (n_real*3, n_real*3) + if DP_HESSIAN_HVP_BATCH > 1: + hess = _hessian_graph_batched_hvp( + model=model, + kk=kk, + ci=ci, + nloc=n_real, + coord_flat=coord_flat, + 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, + batch=DP_HESSIAN_HVP_BATCH, + create_graph=create_graph, + ) # (n_real*3, n_real*3) + else: + hess = torch.autograd.functional.hessian( + wrapper, + coord_flat, + create_graph=create_graph, + ) # (n_real*3, n_real*3) 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..357e82a386 100644 --- a/deepmd/pt_expt/utils/env.py +++ b/deepmd/pt_expt/utils/env.py @@ -30,6 +30,11 @@ 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. 0 or 1 keeps the one-row-at-a-time +# path. Measured on an H20: 8-24 is the sweet spot, beyond that only memory grows. +DP_HESSIAN_HVP_BATCH = int(os.environ.get("DP_HESSIAN_HVP_BATCH", "16")) try: # only linux ncpus = len(os.sched_getaffinity(0)) From 01f450ee7199c2c892f8a3493911c504d2a3e2da Mon Sep 17 00:00:00 2001 From: Duo <50307526+iProzd@users.noreply.github.com> Date: Sun, 20 Sep 2026 22:33:40 -0400 Subject: [PATCH 02/15] perf(pt_expt): size the Hessian HVP batch from free memory Batching the Hessian-vector products needs a batch size, and there is no good fixed one. Peak memory is linear in it while the speedup is not: batching recovers kernel-launch overhead, which stops mattering once a single Hessian-vector product already saturates the device. Measured on one H20 with DPA-4.0.1-Pro-MPtrj in eval mode, float32, TF32 off, over 18 points spanning 54 to 19008 neighbour pairs and 8 to 512 atoms, peak(MiB) = 782 + [4.50 + 3.27*(B-1)]*edges + [10.9 + 7.96*(B-1)]*natoms fits every measurement to 0.2%, while the speedup falls from 5.06x at 72 edges to 1.24x at 5832. A fixed batch of 8 would cut what fits on a 96 GiB card at fcc-solid density from ~381 atoms to ~63, and buy 1.2x on the systems that large: it would turn systems that ran into systems that do not. So size it per call, as DP_INFER_BATCH_SIZE already does for inference batches. One Hessian-vector product is run to measure what a replica costs, the batch becomes what the free memory affords, and it is capped at 8, where the speedup has flattened on every system measured. The measurement covers a whole product while each step past the first adds only its marginal share, so the estimate reads high and the batch comes out conservative. That is the direction to err: a batch that does not fit is recovered by halving and retrying, and nothing recovers the time lost to one that was too small. Reaching 1 hands over to the original one-row-at-a-time path, so the fallback bottoms out in exactly the code a user who asked for 1 would take. An explicit DP_HESSIAN_HVP_BATCH is used as given, including above the cap. The out-of-memory fallback still applies to it, because halving changes how the Hessian is computed and not what it is, and the alternative is ending a run that could have finished. Without CUDA there is no allocator to size against and no recoverable out-of-memory error to catch, so the automatic choice is 1. This also corrects the measurements the comments quoted. They were taken with the model left in training mode, where this checkpoint's use_amp silently enables a bfloat16 autocast, so they described the bf16 path and overstated peak memory by roughly 18x. Likewise the equivalence claim: against the DPA-4.0.1-Pro-MPtrj checkpoint in float64 the two routes agree to 2.8e-14 relative RMS, not the 1e-16 measured on the smaller example model. --- deepmd/pt_expt/model/make_model.py | 176 +++++++++++++++++++++++------ deepmd/pt_expt/utils/env.py | 35 +++++- doc/env.md | 26 +++++ 3 files changed, 202 insertions(+), 35 deletions(-) diff --git a/deepmd/pt_expt/model/make_model.py b/deepmd/pt_expt/model/make_model.py index 9f76eecc94..0f061d8a2d 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 ( @@ -34,6 +35,8 @@ ) 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, @@ -379,6 +382,8 @@ 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, @@ -396,6 +401,7 @@ def _hessian_graph_batched_hvp( 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. @@ -408,10 +414,19 @@ def _hessian_graph_batched_hvp( 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: on DPA-4 in float64 the - result matches the one-row-at-a-time path to ~1e-16, i.e. machine precision. - - Returns the ``(nloc * 3, nloc * 3)`` Hessian for output component ``ci``. + 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)) @@ -443,7 +458,8 @@ def _hessian_graph_batched_hvp( eye = torch.eye(ndof, dtype=x.dtype, device=x.device) rows: list[torch.Tensor] = [] - for start in range(0, ndof, nb): + wanted = ndof if max_rows is None else min(max_rows, ndof) + for start in range(0, wanted, nb): seeds = eye[start : start + nb] if seeds.shape[0] < nb: # pad so the seed batch keeps the graph's shape seeds = torch.cat([seeds, seeds.new_zeros(nb - seeds.shape[0], ndof)]) @@ -454,8 +470,90 @@ def _hessian_graph_batched_hvp( retain_graph=True, create_graph=create_graph, ) - rows.append(hvp[: min(nb, ndof - start)]) - return torch.cat(rows) + 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, +) -> torch.Tensor: + """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. + """ + while batch > 1: + try: + return _hessian_graph_batched_hvp( + coord_flat=coord_flat, + batch=batch, + create_graph=create_graph, + **kwargs, + ) + except torch.OutOfMemoryError: + batch //= 2 + if torch.cuda.is_available(): + torch.cuda.empty_cache() + 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, + ) def _cal_hessian_ext_graph( @@ -495,6 +593,9 @@ def _cal_hessian_ext_graph( if charge_spin is not None and charge_spin.ndim == 1 else charge_spin ) + # Resolved once and reused: what a replica costs is set by the neighbour + # count, which is the same for every component of one call. + hvp_batch = DP_HESSIAN_HVP_BATCH hessians = [] for ii in range(nf): node_index = torch.nonzero(atype[ii] >= 0, as_tuple=False).reshape(-1) @@ -525,31 +626,42 @@ def _cal_hessian_ext_graph( spin=spin_frame, charge_spin=charge_spin_frame, ) - if DP_HESSIAN_HVP_BATCH > 1: - hess = _hessian_graph_batched_hvp( - model=model, - kk=kk, - ci=ci, - nloc=n_real, - coord_flat=coord_flat, - 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, - batch=DP_HESSIAN_HVP_BATCH, - create_graph=create_graph, - ) # (n_real*3, n_real*3) - else: - hess = torch.autograd.functional.hessian( - wrapper, - coord_flat, - create_graph=create_graph, - ) # (n_real*3, n_real*3) + hvp_kwargs = { + "model": model, + "kk": kk, + "ci": ci, + "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, + } + if hvp_batch is None and n_real: + 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, + **hvp_kwargs, + ), + ) + log.debug( + "Hessian-vector products batched %d rows at a time", hvp_batch + ) + hess = _hessian_graph_row_block( + batch=hvp_batch if hvp_batch is not None else 1, + wrapper=wrapper, + coord_flat=coord_flat, + create_graph=create_graph, + **hvp_kwargs, + ) # (n_real*3, n_real*3) 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 357e82a386..c930bb8e45 100644 --- a/deepmd/pt_expt/utils/env.py +++ b/deepmd/pt_expt/utils/env.py @@ -32,9 +32,38 @@ 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. 0 or 1 keeps the one-row-at-a-time -# path. Measured on an H20: 8-24 is the sweet spot, beyond that only memory grows. -DP_HESSIAN_HVP_BATCH = int(os.environ.get("DP_HESSIAN_HVP_BATCH", "16")) +# 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 call: 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..1061c03879 100644 --- a/doc/env.md +++ b/doc/env.md @@ -70,6 +70,32 @@ 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 to measure what a replica +costs, then takes the batch the free memory affords, capped at 8. 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 and warning which +batch was used. The result does not depend on the batch: it changes how the +Hessian is computed, not what it is. +::: + :::{envvar} DP_BACKEND **Default**: `tensorflow` From 031719fad789512e6451efccf10e0c3c83a39311 Mon Sep 17 00:00:00 2001 From: Duo <50307526+iProzd@users.noreply.github.com> Date: Sun, 20 Sep 2026 22:33:40 -0400 Subject: [PATCH 03/15] test(pt_expt): pin DP_HESSIAN_HVP_BATCH to the unbatched Hessian The batched Hessian-vector product path was reached only incidentally. test_dpa2_graph_lower exercises it, but only because the batch it happens to get exceeds 1: a batch of 1 would turn that into coverage of the unbatched path alone, without any test failing. Nothing treated the batch as the object under test, so no test compared one batch against another, and none asserted which branch had run. Add a float64 DPA-1 model that meets the graph-route gate (mixed types plus graph lower) and check, for batches 2, 3, 4, 8 and 16, that the Hessian matches the one produced with a batch of 1. 3*nloc is 15, so 2, 4, 8 and 16 leave a partial final chunk and cover the zero-padding branch, while 3 divides evenly and 16 exceeds 3*nloc and is clamped. A counter over the three implementations asserts that batches of 0 and 1 take the unbatched path, that the larger ones take the batched helper, and that the dense route neither sees the setting nor changes with it -- the dense Hessian doubles as an independent cross-check, since it shares no code with the graph wrapper. Cover the automatic choice too: that it stays within [1, cap] whatever the free memory, that it never shrinks as memory grows (with both ends pinned, so monotonicity cannot pass vacuously), that it stays at 1 without CUDA without running the probe, and that the probe prices one Hessian-vector product rather than the whole Hessian. Cover the fallback by making the helper refuse batches above 2 and asserting the retry walks 8, 4, 2 and still returns the unbatched answer, and by making every batch fail and asserting it lands on the original path. Verified non-vacuous by mutation. Of the 23 tests, these many fail when the implementation is broken in each way: dropping the trim that discards padded rows, 10; shifting the seed vectors by one row, 8; letting a batch of 1 enter the batched helper, 6; removing the batch cap, 5; ignoring free memory and always taking the cap, 1; retrying at 1 instead of halving, 1; running the probe without CUDA, 1; ignoring max_rows so the probe prices the whole Hessian, 1. --- .../pt_expt/model/test_hessian_hvp_batch.py | 363 ++++++++++++++++++ 1 file changed, 363 insertions(+) create mode 100644 source/tests/pt_expt/model/test_hessian_hvp_batch.py 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..b9fa7ff875 --- /dev/null +++ b/source/tests/pt_expt/model/test_hessian_hvp_batch.py @@ -0,0 +1,363 @@ +# 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 call 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() + + 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_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) + ) + + 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_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 + From 973ad0cef5de08ca82f0627bedacf4398dc1742b Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Mon, 21 Sep 2026 14:51:54 +0000 Subject: [PATCH 04/15] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- deepmd/pt_expt/model/make_model.py | 6 ++---- .../pt_expt/model/test_hessian_hvp_batch.py | 21 ++++++++----------- 2 files changed, 11 insertions(+), 16 deletions(-) diff --git a/deepmd/pt_expt/model/make_model.py b/deepmd/pt_expt/model/make_model.py index 0f061d8a2d..e7e7249444 100644 --- a/deepmd/pt_expt/model/make_model.py +++ b/deepmd/pt_expt/model/make_model.py @@ -384,6 +384,7 @@ def __call__(self, coord_flat: torch.Tensor) -> torch.Tensor: log = logging.getLogger(__name__) + def _hessian_graph_batched_hvp( model: Any, kk: str, @@ -476,7 +477,6 @@ def _hessian_graph_batched_hvp( 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) @@ -510,9 +510,7 @@ def _auto_hvp_batch(device: torch.device, probe: Any) -> int: 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 - ) + 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)) diff --git a/source/tests/pt_expt/model/test_hessian_hvp_batch.py b/source/tests/pt_expt/model/test_hessian_hvp_batch.py index b9fa7ff875..3132e8c306 100644 --- a/source/tests/pt_expt/model/test_hessian_hvp_batch.py +++ b/source/tests/pt_expt/model/test_hessian_hvp_batch.py @@ -152,7 +152,9 @@ def _hessian(self, model: EnergyModel) -> torch.Tensor: ) return out["hessian"].reshape(NDOF, NDOF) - def test_batch_one_takes_the_unbatched_path(self, route_counts, monkeypatch) -> None: + 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): @@ -197,9 +199,7 @@ def test_padded_final_chunk_is_discarded(self, batch, monkeypatch) -> None: # zero row in the middle; neither is possible for a real Hessian here. assert (hessian.abs().sum(dim=1) > 0).all() - def test_probe_prices_one_product_not_the_whole_hessian( - self, monkeypatch - ) -> None: + 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; @@ -317,18 +317,14 @@ 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 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] - ) + @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( @@ -345,7 +341,9 @@ def test_auto_batch_rises_with_free_memory(self, monkeypatch) -> None: 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) + 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 @@ -360,4 +358,3 @@ 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 - From 8c703fc69ff840a479a92f9dbfa51d4f38e86126 Mon Sep 17 00:00:00 2001 From: Duo <50307526+iProzd@users.noreply.github.com> Date: Mon, 21 Sep 2026 12:49:46 -0400 Subject: [PATCH 05/15] test(pt_expt): skip the probe-cost test without CUDA The automatic batch choice returns 1 immediately on a non-CUDA device and never runs the probe -- there is no allocator to size against and no recoverable out-of-memory error to catch, which is what test_auto_batch_is_one_without_cuda asserts. The test that checks the probe prices a single Hessian-vector product therefore cannot pass there, and it was missing the skip its two sibling tests already carry. On CPU: 1 failed, 14 passed, 8 skipped -> 14 passed, 9 skipped. On CUDA: 23 passed, unchanged. --- source/tests/pt_expt/model/test_hessian_hvp_batch.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/source/tests/pt_expt/model/test_hessian_hvp_batch.py b/source/tests/pt_expt/model/test_hessian_hvp_batch.py index 3132e8c306..4966870af5 100644 --- a/source/tests/pt_expt/model/test_hessian_hvp_batch.py +++ b/source/tests/pt_expt/model/test_hessian_hvp_batch.py @@ -199,6 +199,9 @@ def test_padded_final_chunk_is_discarded(self, batch, monkeypatch) -> None: # 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. From be7432245ee5be03ee4e885e8ccd3682c8238e38 Mon Sep 17 00:00:00 2001 From: Duo <50307526+iProzd@users.noreply.github.com> Date: Tue, 22 Sep 2026 02:47:53 -0400 Subject: [PATCH 06/15] fix(pt_expt): correct four defects in the batched Hessian route Review found four problems with the batched Hessian-vector products, all of them in how the batched path differs from the `torch.autograd.functional.hessian` call it replaces. 1. `create_graph` no longer detaches the coordinates. The batched path built its replicas from `coord_flat.detach()` whatever the caller asked for, so with `create_graph=True` the Hessian lost its dependence on the input coordinates. The parameter path was unaffected -- training on Hessian labels still trained -- which is what made the loss silent. `functional.hessian` keeps the input in the graph in this case (`_grad_preprocess`), and now so does this. 2. A constant or linear coordinate dependence yields zeros instead of raising. `functional.hessian` materialises the zero block under its default `strict=False`; differentiating a constant first derivative again raises instead, so a legitimate model became a crash. 3. The memory probe no longer builds an `ndof x ndof` identity. Seeds are now cut one block at a time. The identity is the allocation batching exists to avoid, and building it to decide the batch both risked the out-of-memory error being measured for and inflated the measurement, biasing the chosen batch too small. 4. An out-of-memory error from the probe falls back to one row at a time. Pricing a product is itself a product, so it can be the allocation that does not fit; that escaped before the halving path could run, ending a job that the unbatched route might have completed. Each fix has a test that fails without it: reverting any one of the four fails exactly its own test and leaves the other three green. --- deepmd/pt_expt/model/make_model.py | 75 +++++++++--- .../pt_expt/model/test_hessian_hvp_batch.py | 115 ++++++++++++++++++ 2 files changed, 172 insertions(+), 18 deletions(-) diff --git a/deepmd/pt_expt/model/make_model.py b/deepmd/pt_expt/model/make_model.py index e7e7249444..45fb9152ad 100644 --- a/deepmd/pt_expt/model/make_model.py +++ b/deepmd/pt_expt/model/make_model.py @@ -431,8 +431,15 @@ def _hessian_graph_batched_hvp( """ ndof = nloc * 3 nb = max(1, min(batch, ndof)) - x = coord_flat.detach().reshape(1, ndof).expand(nb, ndof).contiguous() - x = x.requires_grad_(True) + 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() @@ -455,21 +462,39 @@ def _hessian_graph_batched_hvp( # 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() - (grad,) = torch.autograd.grad(total, x, create_graph=True) + (grad,) = torch.autograd.grad( + total, x, create_graph=True, allow_unused=True, materialize_grads=True + ) - eye = torch.eye(ndof, dtype=x.dtype, device=x.device) - rows: list[torch.Tensor] = [] wanted = ndof if max_rows is None else min(max_rows, ndof) + if not grad.requires_grad: + # The reduced output is constant or linear in the coordinates, 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, which would turn a + # legitimate model into a crash. + return x.new_zeros(wanted, ndof) + + rows: list[torch.Tensor] = [] for start in range(0, wanted, nb): - seeds = eye[start : start + nb] - if seeds.shape[0] < nb: # pad so the seed batch keeps the graph's shape - seeds = torch.cat([seeds, seeds.new_zeros(nb - seeds.shape[0], ndof)]) + 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 @@ -640,16 +665,30 @@ def _cal_hessian_ext_graph( "charge_spin": charge_spin_frame, } if hvp_batch is None and n_real: - 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, - **hvp_kwargs, - ), - ) + 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, + **hvp_kwargs, + ), + ) + except torch.OutOfMemoryError: + # 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 torch.cuda.is_available(): + torch.cuda.empty_cache() + 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 = 1 log.debug( "Hessian-vector products batched %d rows at a time", hvp_batch ) diff --git a/source/tests/pt_expt/model/test_hessian_hvp_batch.py b/source/tests/pt_expt/model/test_hessian_hvp_batch.py index 4966870af5..341b2fa8f5 100644 --- a/source/tests/pt_expt/model/test_hessian_hvp_batch.py +++ b/source/tests/pt_expt/model/test_hessian_hvp_batch.py @@ -227,6 +227,121 @@ def record(*args, **kwargs): assert probe_max_rows == 1 assert probe_rows == 1, "the probe computed more than one row" + 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() + + def test_a_linear_energy_gives_a_zero_hessian_not_a_crash( + self, monkeypatch + ) -> None: + """Constant or linear coordinate dependence must yield zeros. + + ``functional.hessian`` materialises the zero block under its default + ``strict=False``. Differentiating a constant first derivative again + instead raises, which turns a legitimate model into a crash. + """ + + class _LinearAtomicModel: + """Energy exactly linear in the coordinates: d2E/dx2 == 0.""" + + def forward_common_atomic_graph(self, graph, atype_flat, **kwargs): + nb = graph.shape[0] + energy = (3.0 * graph).sum(-1).reshape(nb * graph.shape[1], 1) + return {"energy": energy} + + class _Model: + atomic_model = _LinearAtomicModel() + + 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, "a linear 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.skipif( + not torch.cuda.is_available(), reason="the automatic choice needs CUDA" + ) + def test_an_out_of_memory_probe_falls_back_instead_of_escaping( + self, 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. + """ + monkeypatch.setattr(mm, "DP_HESSIAN_HVP_BATCH", None) + + def always_oom(device, probe): + 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) From 5fed0e6eae5de00b7c7df257142c389dde349f12 Mon Sep 17 00:00:00 2001 From: Duo <50307526+iProzd@users.noreply.github.com> Date: Tue, 22 Sep 2026 07:18:02 -0400 Subject: [PATCH 07/15] fix(pt_expt): a coordinate-independent output also yields a zero Hessian The previous commit handled only half of this. It guarded on `grad.requires_grad`, which catches an output that is linear in the coordinates -- the first derivative exists but is constant. An output that does not depend on the coordinates at all never gets that far: `total.requires_grad` is False, and autograd refuses to differentiate an output carrying no graph, so it raised at the *first* `autograd.grad`, before the guard could run. `torch.autograd.functional.hessian` returns an all-zero Hessian for both shapes under its default `strict=False`, so both have to be materialised here. The check now sits before the first derivative. The regression is parametrised over both shapes rather than only the linear one, which is what let the constant case through: the guard was never exercised with the input it was supposed to catch. Removing this guard alone fails the `constant` case and leaves `linear` green. --- deepmd/pt_expt/model/make_model.py | 19 +++++++--- .../pt_expt/model/test_hessian_hvp_batch.py | 38 +++++++++++++------ 2 files changed, 39 insertions(+), 18 deletions(-) diff --git a/deepmd/pt_expt/model/make_model.py b/deepmd/pt_expt/model/make_model.py index 45fb9152ad..481a296150 100644 --- a/deepmd/pt_expt/model/make_model.py +++ b/deepmd/pt_expt/model/make_model.py @@ -462,17 +462,24 @@ def _hessian_graph_batched_hvp( # 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 ) - wanted = ndof if max_rows is None else min(max_rows, ndof) if not grad.requires_grad: - # The reduced output is constant or linear in the coordinates, 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, which would turn a - # legitimate model into a crash. + # 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] = [] diff --git a/source/tests/pt_expt/model/test_hessian_hvp_batch.py b/source/tests/pt_expt/model/test_hessian_hvp_batch.py index 341b2fa8f5..770c51d358 100644 --- a/source/tests/pt_expt/model/test_hessian_hvp_batch.py +++ b/source/tests/pt_expt/model/test_hessian_hvp_batch.py @@ -249,26 +249,38 @@ def test_create_graph_keeps_the_path_to_the_coordinates(self, monkeypatch) -> No assert back is not None, "the Hessian no longer depends on the coordinates" assert torch.isfinite(back).all() - def test_a_linear_energy_gives_a_zero_hessian_not_a_crash( - self, monkeypatch + @pytest.mark.parametrize("dependence", ["linear", "constant"]) + def test_a_curvature_free_energy_gives_a_zero_hessian_not_a_crash( + self, dependence, monkeypatch ) -> None: - """Constant or linear coordinate dependence must yield zeros. + """Constant *and* linear coordinate dependence must yield zeros. ``functional.hessian`` materialises the zero block under its default - ``strict=False``. Differentiating a constant first derivative again - instead raises, which turns a legitimate model into a crash. + ``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 _LinearAtomicModel: - """Energy exactly linear in the coordinates: d2E/dx2 == 0.""" + 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 = graph.shape[0] - energy = (3.0 * graph).sum(-1).reshape(nb * graph.shape[1], 1) - return {"energy": energy} + 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 = _LinearAtomicModel() + atomic_model = _CurvatureFreeAtomicModel(dependence) monkeypatch.setattr( mm, @@ -296,7 +308,9 @@ class _Model: create_graph=False, ) assert hessian.shape == (NDOF, NDOF) - assert torch.count_nonzero(hessian) == 0, "a linear energy has no curvature" + 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. From 2c4659ba8a4e40990b7321533da5524c9cac5dd8 Mon Sep 17 00:00:00 2001 From: Duo <50307526+iProzd@users.noreply.github.com> Date: Thu, 24 Sep 2026 21:29:00 -0400 Subject: [PATCH 08/15] fix(pt_expt): recognise wrapped out-of-memory errors in Hessian fallback The batched-Hessian ladder and the pricing probe caught only torch.OutOfMemoryError, but the allocator failure does not always arrive with that type: AOTInductor rewraps it in a plain RuntimeError, keeping the original text in the message, only in the __cause__ chain, or stripping both behind its run_func_ signature. A catch keyed on the exception type lets every one of those forms end the run, which the fallback exists to prevent. Both sites now delegate to the repository's is_oom_error (deepmd/pt/utils/auto_batch_size.py), which walks the exception chain and knows the wrapper signatures; it also releases the allocator cache itself on a positive match. The probe fallback and the ladder behave exactly as before for the unwrapped form. The tests feed the ladder all three wrapped forms and assert it still halves 8, 4, 2 and returns the unbatched answer, and that an unrelated RuntimeError propagates instead of being retried smaller; the probe fallback test is parametrised over the plain and wrapped forms. --- deepmd/pt_expt/model/make_model.py | 20 ++++-- .../pt_expt/model/test_hessian_hvp_batch.py | 65 ++++++++++++++++++- 2 files changed, 77 insertions(+), 8 deletions(-) diff --git a/deepmd/pt_expt/model/make_model.py b/deepmd/pt_expt/model/make_model.py index 481a296150..36960f2912 100644 --- a/deepmd/pt_expt/model/make_model.py +++ b/deepmd/pt_expt/model/make_model.py @@ -25,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, @@ -561,6 +564,11 @@ def _hessian_graph_row_block( 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. """ while batch > 1: try: @@ -570,10 +578,10 @@ def _hessian_graph_row_block( create_graph=create_graph, **kwargs, ) - except torch.OutOfMemoryError: + except Exception as e: + if not AutoBatchSize(silent=True).is_oom_error(e): + raise batch //= 2 - if torch.cuda.is_available(): - torch.cuda.empty_cache() log.warning( "Hessian-vector product batch did not fit in memory; retrying " "with %d. Set DP_HESSIAN_HVP_BATCH to choose it yourself.", @@ -683,13 +691,13 @@ def _cal_hessian_ext_graph( **hvp_kwargs, ), ) - except torch.OutOfMemoryError: + 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 torch.cuda.is_available(): - torch.cuda.empty_cache() + 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 " diff --git a/source/tests/pt_expt/model/test_hessian_hvp_batch.py b/source/tests/pt_expt/model/test_hessian_hvp_batch.py index 770c51d358..6d9c3990cc 100644 --- a/source/tests/pt_expt/model/test_hessian_hvp_batch.py +++ b/source/tests/pt_expt/model/test_hessian_hvp_batch.py @@ -337,17 +337,21 @@ def record_eye(n, *args, **kwargs): @pytest.mark.skipif( not torch.cuda.is_available(), reason="the automatic choice needs CUDA" ) + @pytest.mark.parametrize("wrapped", [False, True]) def test_an_out_of_memory_probe_falls_back_instead_of_escaping( - self, monkeypatch + 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. + 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) @@ -404,6 +408,63 @@ def only_small_batches_fit(*args, **kwargs): 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: From 3049b57a3606e7b1a05e5084eff98f34eb4d9266 Mon Sep 17 00:00:00 2001 From: Duo <50307526+iProzd@users.noreply.github.com> Date: Thu, 24 Sep 2026 23:22:34 -0400 Subject: [PATCH 09/15] fix(pt_expt): keep the batch the Hessian OOM ladder discovers The halved batch stayed local to _hessian_graph_row_block and was never written back to hvp_batch, so every frame and every output component restarted the ladder from the top: three frames with only batches up to 2 fitting attempted [8, 4, 2, 8, 4, 2, 8, 4, 2] and printed six warnings. Each failed attempt pays a full forward plus a first-order backward, so in exactly the regime the fallback exists for it spent the speedup batching was bought for. The helper now returns the batch that fit alongside the Hessian, and the caller keeps it in hvp_batch, so later calls start where the last one survived. The two-frame test asserts the attempts are [8, 4, 2, 2] and that both frames still return the unbatched answer. --- deepmd/pt_expt/model/make_model.py | 33 ++++++++++------ doc/env.md | 7 ++-- .../pt_expt/model/test_hessian_hvp_batch.py | 39 +++++++++++++++++++ 3 files changed, 65 insertions(+), 14 deletions(-) diff --git a/deepmd/pt_expt/model/make_model.py b/deepmd/pt_expt/model/make_model.py index 36960f2912..590fb562ff 100644 --- a/deepmd/pt_expt/model/make_model.py +++ b/deepmd/pt_expt/model/make_model.py @@ -556,7 +556,7 @@ def _hessian_graph_row_block( coord_flat: torch.Tensor, create_graph: bool, **kwargs: Any, -) -> torch.Tensor: +) -> 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 @@ -569,14 +569,22 @@ def _hessian_graph_row_block( 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, + 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): @@ -587,10 +595,13 @@ def _hessian_graph_row_block( "with %d. Set DP_HESSIAN_HVP_BATCH to choose it yourself.", batch, ) - return torch.autograd.functional.hessian( - wrapper, - coord_flat, - create_graph=create_graph, + return ( + torch.autograd.functional.hessian( + wrapper, + coord_flat, + create_graph=create_graph, + ), + 1, ) @@ -707,7 +718,7 @@ def _cal_hessian_ext_graph( log.debug( "Hessian-vector products batched %d rows at a time", hvp_batch ) - hess = _hessian_graph_row_block( + hess, hvp_batch = _hessian_graph_row_block( batch=hvp_batch if hvp_batch is not None else 1, wrapper=wrapper, coord_flat=coord_flat, diff --git a/doc/env.md b/doc/env.md index 1061c03879..0802ec9866 100644 --- a/doc/env.md +++ b/doc/env.md @@ -91,9 +91,10 @@ 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 and warning which -batch was used. The result does not depend on the batch: it changes how the -Hessian is computed, not what it is. +out-of-memory errors can still reduce the batch, halving it, 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 diff --git a/source/tests/pt_expt/model/test_hessian_hvp_batch.py b/source/tests/pt_expt/model/test_hessian_hvp_batch.py index 6d9c3990cc..76bc889289 100644 --- a/source/tests/pt_expt/model/test_hessian_hvp_batch.py +++ b/source/tests/pt_expt/model/test_hessian_hvp_batch.py @@ -490,6 +490,45 @@ def nothing_fits(*args, **kwargs): 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) + ) + def test_hessian_is_symmetric(self, monkeypatch) -> None: """Batching changes the summation order; it must not break symmetry.""" model = self._make_model() From 923bed525ce33b2a306ab90595baba1edfe7cd06 Mon Sep 17 00:00:00 2001 From: Duo <50307526+iProzd@users.noreply.github.com> Date: Thu, 24 Sep 2026 23:27:23 -0400 Subject: [PATCH 10/15] fix(pt_expt): price the Hessian batch per frame The automatic batch was priced once, on the first frame with real atoms, and reused for every later frame, but each frame has its own neighbour count and its own count of real atoms: a batch priced on a sparse frame is then spent on a dense one, which only the OOM ladder recovers, and a batch priced on a dense frame idles on sparse ones, which nothing recovers. The Hessians accumulated for earlier frames also shrink the free memory a later frame has, which a price taken up front cannot know. Pricing costs a single Hessian-vector product (one row out of 3 * nloc), so run it once per frame; within a frame every output component still shares one price. The CUDA test asserts a two-frame call prices twice and still returns the unbatched answer. --- deepmd/pt_expt/model/make_model.py | 93 ++++++++++--------- doc/env.md | 6 +- .../pt_expt/model/test_hessian_hvp_batch.py | 45 ++++++++- 3 files changed, 95 insertions(+), 49 deletions(-) diff --git a/deepmd/pt_expt/model/make_model.py b/deepmd/pt_expt/model/make_model.py index 590fb562ff..ceb3403046 100644 --- a/deepmd/pt_expt/model/make_model.py +++ b/deepmd/pt_expt/model/make_model.py @@ -642,9 +642,11 @@ def _cal_hessian_ext_graph( if charge_spin is not None and charge_spin.ndim == 1 else charge_spin ) - # Resolved once and reused: what a replica costs is set by the neighbour - # count, which is the same for every component of one call. + # 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. hvp_batch = DP_HESSIAN_HVP_BATCH + auto_batch = hvp_batch is None hessians = [] for ii in range(nf): node_index = torch.nonzero(atype[ii] >= 0, as_tuple=False).reshape(-1) @@ -659,6 +661,49 @@ 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, + } + if 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 = 1 + log.debug("Hessian-vector products batched %d rows at a time", hvp_batch) for ci in range(vsize): wrapper = _WrapperForwardEnergyGraph( model=model, @@ -675,54 +720,12 @@ def _cal_hessian_ext_graph( spin=spin_frame, charge_spin=charge_spin_frame, ) - hvp_kwargs = { - "model": model, - "kk": kk, - "ci": ci, - "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, - } - if hvp_batch is None and n_real: - 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, - **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 = 1 - log.debug( - "Hessian-vector products batched %d rows at a time", hvp_batch - ) hess, hvp_batch = _hessian_graph_row_block( batch=hvp_batch if hvp_batch is not None else 1, wrapper=wrapper, coord_flat=coord_flat, create_graph=create_graph, + ci=ci, **hvp_kwargs, ) # (n_real*3, n_real*3) if n_real != nloc: diff --git a/doc/env.md b/doc/env.md index 0802ec9866..b3cd74d5ac 100644 --- a/doc/env.md +++ b/doc/env.md @@ -81,8 +81,10 @@ 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 to measure what a replica -costs, then takes the batch the free memory affords, capped at 8. Peak memory +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. 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 diff --git a/source/tests/pt_expt/model/test_hessian_hvp_batch.py b/source/tests/pt_expt/model/test_hessian_hvp_batch.py index 76bc889289..296d2801c3 100644 --- a/source/tests/pt_expt/model/test_hessian_hvp_batch.py +++ b/source/tests/pt_expt/model/test_hessian_hvp_batch.py @@ -13,8 +13,8 @@ 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 call from free memory", so the -choice itself and the out-of-memory fallback are tested too. +``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 @@ -227,6 +227,47 @@ def record(*args, **kwargs): assert probe_max_rows == 1 assert probe_rows == 1, "the probe computed more than one row" + @pytest.mark.skipif( + not torch.cuda.is_available(), reason="the automatic choice needs CUDA" + ) + 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. From 228c8c4614672ff9d2119ca150e2a2c4c2cbcab5 Mon Sep 17 00:00:00 2001 From: Duo <50307526+iProzd@users.noreply.github.com> Date: Mon, 5 Oct 2026 07:31:11 -0400 Subject: [PATCH 11/15] test(pt_expt): run the per-frame pricing test without CUDA test_the_batch_is_repriced_for_each_frame only counts calls to _auto_hvp_batch, which _cal_hessian_ext_graph makes once per frame on every device; on a device without CUDA it returns 1 without touching the allocator. Gated on torch.cuda.is_available(), the regression test for per-frame pricing could run only in the CUDA jobs, which skip on pull requests unless labelled. On CPU it fails with make_model.py from before per-frame pricing ("priced 1 times for 2 frames") and passes with it. --- source/tests/pt_expt/model/test_hessian_hvp_batch.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/source/tests/pt_expt/model/test_hessian_hvp_batch.py b/source/tests/pt_expt/model/test_hessian_hvp_batch.py index 296d2801c3..8890f609ac 100644 --- a/source/tests/pt_expt/model/test_hessian_hvp_batch.py +++ b/source/tests/pt_expt/model/test_hessian_hvp_batch.py @@ -227,9 +227,6 @@ def record(*args, **kwargs): assert probe_max_rows == 1 assert probe_rows == 1, "the probe computed more than one row" - @pytest.mark.skipif( - not torch.cuda.is_available(), reason="the automatic choice needs CUDA" - ) def test_the_batch_is_repriced_for_each_frame(self, monkeypatch) -> None: """One price does not fit all frames: neighbour counts differ per frame. From 55f8a72c82f5614d9d15f8f0d2f849b06684f538 Mon Sep 17 00:00:00 2001 From: Duo <50307526+iProzd@users.noreply.github.com> Date: Mon, 5 Oct 2026 07:31:11 -0400 Subject: [PATCH 12/15] test(pt_expt): run the probe out-of-memory fallback test without CUDA test_an_out_of_memory_probe_falls_back_instead_of_escaping replaces _auto_hvp_batch with a stub that raises, so it exercises only the handler in _cal_hessian_ext_graph, which is the same on every device; the CUDA gate only kept it out of unlabelled pull-request runs. On CPU it fails when the handler re-raises and passes with it. --- source/tests/pt_expt/model/test_hessian_hvp_batch.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/source/tests/pt_expt/model/test_hessian_hvp_batch.py b/source/tests/pt_expt/model/test_hessian_hvp_batch.py index 8890f609ac..f9c262b089 100644 --- a/source/tests/pt_expt/model/test_hessian_hvp_batch.py +++ b/source/tests/pt_expt/model/test_hessian_hvp_batch.py @@ -372,9 +372,6 @@ def record_eye(n, *args, **kwargs): f"an identity of size {NDOF} was built; sizes seen: {seen}" ) - @pytest.mark.skipif( - not torch.cuda.is_available(), reason="the automatic choice needs CUDA" - ) @pytest.mark.parametrize("wrapped", [False, True]) def test_an_out_of_memory_probe_falls_back_instead_of_escaping( self, wrapped, monkeypatch From efa3dc17b79b9761dc741cd73bf0753b6de0b5db Mon Sep 17 00:00:00 2001 From: Duo <50307526+iProzd@users.noreply.github.com> Date: Mon, 5 Oct 2026 07:31:11 -0400 Subject: [PATCH 13/15] fix(pt_expt): cap later Hessian prices at the batch the ladder kept In automatic mode every frame is priced afresh, and the fresh price overwrote the batch the previous frame's out-of-memory ladder had settled on. A price is an estimate of what the device holds; an out-of-memory error is a measurement of it. On a call whose real ceiling sits below the estimate, each frame walked the ladder down again, paying a failed forward and first-order backward per step and printing the warnings again. Once a batch has been refused, later frames with the same number of real atoms are now priced no higher than the batch that survived. The cap only lowers a price, it is set only when a batch was actually refused, and every frame is still priced, so a sparser frame keeps its own lower price. It is keyed by the real-atom count, the only size known before the graph is built, so a frame of another size is priced on its own instead of idling at a batch learned on a different one -- the mismatch per-frame pricing was introduced to avoid. A probe that itself runs out of memory counts as a refusal down to one row, and once one row is all that fits, later frames of that size skip the probe: it can no longer change the batch, and it is the allocation that failed. Explicit values keep their existing behaviour. The env.py comment still said the batch was chosen per call; since 923bed525 it is chosen per frame. The new tests drive automatic mode with stub prices on any device: [8, 8] with batches up to 2 fitting attempts [8, 4, 2, 2] rather than [8, 4, 2, 8, 4, 2]; a lower later price still wins over the cap; a call that never ran out of memory caps nothing; a second frame with one atom fewer keeps its own price ([8, 4, 2, 8]); and after a refused probe or a ladder that bottomed out, the remaining frames of a three-frame call are not priced again. --- deepmd/pt_expt/model/make_model.py | 25 +++- deepmd/pt_expt/utils/env.py | 2 +- doc/env.md | 5 +- .../pt_expt/model/test_hessian_hvp_batch.py | 116 ++++++++++++++++++ 4 files changed, 142 insertions(+), 6 deletions(-) diff --git a/deepmd/pt_expt/model/make_model.py b/deepmd/pt_expt/model/make_model.py index ceb3403046..ac73cf45b7 100644 --- a/deepmd/pt_expt/model/make_model.py +++ b/deepmd/pt_expt/model/make_model.py @@ -644,9 +644,17 @@ def _cal_hessian_ext_graph( ) # 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. + # 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) @@ -675,7 +683,11 @@ def _cal_hessian_ext_graph( "spin": spin_frame, "charge_spin": charge_spin_frame, } - if auto_batch and n_real: + 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: @@ -702,7 +714,9 @@ def _cal_hessian_ext_graph( "product; falling back to one row at a time. Set " "DP_HESSIAN_HVP_BATCH to choose the batch yourself." ) - hvp_batch = 1 + 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( @@ -720,14 +734,17 @@ def _cal_hessian_ext_graph( spin=spin_frame, charge_spin=charge_spin_frame, ) + batch = hvp_batch if hvp_batch is not None else 1 hess, hvp_batch = _hessian_graph_row_block( - batch=hvp_batch if hvp_batch is not None else 1, + 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 c930bb8e45..fbbc962c86 100644 --- a/deepmd/pt_expt/utils/env.py +++ b/deepmd/pt_expt/utils/env.py @@ -35,7 +35,7 @@ # 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 call: one Hessian-vector product is run to +# 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 diff --git a/doc/env.md b/doc/env.md index b3cd74d5ac..353a1e99e2 100644 --- a/doc/env.md +++ b/doc/env.md @@ -84,7 +84,10 @@ 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. Peak 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 diff --git a/source/tests/pt_expt/model/test_hessian_hvp_batch.py b/source/tests/pt_expt/model/test_hessian_hvp_batch.py index f9c262b089..dcb4873b6c 100644 --- a/source/tests/pt_expt/model/test_hessian_hvp_batch.py +++ b/source/tests/pt_expt/model/test_hessian_hvp_batch.py @@ -564,6 +564,122 @@ def oom_above_two(*args, **kwargs): 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() From 38b1b28e7ffa3a4e0d885c9ec4d16c8679408b85 Mon Sep 17 00:00:00 2001 From: Duo <50307526+iProzd@users.noreply.github.com> Date: Mon, 5 Oct 2026 09:54:57 -0400 Subject: [PATCH 14/15] fix(pt_expt): round the Hessian batch up when halving after OOM The out-of-memory ladder halved the batch by rounding down, so an odd batch skipped the one below it: 3 went straight to 1, the one-row-at-a-time path, and a batch of 2 that would have fit was never tried. The automatic price lands on odd values, so this is the common case: on DPA-4 a price of 7 under a capped device went 7, 3, 1 and took about twice as long as a batch of 2 would have. Round up instead, so 7 goes to 4 and then 2, and 1 is reached only from 2. Even batches halve exactly as before. The test starts the ladder from 7, 3 and 5 and asserts [7, 4, 2], [3, 2] and [5, 3]; rounding down fails all three. --- deepmd/pt_expt/model/make_model.py | 4 +- doc/env.md | 3 +- .../pt_expt/model/test_hessian_hvp_batch.py | 41 +++++++++++++++++++ 3 files changed, 46 insertions(+), 2 deletions(-) diff --git a/deepmd/pt_expt/model/make_model.py b/deepmd/pt_expt/model/make_model.py index ac73cf45b7..6d2dedfb80 100644 --- a/deepmd/pt_expt/model/make_model.py +++ b/deepmd/pt_expt/model/make_model.py @@ -589,7 +589,9 @@ def _hessian_graph_row_block( except Exception as e: if not AutoBatchSize(silent=True).is_oom_error(e): raise - batch //= 2 + # 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.", diff --git a/doc/env.md b/doc/env.md index 353a1e99e2..586c610c69 100644 --- a/doc/env.md +++ b/doc/env.md @@ -96,7 +96,8 @@ 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, warning which +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. diff --git a/source/tests/pt_expt/model/test_hessian_hvp_batch.py b/source/tests/pt_expt/model/test_hessian_hvp_batch.py index dcb4873b6c..ec495918f2 100644 --- a/source/tests/pt_expt/model/test_hessian_hvp_batch.py +++ b/source/tests/pt_expt/model/test_hessian_hvp_batch.py @@ -443,6 +443,47 @@ def only_small_batches_fit(*args, **kwargs): 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 From 0e01b77bc18147e001e411767fb1e284fc318d66 Mon Sep 17 00:00:00 2001 From: Duo <50307526+iProzd@users.noreply.github.com> Date: Tue, 6 Oct 2026 05:30:54 -0400 Subject: [PATCH 15/15] test(pt_expt): keep the allocator cache out of the batch-policy tests TestHvpBatchPolicy simulates free memory by patching mem_get_info, but _auto_hvp_batch also counts what the caching allocator holds and is not using (memory_reserved - memory_allocated). Late in a long session that cache is real: in the CUDA job on a V100, test_auto_batch_rises_with_free_memory chose 4 at a simulated 8 MiB free, because the stale cache rather than the simulated free memory set the budget. Patching memory_reserved to match memory_allocated pins the cache at zero, so the tests measure the arithmetic they claim to. Reproduced locally by leaving 254 MiB of unreleasable cache before the class runs: the old tests fail, these pass. Pinning the cache at zero would leave that term untested, so test_cached_memory_counts_as_free covers it: with no free memory and 1 GiB cached, a 32 MiB product is still batched to the cap. It fails if the budget ignores the cache. --- .../pt_expt/model/test_hessian_hvp_batch.py | 49 ++++++++++++++++--- 1 file changed, 41 insertions(+), 8 deletions(-) diff --git a/source/tests/pt_expt/model/test_hessian_hvp_batch.py b/source/tests/pt_expt/model/test_hessian_hvp_batch.py index ec495918f2..f69aae3011 100644 --- a/source/tests/pt_expt/model/test_hessian_hvp_batch.py +++ b/source/tests/pt_expt/model/test_hessian_hvp_batch.py @@ -745,15 +745,31 @@ def probe(): return probe + def _free(self, monkeypatch, free_mib: int) -> None: + """Make ``free_mib`` the only memory the automatic choice can see. + + The budget also counts what the caching allocator holds but is not + using, and late in a long test session that cache can dwarf the free + memory being simulated. Reporting the reservation as exactly what is + allocated makes that cache zero and keeps the device's real state out + of the arithmetic under test. + """ + monkeypatch.setattr( + torch.cuda, "mem_get_info", lambda *a, **k: (free_mib * self.MIB, 0) + ) + monkeypatch.setattr( + torch.cuda, + "memory_reserved", + lambda *a, **k: torch.cuda.memory_allocated(*a, **k), + ) + @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) - ) + self._free(monkeypatch, free_mib) batch = mm._auto_hvp_batch(env.DEVICE, self._probe(32 * self.MIB)) assert 1 <= batch <= mm.DP_HESSIAN_HVP_BATCH_CAP @@ -764,17 +780,34 @@ 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), - ) + self._free(monkeypatch, free_mib) 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 + @pytest.mark.skipif( + not torch.cuda.is_available(), reason="the automatic batch needs CUDA" + ) + def test_cached_memory_counts_as_free(self, monkeypatch) -> None: + """Memory the caching allocator holds but is not using is usable. + + The driver reports it as taken, so pricing on ``mem_get_info`` alone + would starve the batch on a device whose memory sits in PyTorch's own + cache. With no free memory and 1 GiB cached, a 32 MiB product must + still be batched. + """ + self._free(monkeypatch, 0) + cached = 1024 * self.MIB + monkeypatch.setattr( + torch.cuda, + "memory_reserved", + lambda *a, **k: torch.cuda.memory_allocated(*a, **k) + cached, + ) + batch = mm._auto_hvp_batch(env.DEVICE, self._probe(32 * self.MIB)) + assert batch == 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."""