Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
286 changes: 283 additions & 3 deletions deepmd/pt_expt/model/make_model.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
# SPDX-License-Identifier: LGPL-3.0-or-later
import logging
import math
import types
from typing import (
Expand All @@ -24,6 +25,9 @@
compact_nodes,
expand_node_values,
)
from deepmd.pt.utils.auto_batch_size import (
AutoBatchSize,
)
from deepmd.pt_expt.common import (
auto_wrapped_class,
torch_module,
Expand All @@ -32,6 +36,11 @@
fused_energy_force_enabled,
fused_operators_enabled,
)
from deepmd.pt_expt.utils.env import (
DP_HESSIAN_HVP_BATCH,
DP_HESSIAN_HVP_BATCH_CAP,
DP_HESSIAN_HVP_MEMORY_FRACTION,
)
from deepmd.pt_expt.utils.graph_builder import (
build_neighbor_graph_for_method,
build_ragged_neighbor_graph,
Expand Down Expand Up @@ -376,6 +385,226 @@ def __call__(self, coord_flat: torch.Tensor) -> torch.Tensor:
return atom_out.sum(dim=0).reshape(-1)[self.ci]


log = logging.getLogger(__name__)


def _hessian_graph_batched_hvp(
model: Any,
kk: str,
ci: int,
nloc: int,
coord_flat: torch.Tensor,
atype: torch.Tensor,
box: torch.Tensor | None,
method: str,
pair_excl: Any,
rcut: float,
fparam: torch.Tensor | None,
aparam: torch.Tensor | None,
spin: torch.Tensor | None,
charge_spin: torch.Tensor | None,
batch: int,
create_graph: bool,
max_rows: int | None = None,
) -> torch.Tensor:
"""Hessian of one reduced-output component via batched Hessian-vector products.

Exactly equivalent to ``torch.autograd.functional.hessian`` on
:class:`_WrapperForwardEnergyGraph`, but evaluates ``batch`` rows per
second-order backward instead of one. The structure is replicated ``batch``
times along the frame axis the carry-all graph already has; frames are
independent, so the energy of the replicated system is a sum of independent
terms and its Hessian is block diagonal. One forward and one first-order
backward build a graph that every chunk of seed vectors then reuses.

No ``vmap`` is involved, so custom autograd Functions that lack a batching
rule keep working, and no approximation is made: in float64 the result
matches the one-row-at-a-time path to 2.8e-14 relative RMS on the
DPA-4.0.1-Pro-MPtrj checkpoint (1e-16 on the smaller example model), against
a 6.4e-16 asymmetry in the unbatched path's own answer. In float32 the two
differ by ~3e-6 relative RMS, a few times the unbatched path's own ~5e-7
asymmetry -- the longer summation chain, not a different computation.

``max_rows`` stops after that many leading rows. Only the memory probe
uses it: pricing one Hessian-vector product must not pay for the whole
Hessian first.

Returns the ``(nloc * 3, nloc * 3)`` Hessian for output component ``ci``
(or its first ``max_rows`` rows).
"""
ndof = nloc * 3
nb = max(1, min(batch, ndof))
if create_graph and coord_flat.requires_grad:
# Mirror ``torch.autograd.functional._grad_preprocess``: with
# ``create_graph`` the caller means to differentiate the Hessian
# itself, so the path back to the coordinates has to survive.
# Detaching would silently cut it and leave only the parameter path.
x = coord_flat.reshape(1, ndof).expand(nb, ndof).contiguous()
else:
x = coord_flat.detach().reshape(1, ndof).expand(nb, ndof).contiguous()
x = x.requires_grad_(True)
atype_b = atype.reshape(1, nloc).expand(nb, nloc).contiguous()
box_b = (
box.reshape(1, -1).expand(nb, box.numel()).contiguous()
if box is not None
else None
)
graph = build_neighbor_graph_for_method(
method, x.reshape(nb, nloc, 3), atype_b, box_b, rcut, pair_excl
)
atomic_ret = model.atomic_model.forward_common_atomic_graph(
graph,
atype_b.reshape(-1),
fparam=fparam.expand(nb, -1).contiguous() if fparam is not None else None,
aparam=aparam.repeat(nb, 1) if aparam is not None else None,
spin=spin.repeat(nb, 1) if spin is not None else None,
charge_spin=charge_spin.expand(nb, -1).contiguous()
if charge_spin is not None
else None,
)
# flat (nb * nloc, *def) -> one scalar per replica, summed: the replicas are
# independent, so d/dx_b only sees replica b.
total = atomic_ret[kk].reshape(nb, nloc, -1)[..., ci].sum()
wanted = ndof if max_rows is None else min(max_rows, ndof)
if not total.requires_grad:
# The reduced output does not depend on the coordinates at all, so both
# derivatives are zero. autograd refuses to differentiate an output
# that carries no graph, so this has to be caught before the first
# ``grad`` rather than after it.
return x.new_zeros(wanted, ndof)

(grad,) = torch.autograd.grad(
Comment thread
iProzd marked this conversation as resolved.
total, x, create_graph=True, allow_unused=True, materialize_grads=True
)

if not grad.requires_grad:
# The output is linear in the coordinates: the first derivative exists
# but is constant, so every second derivative is zero.
# ``functional.hessian`` materialises that zero block under its default
# ``strict=False``; differentiating a constant ``grad`` again would
# instead raise, turning a legitimate model into a crash.
return x.new_zeros(wanted, ndof)

rows: list[torch.Tensor] = []
for start in range(0, wanted, nb):
stop = min(start + nb, wanted)
# One seed block per chunk rather than one ``ndof x ndof`` identity:
# the identity is the very allocation the batching exists to avoid, and
# a fresh block per chunk keeps ``create_graph`` from retaining a buffer
# that a later chunk would overwrite.
seeds = torch.zeros(nb, ndof, dtype=x.dtype, device=x.device)
seeds[
torch.arange(stop - start, dtype=torch.int64, device=x.device),
torch.arange(start, stop, dtype=torch.int64, device=x.device),
] = 1
(hvp,) = torch.autograd.grad(
grad,
x,
grad_outputs=seeds,
retain_graph=True,
create_graph=create_graph,
allow_unused=True,
materialize_grads=True,
)
rows.append(hvp)
# Padding rows carry a zero seed and sit at the end of the final chunk, so
# one slice drops both them and anything past ``max_rows``.
return torch.cat(rows)[:wanted]


def _hvp_replica_cost(device: torch.device, probe: Any) -> int:
"""Peak memory, in bytes, that one Hessian-vector product costs."""
torch.cuda.synchronize(device)
torch.cuda.reset_peak_memory_stats(device)
base = torch.cuda.memory_allocated(device)
probe()
torch.cuda.synchronize(device)
return max(torch.cuda.max_memory_allocated(device) - base, 1)


def _auto_hvp_batch(device: torch.device, probe: Any) -> int:
"""Pick a batch from what one Hessian-vector product costs and what is free.

Peak memory is linear in the batch, so the batch that fits is the one worth
taking: past it the run dies, and below it the device is idle. Pricing a
single product costs one row out of ``3 * nloc`` and needs no per-model
constants, which a fitted memory model would.

The measured cost covers a whole product, while each batch step beyond the
first adds only its marginal share, so this reads high and the batch comes
out conservative -- the direction to err, since
:func:`_hessian_graph_row_block` can recover from a batch that turns out too
large but nothing recovers the time lost to one that was too small.
"""
if device.type != "cuda":
# Without allocator introspection there is nothing to size against, and
# without a recoverable out-of-memory error nothing to catch if the
# guess is wrong. Stay on the one-row-at-a-time path.
return 1
cost = _hvp_replica_cost(device, probe)
free, _total = torch.cuda.mem_get_info(device)
# Blocks the caching allocator holds but is not using are free to us even
# though the driver counts them as taken.
reusable = torch.cuda.memory_reserved(device) - torch.cuda.memory_allocated(device)
budget = (free + reusable) * DP_HESSIAN_HVP_MEMORY_FRACTION
return max(1, min(int(budget // cost), DP_HESSIAN_HVP_BATCH_CAP))


def _hessian_graph_row_block(
batch: int,
wrapper: Any,
coord_flat: torch.Tensor,
create_graph: bool,
**kwargs: Any,
) -> tuple[torch.Tensor, int]:
"""One output component's Hessian, halving the batch if memory runs out.

Batching changes how the Hessian is computed, never what it is, so a batch
that does not fit can simply be retried smaller instead of ending the run.
Reaching 1 hands over to the original one-row-at-a-time path rather than to
this helper with a batch of one, so the fallback bottoms out in exactly the
code a user who set 1 would have taken.

The out-of-memory test is the repository's ``is_oom_error``: the allocator
failure does not always arrive as ``torch.OutOfMemoryError`` -- AOTInductor
rewraps it in a plain ``RuntimeError`` -- and a catch keyed on the
exception type alone lets the wrapped form end the run.

Returns the Hessian and the batch that fit, so the caller starts the next
component there instead of re-climbing the ladder from the top: each failed
attempt is a full forward plus a first-order backward, spent in exactly the
regime the fallback exists for.
"""
while batch > 1:
try:
return (
_hessian_graph_batched_hvp(
coord_flat=coord_flat,
batch=batch,
create_graph=create_graph,
**kwargs,
),
batch,
)
except Exception as e:
if not AutoBatchSize(silent=True).is_oom_error(e):
raise
batch //= 2
Comment thread
iProzd marked this conversation as resolved.
log.warning(
"Hessian-vector product batch did not fit in memory; retrying "
"with %d. Set DP_HESSIAN_HVP_BATCH to choose it yourself.",
batch,
)
return (
torch.autograd.functional.hessian(
wrapper,
coord_flat,
create_graph=create_graph,
),
1,
)


def _cal_hessian_ext_graph(
model: Any,
kk: str,
Expand Down Expand Up @@ -413,6 +642,11 @@ def _cal_hessian_ext_graph(
if charge_spin is not None and charge_spin.ndim == 1
else charge_spin
)
# Priced per frame: within one frame every output component shares the
# neighbour count, but each frame has its own, so a batch priced on one
# frame does not price another.
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)
Expand All @@ -427,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:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Not blocking, but worth a decision in the code. auto_batch is fixed at line 649, so in the default configuration every frame overwrites hvp_batch with a fresh probe result, discarding whatever the previous frame's OOM ladder found. That is the batch commit 3049b57 was added to preserve: on a multi-frame call whose true ceiling is below the probe's estimate, each frame re-climbs 8 -> 4 -> 2 and pays a forward plus a first-order backward per failed attempt. test_the_surviving_batch_is_kept_for_the_next_frame does not see this because it pins DP_HESSIAN_HVP_BATCH=8, so it exercises only the explicit-value path.

I asked for per-frame pricing, so this is a trade-off between my two earlier points rather than a defect, and repricing per frame is defensible since frames genuinely differ. Either seed the probe with the previous frame's surviving batch (for instance min(probe, hvp_batch) once a ladder has reduced it) or state in the comment that a per-frame re-climb is accepted, and add a test in auto mode either way.

# 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,
Expand All @@ -443,10 +720,13 @@ def _cal_hessian_ext_graph(
spin=spin_frame,
charge_spin=charge_spin_frame,
)
hess = torch.autograd.functional.hessian(
wrapper,
coord_flat,
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:
component_index = (
Expand Down
34 changes: 34 additions & 0 deletions deepmd/pt_expt/utils/env.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,40 @@

SAMPLER_RECORD = os.environ.get("SAMPLER_RECORD", False)
DP_DTYPE_PROMOTION_STRICT = os.environ.get("DP_DTYPE_PROMOTION_STRICT", "0") == "1"
# Number of Hessian rows evaluated per second-order backward pass. The Hessian
# is built from Hessian-vector products; batching them over replicated frames
# trades memory for far fewer kernel launches. 1 keeps the one-row-at-a-time
# path and reproduces the pre-batching behaviour exactly.
#
# Left unset, the batch is chosen per 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 = (
Comment thread
iProzd marked this conversation as resolved.
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
Comment thread
iProzd marked this conversation as resolved.
# 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
Comment thread
iProzd marked this conversation as resolved.
try:
# only linux
ncpus = len(os.sched_getaffinity(0))
Expand Down
Loading
Loading