Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 40 additions & 0 deletions python/cudf_polars/cudf_polars/streaming/benchmarks/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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)
Expand Down Expand Up @@ -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,
)


Expand All @@ -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,
Expand All @@ -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,
)


Expand All @@ -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,
)
Expand Down Expand Up @@ -1437,19 +1460,23 @@ 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(
engine=engine,
run_id=run_config.run_id,
collect_traces=run_config.collect_traces,
)
shutdown_duration_ms = _elapsed_ms(shutdown_time_begin)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
_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,
)


Expand Down Expand Up @@ -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,
Expand All @@ -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,
)


Expand All @@ -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)
Expand Down Expand Up @@ -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
Expand All @@ -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,
Expand All @@ -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,
)


Expand Down
Loading