Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
21 commits
Select commit Hold shift + click to select a range
6b4ca59
fix(pt/pd): fix incompatibility between AutoBatchSize and eval hooks
njzjz Jan 29, 2026
404d1ac
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Jan 29, 2026
01de666
apply Copilot's suggestions
njzjz Jan 29, 2026
232651b
Merge branch 'eval_desc_auto_batch_size' of https://github.com/njzjz/…
njzjz Jan 29, 2026
72c4e36
rm retry
njzjz Jan 29, 2026
f785d61
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Jan 29, 2026
fb6fff5
Apply suggestions from code review
njzjz May 22, 2026
86a137a
test(infer): cover OOM retry hook cleanup
njzjz-bot May 23, 2026
3b640c9
Merge pull request #229 from njzjz-bothub/pr-5181-oom-retry-tests
njzjz May 23, 2026
b5f789a
fix(infer): use iterative OOM hook retries
njzjz-bot May 24, 2026
42771a8
Merge pull request #230 from njzjz-bothub/pr-5181-iterative-retry
njzjz May 24, 2026
31e42cb
test(oom): exercise production eval retry paths
njzjz-bot May 25, 2026
d19993b
Merge pull request #231 from njzjz-bothub/pr-5181-production-tests
njzjz May 25, 2026
a80ce64
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] May 25, 2026
a1ba195
Potential fix for pull request finding
njzjz May 25, 2026
bd871c3
test(oom): return floating mock outputs
njzjz-bot May 25, 2026
23340af
Merge pull request #232 from njzjz-bothub/pr-5181-production-tests
njzjz May 25, 2026
80025d0
Merge branch 'master' into eval_desc_auto_batch_size
njzjz May 27, 2026
99f812d
test(oom): move backend retry tests out of common
njzjz-bot May 28, 2026
3e79ef3
test: move backend oom retry tests to backend roots
njzjz-bot May 28, 2026
0abacb2
Merge branch 'master' into eval_desc_auto_batch_size
njzjz May 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
77 changes: 51 additions & 26 deletions deepmd/pd/infer/deep_eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,9 @@
to_numpy_array,
to_paddle_tensor,
)
from deepmd.utils.batch_size import (
RetrySignal,
)
from deepmd.utils.econf_embd import (
sort_element_type,
)
Expand Down Expand Up @@ -830,19 +833,30 @@ def eval_descriptor(
model = (
self.dp.model["Default"] if isinstance(self.dp, ModelWrapper) else self.dp
)
model.set_eval_descriptor_hook(True)
self.eval(
coords,
cells,
atom_types,
atomic=False,
fparam=fparam,
aparam=aparam,
**kwargs,
)
descriptor = model.eval_descriptor()
model.set_eval_descriptor_hook(False)
return to_numpy_array(descriptor)
while True:
if self.auto_batch_size is not None:
self.auto_batch_size.set_oom_retry_mode(True)
model.set_eval_descriptor_hook(True)
retry = False
try:
self.eval(
coords,
cells,
atom_types,
atomic=False,
fparam=fparam,
aparam=aparam,
**kwargs,
)
descriptor = model.eval_descriptor()
except RetrySignal:
retry = True
finally:
model.set_eval_descriptor_hook(False)
if self.auto_batch_size is not None:
self.auto_batch_size.set_oom_retry_mode(False)
if not retry:
return to_numpy_array(descriptor)

def eval_fitting_last_layer(
self,
Expand Down Expand Up @@ -885,16 +899,27 @@ def eval_fitting_last_layer(
Fitting output before last layer.
"""
model = self.dp.model["Default"]
model.set_eval_fitting_last_layer_hook(True)
self.eval(
coords,
cells,
atom_types,
atomic=False,
fparam=fparam,
aparam=aparam,
**kwargs,
)
fitting_net = model.eval_fitting_last_layer()
model.set_eval_fitting_last_layer_hook(False)
return to_numpy_array(fitting_net)
while True:
if self.auto_batch_size is not None:
self.auto_batch_size.set_oom_retry_mode(True)
model.set_eval_fitting_last_layer_hook(True)
retry = False
try:
self.eval(
coords,
cells,
atom_types,
atomic=False,
fparam=fparam,
aparam=aparam,
**kwargs,
)
fitting_net = model.eval_fitting_last_layer()
except RetrySignal:
retry = True
finally:
model.set_eval_fitting_last_layer_hook(False)
if self.auto_batch_size is not None:
self.auto_batch_size.set_oom_retry_mode(False)
if not retry:
return to_numpy_array(fitting_net)
77 changes: 51 additions & 26 deletions deepmd/pt/infer/deep_eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,9 @@
to_numpy_array,
to_torch_tensor,
)
from deepmd.utils.batch_size import (
RetrySignal,
)
from deepmd.utils.econf_embd import (
sort_element_type,
)
Expand Down Expand Up @@ -847,19 +850,30 @@ def eval_descriptor(
Descriptors.
"""
model = self.dp.model["Default"]
model.set_eval_descriptor_hook(True)
self.eval(
coords,
cells,
atom_types,
atomic=False,
fparam=fparam,
aparam=aparam,
**kwargs,
)
descriptor = model.eval_descriptor()
model.set_eval_descriptor_hook(False)
return to_numpy_array(descriptor)
while True:
if self.auto_batch_size is not None:
self.auto_batch_size.set_oom_retry_mode(True)
model.set_eval_descriptor_hook(True)
retry = False
try:
self.eval(
coords,
cells,
atom_types,
atomic=False,
fparam=fparam,
aparam=aparam,
**kwargs,
)
descriptor = model.eval_descriptor()
except RetrySignal:
retry = True
finally:
model.set_eval_descriptor_hook(False)
if self.auto_batch_size is not None:
self.auto_batch_size.set_oom_retry_mode(False)
if not retry:
return to_numpy_array(descriptor)

def eval_fitting_last_layer(
self,
Expand Down Expand Up @@ -902,16 +916,27 @@ def eval_fitting_last_layer(
Fitting output before last layer.
"""
model = self.dp.model["Default"]
model.set_eval_fitting_last_layer_hook(True)
self.eval(
coords,
cells,
atom_types,
atomic=False,
fparam=fparam,
aparam=aparam,
**kwargs,
)
fitting_net = model.eval_fitting_last_layer()
model.set_eval_fitting_last_layer_hook(False)
return to_numpy_array(fitting_net)
while True:
if self.auto_batch_size is not None:
self.auto_batch_size.set_oom_retry_mode(True)
model.set_eval_fitting_last_layer_hook(True)
retry = False
try:
self.eval(
coords,
cells,
atom_types,
atomic=False,
fparam=fparam,
aparam=aparam,
**kwargs,
)
fitting_net = model.eval_fitting_last_layer()
except RetrySignal:
retry = True
finally:
model.set_eval_fitting_last_layer_hook(False)
if self.auto_batch_size is not None:
self.auto_batch_size.set_oom_retry_mode(False)
if not retry:
return to_numpy_array(fitting_net)
24 changes: 24 additions & 0 deletions deepmd/utils/batch_size.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,10 @@
log = logging.getLogger(__name__)


class RetrySignal(Exception):
"""Signal to retry execution after OOM error."""


class AutoBatchSize(ABC):
"""This class allows DeePMD-kit to automatically decide the maximum
batch size that will not cause an OOM error.
Expand Down Expand Up @@ -85,6 +89,7 @@ def __init__(
)

self.factor = factor
self.oom_retry_mode = False

def execute(
self, callable: Callable, start_index: int, natoms: int
Expand Down Expand Up @@ -135,6 +140,8 @@ def execute(
) from e
# adjust the next batch size
self._adjust_batch_size(1.0 / self.factor)
if self.oom_retry_mode:
raise RetrySignal from e
return 0, None
else:
n_tot = n_batch * natoms
Expand Down Expand Up @@ -292,3 +299,20 @@ def is_oom_error(self, e: Exception) -> bool:
bool
True if the exception is an OOM error
"""

def set_oom_retry_mode(self, enable: bool) -> None:
"""Set OOM retry mode.

In OOM retry mode, an OOM during execution may reduce the current
batch size and raise :class:`RetrySignal` to indicate that execution
should be retried.

Callers that want all data to be re-executed must catch
:class:`RetrySignal` and restart the full evaluation themselves.

Parameters
----------
enable : bool
True to enable OOM retry mode
"""
self.oom_retry_mode = enable
50 changes: 50 additions & 0 deletions source/tests/common/test_oom_retry.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
# SPDX-License-Identifier: LGPL-3.0-or-later
import unittest

from deepmd.utils.batch_size import (
AutoBatchSize,
RetrySignal,
)
from deepmd.utils.errors import (
OutOfMemoryError,
)


class CustomizedAutoBatchSizeGPU(AutoBatchSize):
def is_gpu_available(self) -> bool:
return True

def is_oom_error(self, e):
return isinstance(e, OutOfMemoryError)


class TestOOMRetry(unittest.TestCase):
def test_execute_oom_retry_mode_raises_retry_signal(self) -> None:
auto_batch_size = CustomizedAutoBatchSizeGPU(256, 2.0)

oom = OutOfMemoryError("oom")

def executor(batch_size: int, start_index: int) -> tuple[int, None]:
raise oom

auto_batch_size.set_oom_retry_mode(True)
with self.assertRaises(RetrySignal) as context:
auto_batch_size.execute(executor, 0, 1)
self.assertIs(context.exception.__cause__, oom)
self.assertEqual(auto_batch_size.current_batch_size, 128)

def test_execute_oom_retry_mode_false_returns_zero(self) -> None:
auto_batch_size = CustomizedAutoBatchSizeGPU(256, 2.0)

def executor(batch_size: int, start_index: int) -> tuple[int, None]:
raise OutOfMemoryError("oom")

auto_batch_size.set_oom_retry_mode(False)
n_batch, result = auto_batch_size.execute(executor, 0, 1)
self.assertEqual(n_batch, 0)
self.assertIsNone(result)
self.assertEqual(auto_batch_size.current_batch_size, 128)


if __name__ == "__main__":
unittest.main()
Loading
Loading