From f34c513c4149a5244a47d5a8681866268a2a498b Mon Sep 17 00:00:00 2001 From: Alex Manley Date: Fri, 4 Sep 2026 15:08:55 -0700 Subject: [PATCH 1/2] Fix OSU-Bench get_metric bug --- .../osu_bench/report_generation_strategy.py | 9 ++++++- ...st_osu_bench_report_generation_strategy.py | 25 ++++++++++++++++++- 2 files changed, 32 insertions(+), 2 deletions(-) diff --git a/src/cloudai/workloads/osu_bench/report_generation_strategy.py b/src/cloudai/workloads/osu_bench/report_generation_strategy.py index 0dc1a5aed..af515e633 100644 --- a/src/cloudai/workloads/osu_bench/report_generation_strategy.py +++ b/src/cloudai/workloads/osu_bench/report_generation_strategy.py @@ -23,7 +23,7 @@ from pathlib import Path from typing import TYPE_CHECKING -from cloudai.core import ReportGenerationStrategy +from cloudai.core import METRIC_ERROR, MetricValue, ReportGenerationStrategy from cloudai.util.lazy_imports import lazy if TYPE_CHECKING: @@ -148,6 +148,13 @@ 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: + df = extract_osu_bench_data(self.results_file) + if df.empty or metric != "default" 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(): return 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..e1342b027 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,23 @@ 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]) +def test_get_metric_returns_mean_latency(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("default") == 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) + + assert strategy.get_metric("default") is METRIC_ERROR + assert strategy.get_metric("avg_lat") 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) From 28b9426bf9ea7d50489facc768e3abe91ddb83cb Mon Sep 17 00:00:00 2001 From: Alex Manley Date: Wed, 9 Sep 2026 09:48:50 -0700 Subject: [PATCH 2/2] minor improvements to OSU bench reporter --- .../osu_bench/report_generation_strategy.py | 13 +++++++++---- .../test_osu_bench_report_generation_strategy.py | 10 ++++++---- 2 files changed, 15 insertions(+), 8 deletions(-) diff --git a/src/cloudai/workloads/osu_bench/report_generation_strategy.py b/src/cloudai/workloads/osu_bench/report_generation_strategy.py index af515e633..fed09e290 100644 --- a/src/cloudai/workloads/osu_bench/report_generation_strategy.py +++ b/src/cloudai/workloads/osu_bench/report_generation_strategy.py @@ -21,7 +21,7 @@ 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 METRIC_ERROR, MetricValue, ReportGenerationStrategy from cloudai.util.lazy_imports import lazy @@ -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" @@ -149,15 +151,18 @@ def can_handle_directory(self) -> bool: 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 metric != "default" or "avg_lat" not in df.columns: + 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 e1342b027..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 @@ -118,12 +118,13 @@ 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]) -def test_get_metric_returns_mean_latency(tmp_path: Path, slurm_system: SlurmSystem) -> None: +@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("default") == pytest.approx((2.81 + 104.30) / 2) + 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: @@ -131,8 +132,9 @@ def test_get_metric_returns_error_for_unsupported_metric_or_output(tmp_path: Pat tr = TestRun(name="osu", test=Mock(), num_nodes=2, nodes=[], output_path=tmp_path) strategy = OSUBenchReportGenerationStrategy(slurm_system, tr) - assert strategy.get_metric("default") is METRIC_ERROR - assert strategy.get_metric("avg_lat") is METRIC_ERROR + 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: