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
74 changes: 53 additions & 21 deletions deepmd/infer/deep_eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,48 @@
import ase.neighborlist


def _standardize_fparam_aparam(
fparam: np.ndarray | list | None,
aparam: np.ndarray | list | None,
nframes: int,
natoms: int,
dim_fparam: int,
dim_aparam: int,
) -> tuple[np.ndarray | None, np.ndarray | None]:
"""Normalize documented parameter shorthand to frame-major arrays.

This normalization must happen before automatic batching. In particular,
an ``(natoms, dim_aparam)`` shared atomic parameter has an atom axis first;
a batcher would otherwise mistake that axis for frames and slice it.
"""
if fparam is not None:
fparam = np.asarray(fparam)
if fparam.size == nframes * dim_fparam:
fparam = fparam.reshape(nframes, dim_fparam)
elif fparam.size == dim_fparam:
fparam = np.tile(fparam.reshape(1, dim_fparam), (nframes, 1))
else:
raise RuntimeError(
"got wrong size of frame param, should be either "
f"{nframes} x {dim_fparam} or {dim_fparam}"
)
if aparam is not None:
aparam = np.asarray(aparam)
if aparam.size == nframes * natoms * dim_aparam:
aparam = aparam.reshape(nframes, natoms, dim_aparam)
elif aparam.size == natoms * dim_aparam:
aparam = np.tile(aparam.reshape(1, natoms, dim_aparam), (nframes, 1, 1))
elif aparam.size == dim_aparam:
aparam = np.tile(aparam.reshape(1, 1, dim_aparam), (nframes, natoms, 1))
else:
raise RuntimeError(
"got wrong size of atomic param, should be either "
f"{nframes} x {natoms} x {dim_aparam} or "
f"{natoms} x {dim_aparam} or {dim_aparam}"
)
return fparam, aparam


class DeepEvalBackend(ABC):
"""Low-level Deep Evaluator interface.

Expand Down Expand Up @@ -947,28 +989,18 @@ def _standard_input(
coords = coords.reshape(nframes, natoms, 3)
if cells is not None:
cells = cells.reshape(nframes, 3, 3)
if fparam is not None:
fdim = self.get_dim_fparam()
if fparam.size == nframes * fdim:
fparam = np.reshape(fparam, [nframes, fdim])
elif fparam.size == fdim:
fparam = np.tile(fparam.reshape([-1]), [nframes, 1])
else:
raise RuntimeError(
f"got wrong size of frame param, should be either {nframes} x {fdim} or {fdim}"
)
fparam, aparam = _standardize_fparam_aparam(
fparam,
aparam,
nframes,
natoms,
self.get_dim_fparam(),
self.get_dim_aparam(),
)
if aparam is not None:
fdim = self.get_dim_aparam()
if aparam.size == nframes * natoms * fdim:
aparam = np.reshape(aparam, [nframes, natoms * fdim])
elif aparam.size == natoms * fdim:
aparam = np.tile(aparam.reshape([-1]), [nframes, 1])
elif aparam.size == fdim:
aparam = np.tile(aparam.reshape([-1]), [nframes, natoms])
else:
raise RuntimeError(
f"got wrong size of frame param, should be either {nframes} x {natoms} x {fdim} or {natoms} x {fdim} or {fdim}"
)
# Preserve the historical flattened backend ABI used by the public
# wrapper; backend adapters normalize it back to frame-major 3-D.
aparam = aparam.reshape(nframes, natoms * self.get_dim_aparam())
return coords, cells, atom_types, fparam, aparam, nframes, natoms

def get_sel_type(self) -> list[int]:
Expand Down
9 changes: 9 additions & 0 deletions deepmd/jax/infer/deep_eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@
from deepmd.infer.deep_eval import DeepEval as DeepEvalWrapper
from deepmd.infer.deep_eval import (
DeepEvalBackend,
_standardize_fparam_aparam,
)
from deepmd.infer.deep_polar import (
DeepPolar,
Expand Down Expand Up @@ -278,6 +279,14 @@ def eval(
natoms, numb_test = self._get_natoms_and_nframes(
coords, atom_types, len(atom_types.shape) > 1
)
fparam, aparam = _standardize_fparam_aparam(
fparam,
aparam,
numb_test,
natoms,
self.get_dim_fparam(),
self.get_dim_aparam(),
)
request_defs = self._get_request_defs(atomic)
out = self._eval_func(self._eval_model, numb_test, natoms)(
coords, cells, atom_types, fparam, aparam, charge_spin, request_defs
Expand Down
9 changes: 9 additions & 0 deletions deepmd/pd/infer/deep_eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
from deepmd.infer.deep_eval import DeepEval as DeepEvalWrapper
from deepmd.infer.deep_eval import (
DeepEvalBackend,
_standardize_fparam_aparam,
)
from deepmd.infer.deep_polar import (
DeepGlobalPolar,
Expand Down Expand Up @@ -376,6 +377,14 @@ def eval(
natoms, numb_test = self._get_natoms_and_nframes(
coords, atom_types, len(atom_types.shape) > 1
)
fparam, aparam = _standardize_fparam_aparam(
fparam,
aparam,
numb_test,
natoms,
self.get_dim_fparam(),
self.get_dim_aparam(),
)
request_defs = self._get_request_defs(atomic)
if "spin" not in kwargs or kwargs["spin"] is None:
out = self._eval_func(self._eval_model, numb_test, natoms)(
Expand Down
20 changes: 20 additions & 0 deletions deepmd/pt/infer/deep_eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
from deepmd.infer.deep_eval import DeepEval as DeepEvalWrapper
from deepmd.infer.deep_eval import (
DeepEvalBackend,
_standardize_fparam_aparam,
)
from deepmd.infer.deep_polar import (
DeepGlobalPolar,
Expand Down Expand Up @@ -544,6 +545,14 @@ def eval(
natoms, numb_test = self._get_natoms_and_nframes(
coords, atom_types, len(atom_types.shape) > 1
)
fparam, aparam = _standardize_fparam_aparam(
fparam,
aparam,
numb_test,
natoms,
self.get_dim_fparam(),
self.get_dim_aparam(),
)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
request_defs = self._get_request_defs(atomic)
if "spin" not in kwargs or kwargs["spin"] is None:
out = self._eval_func(self._eval_model, numb_test, natoms)(
Expand Down Expand Up @@ -1312,6 +1321,17 @@ def eval_embedding(
natoms, numb_test = self._get_natoms_and_nframes(
coords, atom_types, len(atom_types.shape) > 1
)
# Normalize shared parameter shorthand before auto batching. Otherwise
# a one-dimensional fparam/aparam is passed unchanged to every split,
# and _eval_embedding cannot reshape it to the split frame count.
fparam, aparam = _standardize_fparam_aparam(
fparam,
aparam,
numb_test,
natoms,
self.get_dim_fparam(),
self.get_dim_aparam(),
)
return self._eval_func(self._eval_embedding, numb_test, natoms)(
coords, cells, atom_types, fparam, aparam, charge_spin, dtype
)
Expand Down
9 changes: 9 additions & 0 deletions deepmd/tf2/infer/deep_eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
from deepmd.infer.deep_eval import DeepEval as DeepEvalWrapper
from deepmd.infer.deep_eval import (
DeepEvalBackend,
_standardize_fparam_aparam,
)
from deepmd.infer.deep_polar import (
DeepPolar,
Expand Down Expand Up @@ -292,6 +293,14 @@ def eval(
natoms, numb_test = self._get_natoms_and_nframes(
coords, atom_types, len(atom_types.shape) > 1
)
fparam, aparam = _standardize_fparam_aparam(
fparam,
aparam,
numb_test,
natoms,
self.get_dim_fparam(),
self.get_dim_aparam(),
)
request_defs = self._get_request_defs(atomic)
out = self._eval_func(self._eval_model, numb_test, natoms)(
coords, cells, atom_types, fparam, aparam, request_defs
Expand Down
100 changes: 100 additions & 0 deletions source/tests/common/test_deep_eval_parameter_shorthand.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,100 @@
# SPDX-License-Identifier: LGPL-3.0-or-later
"""Tests for backend-level DeepEval parameter normalization."""

import numpy as np
import pytest

from deepmd.infer.deep_eval import (
_standardize_fparam_aparam,
)

NFRAMES = 3
NATOMS = 4
DIM_FPARAM = 2
DIM_APARAM = 2
FPARAM = np.array([0.25, -0.5], dtype=np.float64)
APARAM_PER_ATOM = np.arange(NATOMS * DIM_APARAM, dtype=np.float64).reshape(
NATOMS, DIM_APARAM
)
APARAM_ALL_ATOMS = np.array([0.3, -0.2], dtype=np.float64)


@pytest.mark.parametrize(
("fparam", "expected"),
[
(FPARAM.tolist(), np.tile(FPARAM, (NFRAMES, 1))),
(
np.arange(NFRAMES * DIM_FPARAM).reshape(NFRAMES, DIM_FPARAM),
np.arange(NFRAMES * DIM_FPARAM).reshape(NFRAMES, DIM_FPARAM),
),
],
ids=("shared", "per-frame"),
)
def test_standardize_fparam(fparam, expected) -> None:
"""Frame parameters become a canonical frame-major matrix."""
actual, _ = _standardize_fparam_aparam(
fparam,
None,
NFRAMES,
NATOMS,
DIM_FPARAM,
DIM_APARAM,
)

np.testing.assert_array_equal(actual, expected)


@pytest.mark.parametrize(
("aparam", "expected"),
[
(
APARAM_PER_ATOM,
np.tile(APARAM_PER_ATOM, (NFRAMES, 1, 1)),
),
(
APARAM_ALL_ATOMS.tolist(),
np.tile(APARAM_ALL_ATOMS, (NFRAMES, NATOMS, 1)),
),
(
np.arange(NFRAMES * NATOMS * DIM_APARAM).reshape(
NFRAMES, NATOMS, DIM_APARAM
),
np.arange(NFRAMES * NATOMS * DIM_APARAM).reshape(
NFRAMES, NATOMS, DIM_APARAM
),
),
],
ids=("shared-per-atom", "shared-all-atoms", "per-frame"),
)
def test_standardize_aparam(aparam, expected) -> None:
"""Atomic shorthand is expanded before a batcher can slice its atom axis."""
_, actual = _standardize_fparam_aparam(
None,
aparam,
NFRAMES,
NATOMS,
DIM_FPARAM,
DIM_APARAM,
)

np.testing.assert_array_equal(actual, expected)


@pytest.mark.parametrize(
("fparam", "aparam", "message"),
[
(np.zeros(3), None, "wrong size of frame param"),
(None, np.zeros(3), "wrong size of atomic param"),
],
)
def test_invalid_parameter_size_is_rejected(fparam, aparam, message) -> None:
"""Report the documented contract instead of a backend reshape failure."""
with pytest.raises(RuntimeError, match=message):
_standardize_fparam_aparam(
fparam,
aparam,
NFRAMES,
NATOMS,
DIM_FPARAM,
DIM_APARAM,
)
68 changes: 68 additions & 0 deletions source/tests/consistent/io/test_io.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,15 @@ class IOTest:
# property model), skipped by the cross-backend round trips below.
skip_backends: ClassVar[set[str]] = set()

def _has_fparam_aparam(self) -> bool:
"""Whether the serialized fitting requires both parameter families."""
fitting = self.data.get("model_def_script", {}).get("fitting_net", {})
return (
isinstance(fitting, dict)
and fitting.get("numb_fparam", 0) > 0
and fitting.get("numb_aparam", 0) > 0
)

def get_data_from_model(self, model_file: str) -> dict:
"""Get data from a model file.

Expand Down Expand Up @@ -156,6 +165,7 @@ def test_deep_eval(self) -> None:
("jax", 2) if DP_TEST_TF2_ONLY else ("tensorflow", 0),
("tf2", 0) if DP_TEST_TF2_ONLY else (None, None),
("pytorch", 0),
("paddle", 1) if self._has_fparam_aparam() else (None, None),
("dpmodel", 0),
("jax", 0) if DP_TEST_TF2_ONLY else (None, None),
):
Expand All @@ -181,6 +191,10 @@ def test_deep_eval(self) -> None:
aparam = np.ones((nframes, natoms, deep_eval.get_dim_aparam()))
else:
aparam = None
if backend_name in {"pytorch", "jax", "tf2", "paddle"} and (
deep_eval.get_dim_fparam() > 0 and deep_eval.get_dim_aparam() > 0
):
self._assert_backend_parameter_shorthand(model_file, deep_eval)
ret = deep_eval.eval(
self.coords,
self.box,
Expand Down Expand Up @@ -238,6 +252,60 @@ def test_deep_eval(self) -> None:
err_msg=f"backend {idx + 1} for rets_idx {rets_idx}",
)

def _assert_backend_parameter_shorthand(
self, model_file: str, deep_eval: DeepEval
) -> None:
"""Compare backend-direct shorthand with explicit frame-major inputs.

Calling ``deep_eval.deep_eval`` deliberately bypasses the public
``_standard_input`` normalization. A one-frame auto-batch size also
proves that shared per-atom parameters are expanded before the batcher
can mistake their atom axis for a frame axis.
"""
natoms = self.atype.shape[1]
nframes = 2
coords = np.repeat(self.coords, nframes, axis=0)
boxes = np.repeat(self.box, nframes, axis=0)
atom_types = self.atype.reshape(-1)
fparam_shared = np.ones(deep_eval.get_dim_fparam())
aparam_per_atom = np.ones((natoms, deep_eval.get_dim_aparam()))
fparam_full = np.tile(fparam_shared, (nframes, 1))
aparam_full = np.tile(aparam_per_atom, (nframes, 1, 1))
backend = DeepEval(model_file, auto_batch_size=natoms).deep_eval

expected = backend.eval(
coords,
boxes,
atom_types,
fparam=fparam_full,
aparam=aparam_full,
)
shorthand_cases = (
(fparam_shared.tolist(), aparam_per_atom),
(
fparam_shared,
np.ones(deep_eval.get_dim_aparam()),
),
)
for fparam, aparam in shorthand_cases:
actual = backend.eval(
coords,
boxes,
atom_types,
fparam=fparam,
aparam=aparam,
)
self.assertEqual(actual.keys(), expected.keys())
for name in actual:
np.testing.assert_allclose(
actual[name],
expected[name],
rtol=1e-12,
atol=1e-12,
equal_nan=True,
err_msg=f"backend-direct shorthand output {name}",
)


class TestDeepPot(unittest.TestCase, IOTest):
def setUp(self) -> None:
Expand Down
Loading
Loading