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)