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)
 

Reply via email to