diff --git a/python/cudf_polars/cudf_polars/streaming/benchmarks/utils.py b/python/cudf_polars/cudf_polars/streaming/benchmarks/utils.py index d9437dfce82..89b0c05c277 100644 --- a/python/cudf_polars/cudf_polars/streaming/benchmarks/utils.py +++ b/python/cudf_polars/cudf_polars/streaming/benchmarks/utils.py @@ -1284,12 +1284,19 @@ def _run_query_loop( return records, plans, validation_failures, query_failures +def _elapsed_ms(begin: float) -> float: + """Return milliseconds elapsed since ``begin`` (a ``time.monotonic()`` timestamp).""" + return (time.monotonic() - begin) * 1000 + + def _finalize_benchmark_run( args: argparse.Namespace, run_config: RunConfig, validation_failures: list[int], query_failures: list[tuple[int, int]], serializable_engine_config: dict[str, Any], + startup_duration_ms: float | None, + shutdown_duration_ms: float | None, ) -> None: """Summarize, serialize, and exit after a benchmark run.""" if args.summarize: @@ -1312,6 +1319,14 @@ def _finalize_benchmark_run( ) if not validation_failures and not query_failures: print("✅ All validated queries passed.") + + # We modify the serialized engine config here to record the engine's + # start/end duration. We have to do it here, rather than in `RunConfig.serialize()` + # since serialize() needs to run before the engine is shutdown. + serializable_engine_config = dict(serializable_engine_config) + serializable_engine_config["startup_duration_ms"] = startup_duration_ms + serializable_engine_config["shutdown_duration_ms"] = shutdown_duration_ms + args.output.write(json.dumps(serializable_engine_config)) args.output.write("\n") sys.exit(1 if (query_failures or validation_failures) else 0) @@ -1340,6 +1355,8 @@ def run_polars_cpu( validation_failures, query_failures, serializable_engine_config=run_config.serialize(engine=None), + startup_duration_ms=None, + shutdown_duration_ms=None, ) @@ -1357,10 +1374,12 @@ def run_polars_in_memory( "parquet_options": parquet_options, } engine_options.setdefault("raise_on_fail", True) + start_time_begin = time.monotonic() engine = pl.GPUEngine( executor="in-memory", **engine_options, ) + startup_duration_ms = _elapsed_ms(start_time_begin) records, plans, validation_failures, query_failures = _run_query_loop( benchmark, args, @@ -1377,6 +1396,8 @@ def run_polars_in_memory( validation_failures, query_failures, serializable_engine_config=run_config.serialize(engine=engine), + startup_duration_ms=startup_duration_ms, + shutdown_duration_ms=None, ) @@ -1399,11 +1420,13 @@ def run_polars_spmd( "parquet_options": parquet_options, } engine_options.setdefault("raise_on_fail", True) + start_time_begin = time.monotonic() with SPMDEngine( rapidsmpf_options=run_config.streaming_options.to_rapidsmpf_options(), executor_options=executor_options, engine_options=engine_options, ) as engine: + startup_duration_ms = _elapsed_ms(start_time_begin) from cudf_polars.engine.spmd import ( allgather_polars_dataframe, ) @@ -1437,6 +1460,7 @@ def _allgather_result(df: pl.DataFrame) -> pl.DataFrame: ) # We need to create this before StreamingEngine.shutdown(), which clears engine.config serializable_engine_config = run_config.serialize(engine=engine) + shutdown_time_begin = time.monotonic() if is_rank_0: _write_quent_traces( @@ -1444,12 +1468,15 @@ def _allgather_result(df: pl.DataFrame) -> pl.DataFrame: run_id=run_config.run_id, collect_traces=run_config.collect_traces, ) + shutdown_duration_ms = _elapsed_ms(shutdown_time_begin) _finalize_benchmark_run( args, run_config, validation_failures, query_failures, serializable_engine_config=serializable_engine_config, + startup_duration_ms=startup_duration_ms, + shutdown_duration_ms=shutdown_duration_ms, ) @@ -1478,12 +1505,14 @@ def run_polars_ray( if run_config.num_gpus is not None: ray_init_options["num_gpus"] = run_config.num_gpus + start_time_begin = time.monotonic() with RayEngine( rapidsmpf_options=run_config.streaming_options.to_rapidsmpf_options(), executor_options=executor_options, engine_options=engine_options, ray_init_options=ray_init_options, ) as engine: + startup_duration_ms = _elapsed_ms(start_time_begin) run_config = dataclasses.replace(run_config, n_workers=engine.nranks) records, plans, validation_failures, query_failures = _run_query_loop( benchmark, @@ -1497,18 +1526,22 @@ def run_polars_ray( run_config = _consolidate_logs(run_config, engine=engine) # We need to create this before StreamingEngine.shutdown(), which clears engine.config serializable_engine_config = run_config.serialize(engine=engine) + shutdown_time_begin = time.monotonic() _write_quent_traces( engine=engine, run_id=run_config.run_id, collect_traces=run_config.collect_traces, ) + shutdown_duration_ms = _elapsed_ms(shutdown_time_begin) _finalize_benchmark_run( args, run_config, validation_failures, query_failures, serializable_engine_config=serializable_engine_config, + startup_duration_ms=startup_duration_ms, + shutdown_duration_ms=shutdown_duration_ms, ) @@ -1525,6 +1558,8 @@ def run_polars_dask( from cudf_polars.engine.dask import DaskEngine + start_time_begin = time.monotonic() + executor_options = get_executor_options(run_config, benchmark=benchmark) # "cluster" is reserved — DaskEngine sets it executor_options.pop("cluster", None) @@ -1552,6 +1587,7 @@ def run_polars_dask( engine_options=engine_options, dask_client=dask_client, ) as engine: + startup_duration_ms = _elapsed_ms(start_time_begin) run_config = dataclasses.replace(run_config, n_workers=engine.nranks) records, plans, validation_failures, query_failures = _run_query_loop( benchmark, args, run_config, engine, numeric_type, date_type @@ -1562,6 +1598,7 @@ def run_polars_dask( run_config = _consolidate_logs(run_config, engine) # We need to create this before StreamingEngine.shutdown(), which clears engine.config serializable_engine_config = run_config.serialize(engine=engine) + shutdown_time_begin = time.monotonic() _write_quent_traces( engine=engine, @@ -1571,12 +1608,15 @@ def run_polars_dask( finally: if dask_client is not None: dask_client.close() + shutdown_duration_ms = _elapsed_ms(shutdown_time_begin) _finalize_benchmark_run( args, run_config, validation_failures, query_failures, serializable_engine_config=serializable_engine_config, + startup_duration_ms=startup_duration_ms, + shutdown_duration_ms=shutdown_duration_ms, )