This is an automated email from the ASF dual-hosted git repository.
Amar3tto pushed a commit to branch fix-inference-ml
in repository https://gitbox.apache.org/repos/asf/beam.git
The following commit(s) were added to refs/heads/fix-inference-ml by this push:
new fc0b1086a8c Refactor
fc0b1086a8c is described below
commit fc0b1086a8c85f03480cc123b5d9df1558334e50
Author: Vitaly Terentyev <[email protected]>
AuthorDate: Tue Oct 6 14:34:20 2026 +0400
Refactor
---
.../examples/inference/pytorch_image_captioning.py | 13 --
.../inference/pytorch_image_object_detection.py | 190 +++------------------
2 files changed, 23 insertions(+), 180 deletions(-)
diff --git
a/sdks/python/apache_beam/examples/inference/pytorch_image_captioning.py
b/sdks/python/apache_beam/examples/inference/pytorch_image_captioning.py
index 3f41ee7f4b3..63165492b4c 100644
--- a/sdks/python/apache_beam/examples/inference/pytorch_image_captioning.py
+++ b/sdks/python/apache_beam/examples/inference/pytorch_image_captioning.py
@@ -473,19 +473,6 @@ def cleanup_pubsub_resources(
except NotFound:
logging.info(f"Topic already deleted: {topic_path}")
- try:
- subscriber.delete_subscription(
- request={"subscription": full_subscription_path})
- logging.info(f"Deleted subscription: {subscription_name}")
- except NotFound:
- logging.info(f"Subscription already deleted: {subscription_name}")
-
- try:
- publisher.delete_topic(request={"topic": full_topic_path})
- logging.info(f"Deleted topic: {topic_name}")
- except NotFound:
- logging.info(f"Topic already deleted: {topic_name}")
-
def override_or_add(args, flag, value):
if flag in args:
diff --git
a/sdks/python/apache_beam/examples/inference/pytorch_image_object_detection.py
b/sdks/python/apache_beam/examples/inference/pytorch_image_object_detection.py
index dbabcd19fb5..949fc48180f 100644
---
a/sdks/python/apache_beam/examples/inference/pytorch_image_object_detection.py
+++
b/sdks/python/apache_beam/examples/inference/pytorch_image_object_detection.py
@@ -45,14 +45,10 @@ from apache_beam.ml.inference.base import KeyedModelHandler
from apache_beam.ml.inference.base import PredictionResult
from apache_beam.ml.inference.base import RunInference
from apache_beam.ml.inference.pytorch_inference import
PytorchModelHandlerTensor
-from apache_beam.metrics import Metrics
-from apache_beam.metrics import MetricsFilter
-from apache_beam.runners.dataflow.internal.apiclient import
DataflowApplicationClient
from apache_beam.options.pipeline_options import PipelineOptions
from apache_beam.options.pipeline_options import SetupOptions
from apache_beam.options.pipeline_options import StandardOptions
from apache_beam.runners.runner import PipelineResult
-from apache_beam.runners.runner import PipelineState
from apache_beam.transforms import window
from google.api_core.exceptions import NotFound
@@ -152,16 +148,6 @@ def _torchvision_detection_inference_fn(
return outputs
-class CountProcessedDoFn(beam.DoFn):
- """Records image outputs after inference and before the BigQuery sink."""
- def __init__(self):
- self.counter = Metrics.counter('detection_benchmark', 'processed_images')
-
- def process(self, row):
- self.counter.inc()
- yield row
-
-
class PostProcessDoFn(beam.DoFn):
"""PredictionResult -> dict row for BQ."""
def __init__(
@@ -276,7 +262,7 @@ def parse_known_args(argv):
parser.add_argument(
'--feeder_start_delay_sec',
type=int,
- default=900,
+ default=100,
help=(
'Delay before starting the feeder pipeline that reads URIs from GCS '
'and publishes them to Pub/Sub. This delay allows the main streaming
'
@@ -299,14 +285,6 @@ def parse_known_args(argv):
# Batch sizing (no right-fitting)
parser.add_argument('--inference_batch_size', type=int, default=8)
- # A finite benchmark workload on an unbounded Pub/Sub subscription.
- parser.add_argument(
- '--expected_messages', type=int, default=50000,
- help='Number of nonempty image URIs in --input (required for streaming).')
- parser.add_argument('--max_runtime_sec', type=int, default=9000)
- parser.add_argument('--drain_timeout_sec', type=int, default=600)
- parser.add_argument('--metrics_poll_interval_sec', type=int, default=30)
-
# Preprocess
parser.add_argument('--image_size', type=int, default=800)
@@ -404,8 +382,6 @@ def run_load_pipeline(known_args, pipeline_args):
]
pipeline_options = PipelineOptions(pipeline_args)
- # The feeder reads a bounded GCS file, even when the main job is streaming.
- pipeline_options.view_as(StandardOptions).streaming = False
pipeline = beam.Pipeline(options=pipeline_options)
_ = (
@@ -417,56 +393,6 @@ def run_load_pipeline(known_args, pipeline_args):
return pipeline.run()
-# ============ Completion monitoring ============
-
-
-def processed_image_count(result: PipelineResult) -> int:
- """Query the committed metric; do not count tentative/retried bundles."""
- metrics = result.metrics().query(
- MetricsFilter().with_namespace('detection_benchmark').with_name(
- 'processed_images'))
- counters = metrics.get('counters', [])
- return sum(int(m.committed or 0) for m in counters)
-
-
-def wait_for_terminal_state(result: PipelineResult, timeout_sec: int) -> str:
- deadline = time.monotonic() + timeout_sec
- while time.monotonic() < deadline:
- state = result.state
- if PipelineState.is_terminal(state):
- return state
- time.sleep(10)
- raise TimeoutError(
- f'Job {result.job_id()} did not terminate in {timeout_sec}s; '
- f'current state: {result.state}')
-
-
-def wait_until_processed(
- result: PipelineResult,
- expected: int,
- feeder_done: threading.Event,
- feeder_status: dict,
- timeout_sec: int,
- poll_sec: int,
-) -> None:
- deadline = time.monotonic() + timeout_sec
- while time.monotonic() < deadline:
- state = result.state
- if PipelineState.is_terminal(state):
- raise RuntimeError(
- f'Inference job terminated before reaching {expected} images: {state}')
- if feeder_done.is_set():
- if feeder_status['error'] is not None:
- raise RuntimeError('Pub/Sub feeder failed') from feeder_status['error']
- current = processed_image_count(result)
- logging.info('Processed %d / %d images', current, expected)
- if current >= expected:
- return
- time.sleep(poll_sec)
- raise TimeoutError(
- f'Inference did not process {expected} images within {timeout_sec}s.')
-
-
# ============ Main pipeline ============
@@ -475,21 +401,21 @@ def run(
known_args, pipeline_args = parse_known_args(argv)
if known_args.mode == 'streaming':
- if not known_args.expected_messages or known_args.expected_messages <= 0:
- raise ValueError('--expected_messages > 0 is required in streaming mode')
- if known_args.metrics_poll_interval_sec <= 0:
- raise ValueError('--metrics_poll_interval_sec must be positive')
- if known_args.max_runtime_sec <= 0 or known_args.drain_timeout_sec <= 0:
- raise ValueError('Runtime and drain timeouts must be positive')
ensure_pubsub_resources(
project=known_args.project,
topic_path=known_args.pubsub_topic,
subscription_path=known_args.pubsub_subscription)
- # When a benchmark passes test_pipeline, use its existing options.
- pipeline_options = (
- test_pipeline.get_pipeline_options()
- if test_pipeline is not None else PipelineOptions(pipeline_args))
+ # Start feeder thread that reads URIs from GCS and fills Pub/Sub.
+ # Delay is used to allow the main streaming pipeline workers to start
+ # and autoscale before the feeder pipeline begins publishing messages.
+ threading.Thread(
+ target=lambda: (
+ time.sleep(known_args.feeder_start_delay_sec), run_load_pipeline(
+ known_args, pipeline_args)),
+ daemon=True).start()
+
+ pipeline_options = PipelineOptions(pipeline_args)
pipeline_options.view_as(SetupOptions).save_main_session = save_main_session
pipeline_options.view_as(StandardOptions).streaming = (
known_args.mode == 'streaming')
@@ -510,8 +436,7 @@ def run(
inference_fn=_torchvision_detection_inference_fn,
)
- pipeline = test_pipeline if test_pipeline is not None else beam.Pipeline(
- options=pipeline_options)
+ pipeline = test_pipeline or beam.Pipeline(options=pipeline_options)
if known_args.mode == 'batch':
pcoll = (
@@ -554,9 +479,6 @@ def run(
score_threshold=known_args.score_threshold,
max_detections=known_args.max_detections)))
- if known_args.mode == 'streaming':
- results = results | 'CountProcessed' >> beam.ParDo(CountProcessedDoFn())
-
method = (
beam.io.WriteToBigQuery.Method.FILE_LOADS if known_args.mode == 'batch'
else beam.io.WriteToBigQuery.Method.STREAMING_INSERTS)
@@ -573,84 +495,18 @@ def run(
create_disposition=beam.io.BigQueryDisposition.CREATE_IF_NEEDED,
method=method))
- feeder_done = threading.Event()
- stop_feeder = threading.Event()
- feeder_status = {'result': None, 'error': None}
- feeder_thread = None
- result = None
+ result = pipeline.run()
+ result.wait_until_finish(duration=1800000) # 30 min
+ result.cancel()
+ result.wait_until_finish(duration=600000) # up to 10 min to settle cancel
- try:
- result = pipeline.run()
-
- if known_args.mode == 'streaming':
- # Start the feeder only after the inference job has been submitted.
- def start_feeder():
- try:
- if stop_feeder.wait(known_args.feeder_start_delay_sec):
- return
- feeder_status['result'] = run_load_pipeline(known_args,
pipeline_args)
- feeder_state = feeder_status['result'].wait_until_finish()
- if feeder_state != PipelineState.DONE:
- raise RuntimeError(f'Feeder finished in state: {feeder_state}')
- except Exception as exc: # Report errors to the main thread.
- feeder_status['error'] = exc
- logging.exception('Pub/Sub feeder failed')
- finally:
- feeder_done.set()
-
- feeder_thread = threading.Thread(target=start_feeder, daemon=True)
- feeder_thread.start()
- wait_until_processed(
- result,
- known_args.expected_messages,
- feeder_done,
- feeder_status,
- known_args.max_runtime_sec,
- known_args.metrics_poll_interval_sec)
-
- # Drain (do not cancel): stop consuming Pub/Sub and finish in-flight
- # inference and BigQuery writes before collecting benchmark metrics.
- logging.info('Reached %d processed images: draining Dataflow job %s',
- known_args.expected_messages, result.job_id())
- client = DataflowApplicationClient(pipeline_options)
- if not client.modify_job_state(result.job_id(), 'JOB_STATE_DRAINED'):
- raise RuntimeError(f'Failed to drain Dataflow job {result.job_id()}')
- terminal = wait_for_terminal_state(result, known_args.drain_timeout_sec)
- if terminal != PipelineState.DRAINED:
- raise RuntimeError(f'Expected DRAINED state, got {terminal}')
- else:
- # A bounded batch pipeline must finish successfully by itself.
- state = result.wait_until_finish(duration=known_args.max_runtime_sec *
1000)
- if state != PipelineState.DONE:
- raise TimeoutError(f'Batch pipeline did not complete: {state}')
- return result
- finally:
- stop_feeder.set()
- if feeder_thread is not None and feeder_thread.is_alive():
- feeder_result = feeder_status['result']
- if feeder_result is not None and not PipelineState.is_terminal(
- feeder_result.state):
- try:
- feeder_result.cancel()
- except Exception:
- logging.exception('Failed to cancel unfinished feeder job')
- feeder_thread.join(timeout=30)
-
- # A failure or deadline must not leave a paid streaming job running.
- if result is not None and not PipelineState.is_terminal(result.state):
- logging.warning('Cancelling unfinished Dataflow job %s', result.job_id())
- try:
- result.cancel()
- wait_for_terminal_state(result, known_args.drain_timeout_sec)
- except Exception:
- logging.exception('Failed to stop unfinished Dataflow job %s',
- result.job_id())
-
- if known_args.mode == 'streaming':
- cleanup_pubsub_resources(
- project=known_args.project,
- topic_path=known_args.pubsub_topic,
- subscription_path=known_args.pubsub_subscription)
+ if known_args.mode == 'streaming':
+ cleanup_pubsub_resources(
+ project=known_args.project,
+ topic_path=known_args.pubsub_topic,
+ subscription_path=known_args.pubsub_subscription)
+
+ return result
if __name__ == '__main__':