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__':

Reply via email to