From 668dd1f2e7fd0c51c75f4d20ac9af55b6de9a8a7 Mon Sep 17 00:00:00 2001 From: Jinzhe Zeng Date: Wed, 8 Jul 2026 01:05:35 +0800 Subject: [PATCH 1/3] test(dpa4): enable array-api-strict consistency tests --- deepmd/dpmodel/descriptor/dpa4_nn/norm.py | 2 +- source/tests/array_api_strict/common.py | 2 + .../array_api_strict/descriptor/__init__.py | 4 + .../tests/array_api_strict/descriptor/dpa4.py | 46 +++++++++ .../array_api_strict/fitting/__init__.py | 4 + .../array_api_strict/fitting/dpa4_ener.py | 46 +++++++++ .../tests/consistent/descriptor/test_dpa4.py | 72 +++++++++++--- .../consistent/fitting/test_dpa4_ener.py | 96 +++++++++++++++---- 8 files changed, 240 insertions(+), 32 deletions(-) create mode 100644 source/tests/array_api_strict/descriptor/dpa4.py create mode 100644 source/tests/array_api_strict/fitting/dpa4_ener.py diff --git a/deepmd/dpmodel/descriptor/dpa4_nn/norm.py b/deepmd/dpmodel/descriptor/dpa4_nn/norm.py index f436f8c3dc..c6492bd561 100644 --- a/deepmd/dpmodel/descriptor/dpa4_nn/norm.py +++ b/deepmd/dpmodel/descriptor/dpa4_nn/norm.py @@ -546,7 +546,7 @@ def call(self, x: Any) -> Any: if x.ndim == 2: inv_rms = 1.0 / xp.sqrt(xp.mean(x * x, axis=-1, keepdims=True) + self.eps) x = x * inv_rms - x = x * xp_asarray_nodetach(xp, self.adam_scale[...], device=device)[0] + x = x * xp_asarray_nodetach(xp, self.adam_scale[...], device=device)[0, :] return xp.astype(x, in_dtype) inv_rms = 1.0 / xp.sqrt(xp.mean(x * x, axis=-1, keepdims=True) + self.eps) diff --git a/source/tests/array_api_strict/common.py b/source/tests/array_api_strict/common.py index 546e111dd2..b1032c4bd9 100644 --- a/source/tests/array_api_strict/common.py +++ b/source/tests/array_api_strict/common.py @@ -59,8 +59,10 @@ def to_array_api_strict_array(array: np.ndarray | None) -> Any: f"{_PACKAGE_ROOT}.descriptor.dpa2", f"{_PACKAGE_ROOT}.descriptor.repflows", f"{_PACKAGE_ROOT}.descriptor.dpa3", + f"{_PACKAGE_ROOT}.descriptor.dpa4", f"{_PACKAGE_ROOT}.descriptor.hybrid", f"{_PACKAGE_ROOT}.fitting", + f"{_PACKAGE_ROOT}.fitting.dpa4_ener", ) diff --git a/source/tests/array_api_strict/descriptor/__init__.py b/source/tests/array_api_strict/descriptor/__init__.py index 1bbefbea6f..b9235aa8b5 100644 --- a/source/tests/array_api_strict/descriptor/__init__.py +++ b/source/tests/array_api_strict/descriptor/__init__.py @@ -8,6 +8,9 @@ from .dpa3 import ( DescrptDPA3, ) +from .dpa4 import ( + DescrptDPA4, +) from .hybrid import ( DescrptHybrid, ) @@ -31,6 +34,7 @@ "DescrptDPA1", "DescrptDPA2", "DescrptDPA3", + "DescrptDPA4", "DescrptHybrid", "DescrptSeA", "DescrptSeAttenV2", diff --git a/source/tests/array_api_strict/descriptor/dpa4.py b/source/tests/array_api_strict/descriptor/dpa4.py new file mode 100644 index 0000000000..3438f85a08 --- /dev/null +++ b/source/tests/array_api_strict/descriptor/dpa4.py @@ -0,0 +1,46 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +from deepmd.dpmodel.descriptor.dpa4 import DescrptDPA4 as DescrptDPA4DP +from deepmd.dpmodel.descriptor.dpa4_nn.activation import ( + SwiGLU, +) +from deepmd.dpmodel.descriptor.dpa4_nn.grid_net import ( + GridProduct, +) +from deepmd.dpmodel.descriptor.dpa4_nn.radial import ( + BridgingSwitch, + C3CutoffEnvelope, + InnerClamp, +) +from deepmd.dpmodel.descriptor.dpa4_nn.wignerd import ( + WignerDCalculator, +) + +from ..common import ( + array_api_strict_module, + register_dpmodel_mapping, +) +from ..utils import exclude_mask as _strict_exclude_mask # noqa: F401 +from ..utils import network as _strict_network # noqa: F401 +from .base_descriptor import ( + BaseDescriptor, +) + + +@BaseDescriptor.register("SeZM") +@BaseDescriptor.register("sezm") +@BaseDescriptor.register("DPA4") +@BaseDescriptor.register("dpa4") +@array_api_strict_module +class DescrptDPA4(DescrptDPA4DP): + pass + + +for _stateless_cls in ( + BridgingSwitch, + C3CutoffEnvelope, + GridProduct, + InnerClamp, + SwiGLU, + WignerDCalculator, +): + register_dpmodel_mapping(_stateless_cls, lambda v: v) diff --git a/source/tests/array_api_strict/fitting/__init__.py b/source/tests/array_api_strict/fitting/__init__.py index 2041f600ea..81aee58175 100644 --- a/source/tests/array_api_strict/fitting/__init__.py +++ b/source/tests/array_api_strict/fitting/__init__.py @@ -1,4 +1,7 @@ # SPDX-License-Identifier: LGPL-3.0-or-later +from .dpa4_ener import ( + SeZMEnergyFittingNet, +) from .fitting import ( DipoleFittingNet, DOSFittingNet, @@ -13,4 +16,5 @@ "EnergyFittingNet", "PolarFittingNet", "PropertyFittingNet", + "SeZMEnergyFittingNet", ] diff --git a/source/tests/array_api_strict/fitting/dpa4_ener.py b/source/tests/array_api_strict/fitting/dpa4_ener.py new file mode 100644 index 0000000000..1570ec7891 --- /dev/null +++ b/source/tests/array_api_strict/fitting/dpa4_ener.py @@ -0,0 +1,46 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +from typing import ( + ClassVar, +) + +from deepmd.dpmodel.fitting.dpa4_ener import GLUFittingNet as GLUFittingNetDP +from deepmd.dpmodel.fitting.dpa4_ener import ( + SeZMEnergyFittingNet as SeZMEnergyFittingNetDP, +) +from deepmd.dpmodel.fitting.dpa4_ener import ( + SeZMNetworkCollection as SeZMNetworkCollectionDP, +) + +from ..common import ( + array_api_strict_module, + register_dpmodel_mapping, +) +from ..utils import network as _strict_network # noqa: F401 + + +@array_api_strict_module +class GLUFittingNet(GLUFittingNetDP): + pass + + +@array_api_strict_module +class SeZMNetworkCollection(SeZMNetworkCollectionDP): + NETWORK_TYPE_MAP: ClassVar[dict[str, type]] = { + "sezm_fitting_network": GLUFittingNet, + } + + +@array_api_strict_module +class SeZMEnergyFittingNet(SeZMEnergyFittingNetDP): + pass + + +register_dpmodel_mapping( + GLUFittingNetDP, + lambda v: GLUFittingNet.deserialize(v.serialize()), +) + +register_dpmodel_mapping( + SeZMNetworkCollectionDP, + lambda v: SeZMNetworkCollection.deserialize(v.serialize()), +) diff --git a/source/tests/consistent/descriptor/test_dpa4.py b/source/tests/consistent/descriptor/test_dpa4.py index e6f3216bd4..417e03fc43 100644 --- a/source/tests/consistent/descriptor/test_dpa4.py +++ b/source/tests/consistent/descriptor/test_dpa4.py @@ -19,6 +19,7 @@ ) from ..common import ( + INSTALLED_ARRAY_API_STRICT, INSTALLED_PT, INSTALLED_PT_EXPT, CommonTest, @@ -28,14 +29,40 @@ DescriptorTest, ) -if INSTALLED_PT: - from deepmd.pt.model.descriptor.sezm import DescrptSeZM as DescrptDPA4PT -else: - DescrptDPA4PT = None -if INSTALLED_PT_EXPT: - from deepmd.pt_expt.descriptor.dpa4 import DescrptDPA4 as DescrptDPA4PTExpt +DescrptDPA4PT = None +DescrptDPA4PTExpt = None + + +def _get_descrpt_dpa4_pt() -> type | None: + global DescrptDPA4PT + if not INSTALLED_PT: + return None + if DescrptDPA4PT is None: + from deepmd.pt.model.descriptor.sezm import ( + DescrptSeZM, + ) + + DescrptDPA4PT = DescrptSeZM + return DescrptDPA4PT + + +def _get_descrpt_dpa4_pt_expt() -> type | None: + global DescrptDPA4PTExpt + if not INSTALLED_PT_EXPT: + return None + if DescrptDPA4PTExpt is None: + from deepmd.pt_expt.descriptor.dpa4 import ( + DescrptDPA4, + ) + + DescrptDPA4PTExpt = DescrptDPA4 + return DescrptDPA4PTExpt + + +if INSTALLED_ARRAY_API_STRICT: + from ...array_api_strict.descriptor.dpa4 import DescrptDPA4 as DescrptDPA4Strict else: - DescrptDPA4PTExpt = None + DescrptDPA4Strict = None # not implemented DescrptDPA4TF = None @@ -146,22 +173,31 @@ def data(self) -> dict: @property def skip_pt(self) -> bool: - return CommonTest.skip_pt + return CommonTest.skip_pt or _get_descrpt_dpa4_pt() is None + + @property + def skip_pt_expt(self) -> bool: + return not INSTALLED_PT_EXPT or _get_descrpt_dpa4_pt_expt() is None + + @property + def pt_class(self) -> type | None: + return _get_descrpt_dpa4_pt() + + @property + def pt_expt_class(self) -> type | None: + return _get_descrpt_dpa4_pt_expt() skip_dp = False skip_tf = True skip_jax = True skip_pd = True - skip_pt_expt = not INSTALLED_PT_EXPT - skip_array_api_strict = True + skip_array_api_strict = not INSTALLED_ARRAY_API_STRICT tf_class = DescrptDPA4TF dp_class = DescrptDPA4DP - pt_class = DescrptDPA4PT - pt_expt_class = DescrptDPA4PTExpt jax_class = None pd_class = None - array_api_strict_class = None + array_api_strict_class = DescrptDPA4Strict args: ClassVar[list] = [ *descrpt_se_zm_args(), Argument("ntypes", int, optional=False), @@ -234,6 +270,16 @@ def eval_pt_expt(self, pt_expt_obj: Any) -> Any: mixed_types=True, ) + def eval_array_api_strict(self, array_api_strict_obj: Any) -> Any: + return self.eval_array_api_strict_descriptor( + array_api_strict_obj, + self.natoms, + self.coords, + self.atype, + self.box, + mixed_types=True, + ) + def extract_ret(self, ret: Any, backend) -> tuple[np.ndarray, ...]: return (ret[0],) diff --git a/source/tests/consistent/fitting/test_dpa4_ener.py b/source/tests/consistent/fitting/test_dpa4_ener.py index ee10b54343..b1f68aca71 100644 --- a/source/tests/consistent/fitting/test_dpa4_ener.py +++ b/source/tests/consistent/fitting/test_dpa4_ener.py @@ -6,6 +6,9 @@ import numpy as np +from deepmd.dpmodel.common import ( + to_numpy_array, +) from deepmd.dpmodel.fitting.dpa4_ener import SeZMEnergyFittingNet as SeZMEnerFittingDP from deepmd.env import ( GLOBAL_NP_FLOAT_PRECISION, @@ -15,6 +18,7 @@ ) from ..common import ( + INSTALLED_ARRAY_API_STRICT, INSTALLED_PT, INSTALLED_PT_EXPT, CommonTest, @@ -24,22 +28,61 @@ FittingTest, ) -if INSTALLED_PT: - import torch +torch = None +PT_DEVICE = None +PT_EXPT_DEVICE = None +SeZMEnerFittingPT = None +SeZMEnerFittingPTExpt = None - from deepmd.pt.model.task.sezm_ener import SeZMEnergyFittingNet as SeZMEnerFittingPT - from deepmd.pt.utils.env import DEVICE as PT_DEVICE -else: - SeZMEnerFittingPT = None -if INSTALLED_PT_EXPT: - import torch - from deepmd.pt_expt.fitting.dpa4_ener import ( - SeZMEnergyFittingNet as SeZMEnerFittingPTExpt, +def _get_sezm_ener_fitting_pt() -> type | None: + global PT_DEVICE, SeZMEnerFittingPT, torch + if not INSTALLED_PT: + return None + if SeZMEnerFittingPT is None: + import torch as torch_module + + from deepmd.pt.model.task.sezm_ener import ( + SeZMEnergyFittingNet, + ) + from deepmd.pt.utils.env import ( + DEVICE, + ) + + torch = torch_module + PT_DEVICE = DEVICE + SeZMEnerFittingPT = SeZMEnergyFittingNet + return SeZMEnerFittingPT + + +def _get_sezm_ener_fitting_pt_expt() -> type | None: + global PT_EXPT_DEVICE, SeZMEnerFittingPTExpt, torch + if not INSTALLED_PT_EXPT: + return None + if SeZMEnerFittingPTExpt is None: + import torch as torch_module + + from deepmd.pt_expt.fitting.dpa4_ener import ( + SeZMEnergyFittingNet, + ) + from deepmd.pt_expt.utils.env import ( + DEVICE, + ) + + torch = torch_module + PT_EXPT_DEVICE = DEVICE + SeZMEnerFittingPTExpt = SeZMEnergyFittingNet + return SeZMEnerFittingPTExpt + + +if INSTALLED_ARRAY_API_STRICT: + import array_api_strict + + from ...array_api_strict.fitting.dpa4_ener import ( + SeZMEnergyFittingNet as SeZMEnerFittingStrict, ) - from deepmd.pt_expt.utils.env import DEVICE as PT_EXPT_DEVICE else: - SeZMEnerFittingPTExpt = None + SeZMEnerFittingStrict = None # not implemented SeZMEnerFittingTF = None @@ -70,22 +113,31 @@ def data(self) -> dict: @property def skip_pt(self) -> bool: - return CommonTest.skip_pt + return CommonTest.skip_pt or _get_sezm_ener_fitting_pt() is None + + @property + def skip_pt_expt(self) -> bool: + return not INSTALLED_PT_EXPT or _get_sezm_ener_fitting_pt_expt() is None + + @property + def pt_class(self) -> type | None: + return _get_sezm_ener_fitting_pt() + + @property + def pt_expt_class(self) -> type | None: + return _get_sezm_ener_fitting_pt_expt() skip_dp = False skip_tf = True skip_jax = True skip_pd = True - skip_pt_expt = not INSTALLED_PT_EXPT - skip_array_api_strict = True + skip_array_api_strict = not INSTALLED_ARRAY_API_STRICT tf_class = SeZMEnerFittingTF dp_class = SeZMEnerFittingDP - pt_class = SeZMEnerFittingPT - pt_expt_class = SeZMEnerFittingPTExpt jax_class = None pd_class = None - array_api_strict_class = None + array_api_strict_class = SeZMEnerFittingStrict args = fitting_sezm_ener() def setUp(self) -> None: @@ -138,6 +190,14 @@ def eval_pt_expt(self, pt_expt_obj: Any) -> Any: .numpy() ) + def eval_array_api_strict(self, array_api_strict_obj: Any) -> Any: + return to_numpy_array( + array_api_strict_obj( + array_api_strict.asarray(self.inputs), + array_api_strict.asarray(self.atype.reshape(1, -1)), + )["energy"] + ) + def extract_ret(self, ret: Any, backend) -> tuple[np.ndarray, ...]: return (ret,) From a5b9a6da8f767251c1c925277d37b6c507d85cce Mon Sep 17 00:00:00 2001 From: Jinzhe Zeng Date: Wed, 8 Jul 2026 19:55:54 +0800 Subject: [PATCH 2/3] test(dpa4): address array api strict review comments --- .../tests/array_api_strict/descriptor/dpa4.py | 9 +++++++-- .../array_api_strict/fitting/dpa4_ener.py | 18 +++++------------- 2 files changed, 12 insertions(+), 15 deletions(-) diff --git a/source/tests/array_api_strict/descriptor/dpa4.py b/source/tests/array_api_strict/descriptor/dpa4.py index 3438f85a08..02630d6f6a 100644 --- a/source/tests/array_api_strict/descriptor/dpa4.py +++ b/source/tests/array_api_strict/descriptor/dpa4.py @@ -1,4 +1,8 @@ # SPDX-License-Identifier: LGPL-3.0-or-later +from importlib import ( + import_module, +) + from deepmd.dpmodel.descriptor.dpa4 import DescrptDPA4 as DescrptDPA4DP from deepmd.dpmodel.descriptor.dpa4_nn.activation import ( SwiGLU, @@ -19,12 +23,13 @@ array_api_strict_module, register_dpmodel_mapping, ) -from ..utils import exclude_mask as _strict_exclude_mask # noqa: F401 -from ..utils import network as _strict_network # noqa: F401 from .base_descriptor import ( BaseDescriptor, ) +import_module("..utils.exclude_mask", __package__) +import_module("..utils.network", __package__) + @BaseDescriptor.register("SeZM") @BaseDescriptor.register("sezm") diff --git a/source/tests/array_api_strict/fitting/dpa4_ener.py b/source/tests/array_api_strict/fitting/dpa4_ener.py index 1570ec7891..d7211fc4e5 100644 --- a/source/tests/array_api_strict/fitting/dpa4_ener.py +++ b/source/tests/array_api_strict/fitting/dpa4_ener.py @@ -1,4 +1,7 @@ # SPDX-License-Identifier: LGPL-3.0-or-later +from importlib import ( + import_module, +) from typing import ( ClassVar, ) @@ -13,9 +16,9 @@ from ..common import ( array_api_strict_module, - register_dpmodel_mapping, ) -from ..utils import network as _strict_network # noqa: F401 + +import_module("..utils.network", __package__) @array_api_strict_module @@ -33,14 +36,3 @@ class SeZMNetworkCollection(SeZMNetworkCollectionDP): @array_api_strict_module class SeZMEnergyFittingNet(SeZMEnergyFittingNetDP): pass - - -register_dpmodel_mapping( - GLUFittingNetDP, - lambda v: GLUFittingNet.deserialize(v.serialize()), -) - -register_dpmodel_mapping( - SeZMNetworkCollectionDP, - lambda v: SeZMNetworkCollection.deserialize(v.serialize()), -) From 55a1c152bec75f85205f04b6329baa4cc13e5609 Mon Sep 17 00:00:00 2001 From: njzjz-bot Date: Sat, 11 Jul 2026 13:56:36 +0800 Subject: [PATCH 3/3] test(dpa4): simplify backend imports Restore the established eager conditional-import pattern for the PyTorch backends while retaining the array-api-strict coverage. Coding-Agent: Codex Codex-Version: codex-cli 0.144.1 Model: gpt-5.6-sol Reasoning-Effort: xhigh --- .../tests/consistent/descriptor/test_dpa4.py | 55 +++----------- .../consistent/fitting/test_dpa4_ener.py | 76 +++++-------------- 2 files changed, 30 insertions(+), 101 deletions(-) diff --git a/source/tests/consistent/descriptor/test_dpa4.py b/source/tests/consistent/descriptor/test_dpa4.py index 417e03fc43..1cfbf2f9a1 100644 --- a/source/tests/consistent/descriptor/test_dpa4.py +++ b/source/tests/consistent/descriptor/test_dpa4.py @@ -29,36 +29,14 @@ DescriptorTest, ) -DescrptDPA4PT = None -DescrptDPA4PTExpt = None - - -def _get_descrpt_dpa4_pt() -> type | None: - global DescrptDPA4PT - if not INSTALLED_PT: - return None - if DescrptDPA4PT is None: - from deepmd.pt.model.descriptor.sezm import ( - DescrptSeZM, - ) - - DescrptDPA4PT = DescrptSeZM - return DescrptDPA4PT - - -def _get_descrpt_dpa4_pt_expt() -> type | None: - global DescrptDPA4PTExpt - if not INSTALLED_PT_EXPT: - return None - if DescrptDPA4PTExpt is None: - from deepmd.pt_expt.descriptor.dpa4 import ( - DescrptDPA4, - ) - - DescrptDPA4PTExpt = DescrptDPA4 - return DescrptDPA4PTExpt - - +if INSTALLED_PT: + from deepmd.pt.model.descriptor.sezm import DescrptSeZM as DescrptDPA4PT +else: + DescrptDPA4PT = None +if INSTALLED_PT_EXPT: + from deepmd.pt_expt.descriptor.dpa4 import DescrptDPA4 as DescrptDPA4PTExpt +else: + DescrptDPA4PTExpt = None if INSTALLED_ARRAY_API_STRICT: from ...array_api_strict.descriptor.dpa4 import DescrptDPA4 as DescrptDPA4Strict else: @@ -173,28 +151,19 @@ def data(self) -> dict: @property def skip_pt(self) -> bool: - return CommonTest.skip_pt or _get_descrpt_dpa4_pt() is None - - @property - def skip_pt_expt(self) -> bool: - return not INSTALLED_PT_EXPT or _get_descrpt_dpa4_pt_expt() is None - - @property - def pt_class(self) -> type | None: - return _get_descrpt_dpa4_pt() - - @property - def pt_expt_class(self) -> type | None: - return _get_descrpt_dpa4_pt_expt() + return CommonTest.skip_pt skip_dp = False skip_tf = True skip_jax = True skip_pd = True + skip_pt_expt = not INSTALLED_PT_EXPT skip_array_api_strict = not INSTALLED_ARRAY_API_STRICT tf_class = DescrptDPA4TF dp_class = DescrptDPA4DP + pt_class = DescrptDPA4PT + pt_expt_class = DescrptDPA4PTExpt jax_class = None pd_class = None array_api_strict_class = DescrptDPA4Strict diff --git a/source/tests/consistent/fitting/test_dpa4_ener.py b/source/tests/consistent/fitting/test_dpa4_ener.py index b1f68aca71..3d007a9959 100644 --- a/source/tests/consistent/fitting/test_dpa4_ener.py +++ b/source/tests/consistent/fitting/test_dpa4_ener.py @@ -28,53 +28,22 @@ FittingTest, ) -torch = None -PT_DEVICE = None -PT_EXPT_DEVICE = None -SeZMEnerFittingPT = None -SeZMEnerFittingPTExpt = None - - -def _get_sezm_ener_fitting_pt() -> type | None: - global PT_DEVICE, SeZMEnerFittingPT, torch - if not INSTALLED_PT: - return None - if SeZMEnerFittingPT is None: - import torch as torch_module - - from deepmd.pt.model.task.sezm_ener import ( - SeZMEnergyFittingNet, - ) - from deepmd.pt.utils.env import ( - DEVICE, - ) - - torch = torch_module - PT_DEVICE = DEVICE - SeZMEnerFittingPT = SeZMEnergyFittingNet - return SeZMEnerFittingPT - - -def _get_sezm_ener_fitting_pt_expt() -> type | None: - global PT_EXPT_DEVICE, SeZMEnerFittingPTExpt, torch - if not INSTALLED_PT_EXPT: - return None - if SeZMEnerFittingPTExpt is None: - import torch as torch_module - - from deepmd.pt_expt.fitting.dpa4_ener import ( - SeZMEnergyFittingNet, - ) - from deepmd.pt_expt.utils.env import ( - DEVICE, - ) - - torch = torch_module - PT_EXPT_DEVICE = DEVICE - SeZMEnerFittingPTExpt = SeZMEnergyFittingNet - return SeZMEnerFittingPTExpt +if INSTALLED_PT: + import torch + from deepmd.pt.model.task.sezm_ener import SeZMEnergyFittingNet as SeZMEnerFittingPT + from deepmd.pt.utils.env import DEVICE as PT_DEVICE +else: + SeZMEnerFittingPT = None +if INSTALLED_PT_EXPT: + import torch + from deepmd.pt_expt.fitting.dpa4_ener import ( + SeZMEnergyFittingNet as SeZMEnerFittingPTExpt, + ) + from deepmd.pt_expt.utils.env import DEVICE as PT_EXPT_DEVICE +else: + SeZMEnerFittingPTExpt = None if INSTALLED_ARRAY_API_STRICT: import array_api_strict @@ -113,28 +82,19 @@ def data(self) -> dict: @property def skip_pt(self) -> bool: - return CommonTest.skip_pt or _get_sezm_ener_fitting_pt() is None - - @property - def skip_pt_expt(self) -> bool: - return not INSTALLED_PT_EXPT or _get_sezm_ener_fitting_pt_expt() is None - - @property - def pt_class(self) -> type | None: - return _get_sezm_ener_fitting_pt() - - @property - def pt_expt_class(self) -> type | None: - return _get_sezm_ener_fitting_pt_expt() + return CommonTest.skip_pt skip_dp = False skip_tf = True skip_jax = True skip_pd = True + skip_pt_expt = not INSTALLED_PT_EXPT skip_array_api_strict = not INSTALLED_ARRAY_API_STRICT tf_class = SeZMEnerFittingTF dp_class = SeZMEnerFittingDP + pt_class = SeZMEnerFittingPT + pt_expt_class = SeZMEnerFittingPTExpt jax_class = None pd_class = None array_api_strict_class = SeZMEnerFittingStrict