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 a322d50809a Streaming drain
a322d50809a is described below

commit a322d50809ad6a67980dda7c2c51b4e05b774556
Author: Vitaly Terentyev <[email protected]>
AuthorDate: Tue Oct 6 12:41:03 2026 +0400

    Streaming drain
---
 .../inference/pytorch_image_object_detection.py    | 179 ++++++++++++++++++---
 1 file changed, 159 insertions(+), 20 deletions(-)

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 6e49938c14a..8de7cef2d16 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,10 +45,14 @@ 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
@@ -148,6 +152,16 @@ 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__(
@@ -285,6 +299,14 @@ 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)
 
@@ -382,6 +404,8 @@ 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)
 
   _ = (
@@ -393,6 +417,56 @@ 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 ============
 
 
@@ -401,21 +475,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)
 
-    # 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)
+  # 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))
   pipeline_options.view_as(SetupOptions).save_main_session = save_main_session
   pipeline_options.view_as(StandardOptions).streaming = (
       known_args.mode == 'streaming')
@@ -435,7 +509,8 @@ def run(
       inference_fn=_torchvision_detection_inference_fn,
   )
 
-  pipeline = test_pipeline or beam.Pipeline(options=pipeline_options)
+  pipeline = test_pipeline if test_pipeline is not None else beam.Pipeline(
+    options=pipeline_options)
 
   if known_args.mode == 'batch':
     pcoll = (
@@ -478,6 +553,9 @@ 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)
@@ -494,15 +572,78 @@ def run(
             create_disposition=beam.io.BigQueryDisposition.CREATE_IF_NEEDED,
             method=method))
 
-  result = pipeline.run()
+  feeder_done = threading.Event()
+  stop_feeder = threading.Event()
+  feeder_status = {'result': None, 'error': None}
+  feeder_thread = None
+  result = None
+
   try:
-    result.wait_until_finish(duration=9000000)  # 150 min
+    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:
-    try:
-      result.cancel()
-      result.wait_until_finish(duration=600000)  # up to 10 min to settle 
cancel
-    except Exception:
-      logging.debug("Failed to cancel pipeline result.", exc_info=True)
+    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(
@@ -510,8 +651,6 @@ def run(
           topic_path=known_args.pubsub_topic,
           subscription_path=known_args.pubsub_subscription)
 
-  return result
-
 
 if __name__ == '__main__':
   logging.getLogger().setLevel(logging.INFO)

Reply via email to