diff --git a/src/cloudai/workloads/osu_bench/report_generation_strategy.py b/src/cloudai/workloads/osu_bench/report_generation_strategy.py index 0dc1a5aed..fed09e290 100644 --- a/src/cloudai/workloads/osu_bench/report_generation_strategy.py +++ b/src/cloudai/workloads/osu_bench/report_generation_strategy.py @@ -21,9 +21,9 @@ from enum import Enum from functools import cache from pathlib import Path -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, ClassVar -from cloudai.core import ReportGenerationStrategy +from cloudai.core import METRIC_ERROR, MetricValue, ReportGenerationStrategy from cloudai.util.lazy_imports import lazy if TYPE_CHECKING: @@ -140,6 +140,8 @@ def extract_osu_bench_data(stdout_file: Path) -> pd.DataFrame: class OSUBenchReportGenerationStrategy(ReportGenerationStrategy): """Report generation strategy for OSU Bench.""" + metrics: ClassVar[list[str]] = ["default", "avg_lat"] + @property def results_file(self) -> Path: return self.test_run.output_path / "stdout.txt" @@ -148,9 +150,19 @@ def can_handle_directory(self) -> bool: df = extract_osu_bench_data(self.results_file) return not df.empty + def get_metric(self, metric: str) -> MetricValue: + if metric not in self.metrics: + return METRIC_ERROR + + df = extract_osu_bench_data(self.results_file) + if df.empty or "avg_lat" not in df.columns: + return METRIC_ERROR + + return float(lazy.np.mean(df["avg_lat"])) + def generate_report(self) -> None: - if not self.can_handle_directory(): + df = extract_osu_bench_data(self.results_file) + if df.empty: return - df = extract_osu_bench_data(self.results_file) df.to_csv(self.test_run.output_path / "osu_bench.csv", index=False) diff --git a/tests/report_generation_strategy/test_osu_bench_report_generation_strategy.py b/tests/report_generation_strategy/test_osu_bench_report_generation_strategy.py index c13ff9fc0..ae72bae1d 100644 --- a/tests/report_generation_strategy/test_osu_bench_report_generation_strategy.py +++ b/tests/report_generation_strategy/test_osu_bench_report_generation_strategy.py @@ -15,10 +15,16 @@ # limitations under the License. from pathlib import Path +from unittest.mock import Mock import pytest -from cloudai.workloads.osu_bench.report_generation_strategy import extract_osu_bench_data +from cloudai.core import METRIC_ERROR, TestRun +from cloudai.systems.slurm import SlurmSystem +from cloudai.workloads.osu_bench.report_generation_strategy import ( + OSUBenchReportGenerationStrategy, + extract_osu_bench_data, +) OSU_MULTIPLE_BW = """\ # OSU MPI Multiple Bandwidth / Message Rate Test v7.4 @@ -112,6 +118,25 @@ def test_osu_multi_latency_short_header_parsing(tmp_path: Path) -> None: assert df["avg_lat"].tolist() == pytest.approx([1.88, 1.84, 1.88, 1.91, 1.87, 2.01]) +@pytest.mark.parametrize("metric", OSUBenchReportGenerationStrategy.metrics) +def test_get_metric_returns_mean_latency(metric: str, tmp_path: Path, slurm_system: SlurmSystem) -> None: + (tmp_path / "stdout.txt").write_text(OSU_ALLGATHER_LAT) + tr = TestRun(name="osu", test=Mock(), num_nodes=2, nodes=[], output_path=tmp_path) + strategy = OSUBenchReportGenerationStrategy(slurm_system, tr) + + assert strategy.get_metric(metric) == pytest.approx((2.81 + 104.30) / 2) + + +def test_get_metric_returns_error_for_unsupported_metric_or_output(tmp_path: Path, slurm_system: SlurmSystem) -> None: + (tmp_path / "stdout.txt").write_text(OSU_BW) + tr = TestRun(name="osu", test=Mock(), num_nodes=2, nodes=[], output_path=tmp_path) + strategy = OSUBenchReportGenerationStrategy(slurm_system, tr) + + for metric in strategy.metrics: + assert strategy.get_metric(metric) is METRIC_ERROR + assert strategy.get_metric("unsupported") is METRIC_ERROR + + def test_extract_osu_bench_data_file_not_found_returns_empty_dataframe(tmp_path: Path) -> None: missing = tmp_path / "nonexistent.txt" df = extract_osu_bench_data(missing)