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
1 change: 1 addition & 0 deletions CHANGES.md
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,7 @@
* (Go) Fixed a data race on the Prism runner's artifact cache map in JobServices ([#32656](https://github.com/apache/beam/issues/32656)).
* (Go) Fixed the harness leaking Data/State gRPC streams after the worker stops, and a deadlock when Send returns EOF ([#40260](https://github.com/apache/beam/issues/40260)).
* (Go) Fixed a deadlock recreating a data channel while holding the channel lock ([#40414](https://github.com/apache/beam/issues/40414)).
* (Python) Fixed `KeyedModelHandler` not enforcing `max_models_per_worker_hint` when several copies of the model handler run in the same SDK process ([#40468](https://github.com/apache/beam/issues/40468)).
* (Java) Fixed the declared schema of the error output of the Kafka write SchemaTransform, which wrapped the error schema a second time and did not match the rows it emits ([#39760](https://github.com/apache/beam/issues/39760)).
* (Go) Fixed pubsubio importing a `google.golang.org/genproto` package removed in recent releases, which broke builds of Go modules depending on a current `genproto` version ([#40018](https://github.com/apache/beam/issues/40018)).
* (Java) BigQueryIO now treats a 404 when deleting a temporary table or dataset as success, so a replayed work item whose earlier attempt already deleted it no longer retries forever ([#24997](https://github.com/apache/beam/issues/24997)).
Expand Down
21 changes: 16 additions & 5 deletions sdks/python/apache_beam/ml/inference/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -566,6 +566,8 @@ def __init__(self, mh_map: dict[str, ModelHandler]):
model.
"""
self._max_models = None
# Process ids that have already contributed to _max_models
self._incremented_process_ids: set[int] = set()
# Map keys to model handlers
self._mh_map: dict[str, ModelHandler] = mh_map
# Map keys to the last updated model path for that key
Expand Down Expand Up @@ -620,14 +622,22 @@ def load(self, key: str) -> _ModelLoadStats:
return _ModelLoadStats(
tag, end_time - start_time, memory_after - memory_before)

def increment_max_models(self, increment: int):
def increment_max_models(
self, increment: int, process_id: Optional[int] = None):
"""
Increments the number of models that this instance of a
_ModelHandlerManager is able to hold. If it is never called,
no limit is imposed.
Args:
increment: the amount by which we are incrementing the number of models.
process_id: the id of the process requesting the increment. If set,
only the first increment from each process is applied, since many
copies of the model handler can share this manager within a process.
"""
if process_id is not None:
if process_id in self._incremented_process_ids:
return
self._incremented_process_ids.add(process_id)
if self._max_models is None:
self._max_models = 0
self._max_models += increment
Expand Down Expand Up @@ -848,11 +858,12 @@ def run_inference(
self._unkeyed.run_inference(unkeyed_batch, model, inference_args))

# The first time a MultiProcessShared ModelManager is used for inference
# from this process, we should increment its max model count
# from this process, we should increment its max model count. The manager
# may live in another process and be shared by several deserialized copies
# of this handler, so it dedupes increments by process id.
if self._max_models_per_worker_hint is not None:
lock = threading.Lock()
if lock.acquire(blocking=False):
model.increment_max_models(self._max_models_per_worker_hint)
model.increment_max_models(
self._max_models_per_worker_hint, process_id=os.getpid())
self._max_models_per_worker_hint = None

batch_by_key = defaultdict(list)
Expand Down
35 changes: 35 additions & 0 deletions sdks/python/apache_beam/ml/inference/base_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -1758,6 +1758,41 @@ def test_model_handler_manager_evicts_models_after_being_incremented(self):
mh3.load_model, tag=tag3).acquire()
self.assertEqual(8, model3.predict(10))

def test_model_handler_manager_increments_once_per_process(self):
mhs = {
'key1': FakeModelHandler(state=1),
'key2': FakeModelHandler(state=2),
'key3': FakeModelHandler(state=3)
}
mm = base._ModelHandlerManager(mh_map=mhs)
mm.increment_max_models(1, process_id=1)
mm.increment_max_models(1, process_id=1)
mm.load('key1')
mm.load('key2')
mm.load('key3')
self.assertEqual(['key3'], list(mm._tag_map.keys()))

mm.increment_max_models(1, process_id=2)
mm.load('key1')
self.assertEqual(['key3', 'key1'], list(mm._tag_map.keys()))

def test_keyed_model_handler_max_models_hint_with_handler_copies(self):
mhs = [
base.KeyModelMapping([k],
FakeModelHandler(
state=i, multi_process_shared=True))
for i, k in enumerate(['a', 'b', 'c'])
]
keyed_mh = base.KeyedModelHandler(mhs, max_models_per_worker_hint=1)
mm = keyed_mh.load_model()
# Each DoFn instance in a process deserializes its own copy of the model
# handler, but they all share the same _ModelHandlerManager.
for _ in range(5):
mh_copy = pickle.loads(pickle.dumps(keyed_mh))
mh_copy.override_metrics('test_namespace')
list(mh_copy.run_inference([('a', 1), ('b', 2), ('c', 3)], mm))
self.assertEqual(1, len(mm._tag_map))

def test_run_inference_loads_different_models(self):
mh1 = FakeModelHandler(incrementing=True, min_batch_size=3)
with TestPipeline() as pipeline:
Expand Down
Loading