This is an automated email from the ASF dual-hosted git repository.
tvalentyn pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/beam.git
The following commit(s) were added to refs/heads/master by this push:
new 5394822d017 Remove usage of a (leaking) shared handles, which can also
accidentally release the last reference hodling a model in
_SharedMap._keepalive when all DoFns have been gc'ed. (#40441)
5394822d017 is described below
commit 5394822d0176d6ae3c5db0085870c6e3da362acf
Author: tvalentyn <[email protected]>
AuthorDate: Wed Oct 7 15:24:03 2026 -0700
Remove usage of a (leaking) shared handles, which can also accidentally
release the last reference hodling a model in _SharedMap._keepalive when all
DoFns have been gc'ed. (#40441)
---
sdks/python/apache_beam/ml/inference/base.py | 7 +++-
sdks/python/apache_beam/ml/inference/base_test.py | 47 +++++++++++++++++++++++
2 files changed, 52 insertions(+), 2 deletions(-)
diff --git a/sdks/python/apache_beam/ml/inference/base.py
b/sdks/python/apache_beam/ml/inference/base.py
index 875b329f42c..8c4b10f8573 100644
--- a/sdks/python/apache_beam/ml/inference/base.py
+++ b/sdks/python/apache_beam/ml/inference/base.py
@@ -1812,7 +1812,7 @@ class _ModelStatus():
self._pending_hard_delete.append((
tag,
datetime.now() + 2 * timedelta(seconds=min_model_life_seconds)))
- self._active_tags.remove(tag)
+ self._active_tags.discard(tag)
def get_valid_tag(self, tag: str) -> str:
"""Takes in a proposed valid tag and returns a valid one.
@@ -1874,13 +1874,16 @@ class _ModelStatus():
return
+_model_statuses: dict[str, _ModelStatus] = {}
+
+
def load_model_status(
model_tag: str, share_across_processes: bool) -> _ModelStatus:
tag = f'{model_tag}_model_status'
if share_across_processes:
return multi_process_shared.MultiProcessShared(
lambda: _ModelStatus(True), tag=tag, always_proxy=True).acquire()
- return shared.Shared().acquire(lambda: _ModelStatus(False), tag=tag)
+ return _model_statuses.setdefault(tag, _ModelStatus(False))
class _ProxyLoader:
diff --git a/sdks/python/apache_beam/ml/inference/base_test.py
b/sdks/python/apache_beam/ml/inference/base_test.py
index 7bca1fc6338..8cf408db76f 100644
--- a/sdks/python/apache_beam/ml/inference/base_test.py
+++ b/sdks/python/apache_beam/ml/inference/base_test.py
@@ -16,6 +16,7 @@
#
"""Tests for apache_beam.ml.base."""
+import gc
import math
import multiprocessing
import os
@@ -26,6 +27,7 @@ import tempfile
import time
import unittest
import unittest.mock
+import uuid
from collections.abc import Iterable
from collections.abc import Mapping
from collections.abc import Sequence
@@ -120,6 +122,22 @@ class FakeSlowModelHandler(base.ModelHandler[int, int,
FakeModel]):
return {'min_batch_size': 1, 'max_batch_size': 1}
+_LOAD_COUNTS: dict[str, int] = {}
+
+
+class LoadCountingModelHandler(base.ModelHandler[int, int, FakeModel]):
+ """Counts load_model() calls per key in the module-level _LOAD_COUNTS."""
+ def __init__(self, key: str):
+ self._key = key
+
+ def load_model(self):
+ _LOAD_COUNTS[self._key] = _LOAD_COUNTS.get(self._key, 0) + 1
+ return FakeModel()
+
+ def run_inference(self, batch, model, inference_args=None):
+ return [model.predict(x) for x in batch]
+
+
class FakeModelHandler(base.ModelHandler[int, int, FakeModel]):
def __init__(
self,
@@ -1837,6 +1855,35 @@ class RunInferenceBaseTest(unittest.TestCase):
self.assertTrue(ms.is_valid_tag('tag1_reload_2'))
self.assertTrue(ms.is_valid_tag('tag2_reload_2'))
+ def test_load_model_status_returns_same_status_per_tag(self):
+ tag = 'tag1' + uuid.uuid4().hex
+ status = base.load_model_status(tag, False)
+ self.assertIs(status, base.load_model_status(tag, False))
+ self.assertIsNot(status, base.load_model_status(tag + '_other', False))
+
+ self.assertTrue(status.is_valid_tag('model'))
+ base.load_model_status(tag, False).try_mark_current_model_invalid(0)
+ self.assertFalse(base.load_model_status(tag, False).is_valid_tag('model'))
+
+ def test_new_dofn_instances_reuse_previously_loaded_models(self):
+ # Runners may recreate DoFn instances, e.g. when the SDK
+ # harness evicts an idle bundle processor.
+ model_tag = 'tag2' + uuid.uuid4().hex
+ handler = LoadCountingModelHandler(model_tag)
+ dofn = base._RunInferenceDoFn(
+ handler, clock=FakeClock(), metrics_namespace=None,
model_tag=model_tag)
+ serialized_dofn = pickle.dumps(dofn)
+
+ dofn_instance = pickle.loads(serialized_dofn)
+ dofn_instance.setup()
+ del dofn_instance
+ gc.collect()
+
+ dofn_instance = pickle.loads(serialized_dofn)
+ dofn_instance.setup()
+
+ self.assertEqual(1, _LOAD_COUNTS[model_tag])
+
def test_model_status_provides_valid_garbage_collection(self):
ms = base._ModelStatus(True)