Skip to content

Navigation Menu

Sign in
Sign up

fix(pt_expt): reuse the stored min_nbor_dist and batch the neighbor statistics #5956

New issue

Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.

By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.

Already on GitHub? Sign in to your account

Open
yckbz wants to merge 15 commits into deepmodeling:master
base: master
Choose a base branch
Loading
from yckbz:fix-pt-expt-min-nbor-dist
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
15 commits
Select commit Hold shift + click to select a range
15a0b34
fix(pt_expt): read min_nbor_dist from @variables
yckbz Aug 4, 2026
110e60b
fix(pt_expt): batch the neighbor statistics
yckbz Aug 4, 2026
6681b32
Potential fix for pull request finding
yckbz Aug 4, 2026
12e34ad
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Aug 4, 2026
7d14eae
Merge branch 'master' into fix-pt-expt-min-nbor-dist
yckbz Aug 5, 2026
0c88b50
fix(pt_expt): follow the selected device when auto-batching
yckbz Aug 11, 2026
d850ca9
fix(pt_expt): keep min_nbor_dist in the graph-lower output
yckbz Aug 11, 2026
f1423c1
fix(pt_expt): assign min_nbor_dist through the spin backbone
yckbz Aug 11, 2026
cbb5117
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Aug 11, 2026
d491dd3
Merge branch 'master' into fix-pt-expt-min-nbor-dist
yckbz Aug 21, 2026
fcd067c
Merge remote-tracking branch 'origin/fix-pt-expt-min-nbor-dist'
yckbz Aug 21, 2026
c5c6b31
Merge branch 'master' into fix-pt-expt-min-nbor-dist
OutisLi Aug 28, 2026
913f68a
fix(pt_expt): read DEVICE at call time in the auto batch size
yckbz Aug 29, 2026
9a914be
test(pt_expt): import the compress module one way
yckbz Aug 29, 2026
2c69cff
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Aug 29, 2026
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

Some comments aren't visible on the classic Files Changed page.

7 changes: 7 additions & 0 deletions deepmd/main.py
View file Open in desktop
Original file line number Diff line number Diff line change
Expand Up @@ -668,6 +668,13 @@ def main_parser() -> argparse.ArgumentParser:
default=None,
help="The training script of the input frozen model",
)
parser_compress.add_argument(
"--recompute-min-nbor-dist",
action="store_true",
help="(Supported backend: PyTorch Exportable) Ignore the minimal neighbor "
"distance saved in the model and recompute it from the training data. "
"Requires -t,--training-script",
)
parser_compress.add_argument(
"--head",
"--model-branch",
Expand Down
82 changes: 67 additions & 15 deletions deepmd/pt_expt/entrypoints/compress.py
View file Open in desktop
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,9 @@
"""Compress a pt_expt model (.pte) by tabulating embedding nets."""

import logging
from typing import (
Any,
)

from deepmd.pt_expt.utils.serialization import (
deserialize_to_file,
Expand All @@ -11,13 +14,41 @@
log = logging.getLogger(__name__)


def _read_saved_min_nbor_dist(model: Any, model_dict: dict) -> tuple[float | None, str]:
"""Read the stored minimal neighbor distance and where it was read from.

``@variables`` is the cross-backend location of this value: it is written
by :mod:`deepmd.pt.utils.serialization` and read back by the PyTorch and
Paddle backends, so a ``.pt2`` produced by ``dp convert-backend`` carries
the value there rather than inside the serialized model dict.

Returns
-------
float or None
The stored minimal neighbor distance, None if the model has none.
str
Human-readable description of where the value was read from.
"""
min_nbor_dist = model.get_min_nbor_dist()
if min_nbor_dist is not None:
return float(min_nbor_dist), "the model"
min_nbor_dist = model_dict.get("min_nbor_dist")
if min_nbor_dist is not None:
return float(min_nbor_dist), "the model file"
min_nbor_dist = (model_dict.get("@variables") or {}).get("min_nbor_dist")
if min_nbor_dist is not None:
return float(min_nbor_dist), "the model file (@variables)"
return None, ""


def enable_compression(
input_file: str,
output: str,
stride: float = 0.01,
extrapolate: int = 5,
check_frequency: int = -1,
training_script: str | None = None,
recompute_min_nbor_dist: bool = False,
) -> None:
"""Compress a .pte model by tabulating embedding nets.

Expand All @@ -36,6 +67,9 @@ def enable_compression(
training_script : str or None
Path to training script, used to compute min_nbor_dist if not
stored in the model.
recompute_min_nbor_dist : bool
Ignore the min_nbor_dist stored in the model and recompute it from
the training data. Requires training_script.
"""
from deepmd.pt_expt.model.model import (
BaseModel,
Expand All @@ -48,26 +82,40 @@ def enable_compression(
# 2. Get or compute min_nbor_dist. Measuring it is a dense all-pairs pass
# over the training data, so it is only run for models that tabulate
# from the shortest observed distance.
min_nbor_dist = model.get_min_nbor_dist()
if min_nbor_dist is None:
min_nbor_dist = model_dict.get("min_nbor_dist")
if min_nbor_dist is None and not model.compression_needs_min_nbor_dist():
if recompute_min_nbor_dist:
min_nbor_dist, source = None, ""
else:
min_nbor_dist, source = _read_saved_min_nbor_dist(model, model_dict)
Comment thread
yckbz marked this conversation as resolved.
Comment thread
yckbz marked this conversation as resolved.
if min_nbor_dist is not None:
log.info(f"Minimal neighbor distance read from {source}: {min_nbor_dist:f}")
elif not model.compression_needs_min_nbor_dist():
log.info(
"The model tabulates over an analytically bounded domain; "
"skipping the neighbor statistics."
)
min_nbor_dist = 0.0
if min_nbor_dist is None:
log.info(
"Minimal neighbor distance is not saved in the model, "
"compute it from the training data."
)
if training_script is None:
raise ValueError(
"The model does not have a minimum neighbor distance, "
"so the training script and data must be provided "
"(via -t,--training-script)."
else:
if recompute_min_nbor_dist:
log.info(
"Recompute the minimal neighbor distance from the training data, "
"ignoring the one saved in the model."
)
if training_script is None:
raise ValueError(
"Recomputing the minimal neighbor distance requires the "
"training script and data (via -t,--training-script)."
)
else:
log.info(
"Minimal neighbor distance is not saved in the model, "
"compute it from the training data."
)
if training_script is None:
raise ValueError(
"The model does not have a minimum neighbor distance, "
"so the training script and data must be provided "
"(via -t,--training-script)."
)
from deepmd.common import (
j_loader,
)
Expand Down Expand Up @@ -134,7 +182,11 @@ def enable_compression(
log.info("Re-exporting compressed graph...")
deserialize_to_file(
output,
{"model": compressed_model_dict, "model_def_script": model_def_script},
{
"model": compressed_model_dict,
"model_def_script": model_def_script,
"min_nbor_dist": float(min_nbor_dist),
},
lower_kind="auto",
)
else:
Expand Down
1 change: 1 addition & 0 deletions deepmd/pt_expt/entrypoints/main.py
View file Open in desktop
Original file line number Diff line number Diff line change
Expand Up @@ -939,6 +939,7 @@ def main(args: list[str] | argparse.Namespace | None = None) -> None:
extrapolate=FLAGS.extrapolate,
check_frequency=FLAGS.frequency,
training_script=FLAGS.training_script,
recompute_min_nbor_dist=FLAGS.recompute_min_nbor_dist,
)
else:
raise RuntimeError(
Expand Down
14 changes: 14 additions & 0 deletions deepmd/pt_expt/model/spin_model.py
View file Open in desktop
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,20 @@ def __getattr__(self, name: str) -> Any:
return getattr(backbone, name)
raise AttributeError(name)

@property
def min_nbor_dist(self) -> float | None:
"""Minimal neighbor distance, stored on the backbone model.

``__getattr__`` only delegates reads, so without this property an
assignment would land on the wrapper while ``get_min_nbor_dist`` and
``enable_compression`` keep reading the backbone.
"""
return self.backbone_model.min_nbor_dist

@min_nbor_dist.setter
def min_nbor_dist(self, value: float | None) -> None:
self.backbone_model.min_nbor_dist = value

def forward_common_lower_exportable(
self,
extended_coord: torch.Tensor,
Expand Down
27 changes: 27 additions & 0 deletions deepmd/pt_expt/utils/auto_batch_size.py
View file Open in desktop
Original file line number Diff line number Diff line change
@@ -0,0 +1,27 @@
# SPDX-License-Identifier: LGPL-3.0-or-later

from deepmd.pt.utils.auto_batch_size import AutoBatchSize as AutoBatchSizeBase
from deepmd.pt_expt.utils import (
env,
)


class AutoBatchSize(AutoBatchSizeBase):
"""Auto batch size following the device pt_expt runs on.

``env.DEVICE`` is CPU whenever ``DEVICE=cpu`` is set, even on a CUDA host.
Growing the batch there risks a host OOM, which the CUDA-OOM handler
cannot recover from, so the growth policy follows the selected device
rather than CUDA availability. The module is held rather than the value
so that the check reads the device in effect at call time.
"""

def is_gpu_available(self) -> bool:
"""Check if the selected device is a GPU.

Returns
-------
bool
True if pt_expt runs on a CUDA device
"""
return env.DEVICE.type == "cuda"
15 changes: 14 additions & 1 deletion deepmd/pt_expt/utils/neighbor_stat.py
View file Open in desktop
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,9 @@
from deepmd.pt_expt.common import (
torch_module,
)
from deepmd.pt_expt.utils.auto_batch_size import (
AutoBatchSize,
)
from deepmd.pt_expt.utils.env import (
DEVICE,
GLOBAL_PT_FLOAT_PRECISION,
Expand All @@ -28,6 +31,12 @@ class NeighborStatOP(NeighborStatOPDP):
class NeighborStat(BaseNeighborStat):
"""Neighbor statistics using torch on DEVICE.

The statistics are evaluated one frame batch at a time. The intermediate
tensor is of shape ``[nframes, nloc, nall, 3]``, so processing a whole set
at once needs hundreds of GiB for a large set; :class:`AutoBatchSize`
keeps the batch within the available device memory, as the PyTorch,
Paddle, JAX and TensorFlow backends already do.

Parameters
----------
ntypes : int
Expand All @@ -46,6 +55,7 @@ def __init__(
) -> None:
super().__init__(ntypes, rcut, mixed_type)
self.op = NeighborStatOP(ntypes, rcut, mixed_type)
self.auto_batch_size = AutoBatchSize()
Comment thread
yckbz marked this conversation as resolved.

def iterator(
self, data: DeepmdDataSystem
Expand All @@ -65,7 +75,10 @@ def iterator(
for jj in data.data_systems[ii].dirs:
data_set = data.data_systems[ii]
data_set_data = data_set._load_set(jj)
minrr2, max_nnei = self._execute(
minrr2, max_nnei = self.auto_batch_size.execute_all(
self._execute,
data_set_data["coord"].shape[0],
data_set.get_natoms(),
data_set_data["coord"],
data_set_data["type"],
data_set_data["box"] if data_set.pbc else None,
Expand Down
29 changes: 29 additions & 0 deletions source/tests/pt_expt/model/test_model_compression.py
View file Open in desktop
Original file line number Diff line number Diff line change
Expand Up @@ -187,6 +187,35 @@ def test_min_nbor_dist_roundtrip(self) -> None:
finally:
os.unlink(frozen_path)

def test_compress_min_nbor_dist_from_variables(self) -> None:
"""Test that compress recovers min_nbor_dist from @variables and keeps it.

``dp convert-backend`` stores the value under ``@variables``. Compress
must read it from there — without a training script it would otherwise
raise — and carry it into the compressed archive.
"""
from deepmd.pt_expt.entrypoints.compress import (
enable_compression,
)

md = self._make_model()
md.eval()

model_data = {"model": md.serialize(), "@variables": {"min_nbor_dist": 0.5}}
with tempfile.NamedTemporaryFile(suffix=".pte", delete=False) as f:
frozen_path = f.name
with tempfile.NamedTemporaryFile(suffix=".pte", delete=False) as f:
compressed_path = f.name
try:
deserialize_to_file(frozen_path, model_data)
enable_compression(input_file=frozen_path, output=compressed_path)
compressed_data = serialize_from_file(compressed_path)
self.assertAlmostEqual(compressed_data["min_nbor_dist"], 0.5)
finally:
os.unlink(frozen_path)
if os.path.exists(compressed_path):
os.unlink(compressed_path)

def test_compress_state_serialized(self) -> None:
"""Test that compression state persists through serialize/deserialize.

Expand Down
Loading

AltStyle によって変換されたページ (->オリジナル) /