diff --git a/src/srtctl/cli/mixins/benchmark_stage.py b/src/srtctl/cli/mixins/benchmark_stage.py index f3769a9f0..e772359a6 100644 --- a/src/srtctl/cli/mixins/benchmark_stage.py +++ b/src/srtctl/cli/mixins/benchmark_stage.py @@ -9,6 +9,8 @@ import logging import shlex +import signal +import subprocess import threading import time from pathlib import Path @@ -134,10 +136,16 @@ def run_benchmark( logger.info("Running %s benchmark", runner.name) + # Start perf monitoring on all worker nodes (non-fatal if it fails) + perf_procs = self._start_perf_monitor() + # Run the benchmark script benchmark_log = self.runtime.log_dir / "benchmark.out" exit_code = self._run_benchmark_script(runner, benchmark_log, stop_event) + # Stop monitoring regardless of benchmark outcome + self._stop_perf_monitor(perf_procs) + if exit_code != 0: logger.error("Benchmark failed with exit code %d", exit_code) else: @@ -145,6 +153,78 @@ def run_benchmark( return exit_code + def _start_perf_monitor(self) -> list[tuple[str, "subprocess.Popen"]]: + """Start one perfmon process per worker node. + + Failures are non-fatal: a warning is logged and that node is skipped. + + Returns: + List of (node, Popen) pairs for processes that started successfully. + """ + m = self.config.monitoring + if m is None or not m.enabled: + return [] + + worker_nodes = list(self.runtime.nodes.worker) + if not worker_nodes: + logger.warning("No worker nodes to monitor") + return [] + + perfmon_script = Path(__file__).parent.parent.parent / "monitor" / "perfmon.py" + mounts = dict(self.runtime.container_mounts) + mounts[perfmon_script] = Path("/tmp/srt_perfmon.py") + + procs: list[tuple[str, subprocess.Popen]] = [] + for node in worker_nodes: + cmd = [ + "python3", "/tmp/srt_perfmon.py", + "--output-csv", f"/logs/perf_samples_{node}.csv", + "--output-json", f"/logs/perf_summary_{node}.json", + "--interval", str(m.sample_interval), + ] + perf_log = self.runtime.log_dir / f"perf_monitor_{node}.out" + try: + proc = start_srun_process( + command=cmd, + nodelist=[node], + output=str(perf_log), + container_image=str(self.runtime.container_image), + container_mounts=mounts, + ) + procs.append((node, proc)) + logger.info("perf monitor started on %s (interval=%.1fs)", node, m.sample_interval) + except Exception as e: + logger.warning("Failed to start perf monitor on %s: %s - monitoring skipped for this node", node, e) + + return procs + + def _stop_perf_monitor(self, procs: list[tuple[str, "subprocess.Popen"]]) -> None: + """Stop all perfmon processes, allowing each to write its summary JSON. + + Sends SIGINT (triggers perfmon's exit handler) and waits up to 30s. + Falls back to SIGKILL if the process does not exit cleanly. + """ + if not procs: + return + + logger.info("Stopping perf monitoring on %d node(s)", len(procs)) + for node, proc in procs: + if proc.poll() is not None: + logger.warning("perf monitor on %s already exited (code %d)", node, proc.returncode) + continue + try: + proc.send_signal(signal.SIGINT) + except ProcessLookupError: + logger.warning("perf monitor on %s vanished before SIGINT", node) + continue + try: + proc.wait(timeout=30) + logger.info("perf monitor on %s stopped cleanly", node) + except subprocess.TimeoutExpired: + logger.warning("perf monitor on %s did not stop within 30s, killing", node) + proc.kill() + proc.wait() + def _run_benchmark_script( self, runner: "BenchmarkRunner", diff --git a/src/srtctl/core/schema.py b/src/srtctl/core/schema.py index 085db6c82..6b615d74d 100644 --- a/src/srtctl/core/schema.py +++ b/src/srtctl/core/schema.py @@ -801,6 +801,25 @@ class OutputConfig: Schema: ClassVar[type[Schema]] = Schema +@dataclass(frozen=True) +class MonitoringConfig: + """Built-in GPU performance monitoring during benchmark execution. + + When enabled, one perfmon process runs per worker node (excluding the head node) + and writes per-node output files to the job log directory: + - perf_samples_{node}.csv per-second time-series (GPU util, memory, power, temp) + - perf_summary_{node}.json aggregate statistics over the benchmark window + + Uses nvidia-smi — no external dependencies required. + Failures are non-fatal: monitoring is skipped for affected nodes, benchmark continues. + """ + + enabled: bool = True + sample_interval: float = 1.0 + + Schema: ClassVar[type[Schema]] = Schema + + @dataclass(frozen=True) class HealthCheckConfig: """Health check configuration.""" @@ -872,6 +891,9 @@ class SrtConfig: # Reporting configuration (status API, future: logs to S3, etc.) reporting: ReportingConfig | None = None + # Built-in GPU performance monitoring (runs on all worker nodes during benchmark) + monitoring: MonitoringConfig | None = None + Schema: ClassVar[type[Schema]] = Schema def __post_init__(self): diff --git a/src/srtctl/monitor/perfmon.py b/src/srtctl/monitor/perfmon.py new file mode 100644 index 000000000..771476049 --- /dev/null +++ b/src/srtctl/monitor/perfmon.py @@ -0,0 +1,103 @@ +#!/usr/bin/env python3 +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Lightweight GPU performance monitor. + +Polls nvidia-smi at a fixed interval and writes: + - per-second CSV samples (--output-csv) + - aggregate summary JSON (--output-json, written on SIGINT/exit) + +Usage: + python3 perfmon.py --output-csv /logs/perf_samples_node1.csv \\ + --output-json /logs/perf_summary_node1.json \\ + --interval 1.0 +""" + +import argparse +import csv +import json +import signal +import subprocess +import time +from datetime import datetime, timezone +from pathlib import Path + +_QUERY = "index,utilization.gpu,memory.used,memory.total,power.draw,temperature.gpu" +_FIELDS = ["gpu", "util_pct", "mem_used_mb", "mem_total_mb", "power_w", "temp_c"] + + +def _sample() -> list[dict]: + try: + out = subprocess.check_output( + ["nvidia-smi", f"--query-gpu={_QUERY}", "--format=csv,noheader,nounits"], + text=True, + ) + except (FileNotFoundError, subprocess.CalledProcessError): + return [] + rows = [] + for line in out.strip().splitlines(): + parts = [p.strip() for p in line.split(",")] + if len(parts) == len(_FIELDS): + rows.append(dict(zip(_FIELDS, parts))) + return rows + + +def _summarize(samples: list[dict]) -> dict: + by_gpu: dict[str, list[dict]] = {} + for s in samples: + by_gpu.setdefault(s["gpu"], []).append(s) + + summary = {} + for gpu_idx, gpu_samples in by_gpu.items(): + + def avg(field: str, _s: list[dict] = gpu_samples) -> float | None: + vals = [float(s[field]) for s in _s if s.get(field, "").strip() not in ("", "[N/A]")] + return round(sum(vals) / len(vals), 2) if vals else None + + summary[f"gpu_{gpu_idx}"] = { + "samples": len(gpu_samples), + "avg_util_pct": avg("util_pct"), + "avg_mem_used_mb": avg("mem_used_mb"), + "mem_total_mb": avg("mem_total_mb"), + "avg_power_w": avg("power_w"), + "avg_temp_c": avg("temp_c"), + } + return summary + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--output-csv", required=True) + parser.add_argument("--output-json", required=True) + parser.add_argument("--interval", type=float, default=1.0) + args = parser.parse_args() + + samples: list[dict] = [] + stop = False + + def handle_sigint(sig, frame): + nonlocal stop + stop = True + + signal.signal(signal.SIGINT, handle_sigint) + + with Path(args.output_csv).open("w", newline="") as f: + writer: csv.DictWriter | None = None + while not stop: + ts = datetime.now(timezone.utc).isoformat() + for row in _sample(): + record = {"timestamp": ts, **row} + if writer is None: + writer = csv.DictWriter(f, fieldnames=list(record.keys())) + writer.writeheader() + writer.writerow(record) + samples.append(record) + f.flush() + time.sleep(args.interval) + + if samples: + Path(args.output_json).write_text(json.dumps(_summarize(samples), indent=2)) + + +if __name__ == "__main__": + main() diff --git a/tests/test_monitoring.py b/tests/test_monitoring.py new file mode 100644 index 000000000..7e5dc5962 --- /dev/null +++ b/tests/test_monitoring.py @@ -0,0 +1,227 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for built-in GPU performance monitoring configuration and benchmark stage integration.""" + +import signal +import subprocess +from pathlib import Path +from unittest.mock import MagicMock, patch + +import pytest + +from srtctl.core.schema import ( + BenchmarkConfig, + ModelConfig, + MonitoringConfig, + ResourceConfig, + SrtConfig, +) + + +# --------------------------------------------------------------------------- +# Schema +# --------------------------------------------------------------------------- + + +class TestMonitoringConfig: + def test_defaults(self): + m = MonitoringConfig() + assert m.enabled is True + assert m.sample_interval == 1.0 + + def test_disabled(self): + m = MonitoringConfig(enabled=False) + assert m.enabled is False + + def test_custom_interval(self): + m = MonitoringConfig(sample_interval=2.5) + assert m.sample_interval == 2.5 + + def test_yaml_roundtrip(self): + schema = MonitoringConfig.Schema() + result = schema.load({"enabled": True, "sample_interval": 0.5}) + assert result.enabled is True + assert result.sample_interval == 0.5 + + +class TestSrtConfigMonitoringField: + def _base_config(self, **kwargs) -> SrtConfig: + return SrtConfig( + name="test", + model=ModelConfig(path="/model", container="/image.sqsh", precision="fp8"), + resources=ResourceConfig(gpu_type="h100", agg_nodes=1, agg_workers=1), + benchmark=BenchmarkConfig(type="manual"), + **kwargs, + ) + + def test_monitoring_defaults_to_none(self): + assert self._base_config().monitoring is None + + def test_monitoring_accepted(self): + config = self._base_config(monitoring=MonitoringConfig()) + assert config.monitoring is not None + assert config.monitoring.enabled is True + + def test_monitoring_disabled(self): + config = self._base_config(monitoring=MonitoringConfig(enabled=False)) + assert config.monitoring.enabled is False + + +# --------------------------------------------------------------------------- +# BenchmarkStageMixin._start_perf_monitor / _stop_perf_monitor +# --------------------------------------------------------------------------- + + +def _make_orchestrator(monitoring=None): + from srtctl.cli.mixins.benchmark_stage import BenchmarkStageMixin + + config = SrtConfig( + name="test", + model=ModelConfig(path="/model", container="/image.sqsh", precision="fp8"), + resources=ResourceConfig(gpu_type="h100", agg_nodes=2, agg_workers=2, gpus_per_node=8), + benchmark=BenchmarkConfig(type="sa-bench"), + monitoring=monitoring, + ) + + runtime = MagicMock() + runtime.nodes.head = "node0" + runtime.nodes.worker = ("node0", "node1", "node2") + runtime.container_image = Path("/image.sqsh") + runtime.container_mounts = {} + runtime.log_dir = Path("/logs/12345") + + class Orchestrator(BenchmarkStageMixin): + @property + def endpoints(self): + return [] + + @property + def backend_processes(self): + return [] + + orch = Orchestrator.__new__(Orchestrator) + orch.config = config + orch.runtime = runtime + return orch + + +class TestStartPerfMonitor: + def test_returns_empty_when_monitoring_is_none(self): + assert _make_orchestrator(monitoring=None)._start_perf_monitor() == [] + + def test_returns_empty_when_disabled(self): + assert _make_orchestrator(monitoring=MonitoringConfig(enabled=False))._start_perf_monitor() == [] + + def test_starts_one_proc_per_worker_node_including_head(self): + """Starts perfmon on all worker nodes including head (node0).""" + orch = _make_orchestrator(monitoring=MonitoringConfig()) + mock_proc = MagicMock() + + with patch("srtctl.cli.mixins.benchmark_stage.start_srun_process", return_value=mock_proc) as mock_srun: + result = orch._start_perf_monitor() + + assert len(result) == 3 + nodes = [node for node, _ in result] + assert "node0" in nodes + assert "node1" in nodes + assert "node2" in nodes + + def test_output_paths_keyed_by_hostname(self): + orch = _make_orchestrator(monitoring=MonitoringConfig()) + mock_proc = MagicMock() + + with patch("srtctl.cli.mixins.benchmark_stage.start_srun_process", return_value=mock_proc) as mock_srun: + orch._start_perf_monitor() + + all_cmds = [c.kwargs["command"] for c in mock_srun.call_args_list] + for node in ("node0", "node1", "node2"): + node_cmds = [cmd for cmd in all_cmds if any(node in arg for arg in cmd)] + assert any(f"perf_samples_{node}.csv" in arg for arg in node_cmds[0]) + assert any(f"perf_summary_{node}.json" in arg for arg in node_cmds[0]) + + def test_perfmon_script_mounted_in_container(self): + orch = _make_orchestrator(monitoring=MonitoringConfig()) + mock_proc = MagicMock() + + with patch("srtctl.cli.mixins.benchmark_stage.start_srun_process", return_value=mock_proc) as mock_srun: + orch._start_perf_monitor() + + for c in mock_srun.call_args_list: + mounts = c.kwargs["container_mounts"] + assert Path("/tmp/srt_perfmon.py") in mounts.values() + + def test_failed_node_is_skipped(self): + orch = _make_orchestrator(monitoring=MonitoringConfig()) + mock_proc = MagicMock() + + def srun_side_effect(**kwargs): + if kwargs["nodelist"] == ["node1"]: + raise RuntimeError("srun failed") + return mock_proc + + with patch("srtctl.cli.mixins.benchmark_stage.start_srun_process", side_effect=srun_side_effect): + result = orch._start_perf_monitor() + + assert len(result) == 2 + nodes = [node for node, _ in result] + assert "node0" in nodes + assert "node2" in nodes + + +class TestStopPerfMonitor: + def test_noop_on_empty_list(self): + _make_orchestrator()._stop_perf_monitor([]) + + def test_sends_sigint_to_running_process(self): + orch = _make_orchestrator() + mock_proc = MagicMock() + mock_proc.poll.return_value = None + mock_proc.wait.return_value = 0 + + orch._stop_perf_monitor([("node1", mock_proc)]) + + mock_proc.send_signal.assert_called_once_with(signal.SIGINT) + mock_proc.wait.assert_called_once_with(timeout=30) + + def test_skips_already_exited_process(self): + orch = _make_orchestrator() + mock_proc = MagicMock() + mock_proc.poll.return_value = 0 + mock_proc.returncode = 0 + + orch._stop_perf_monitor([("node1", mock_proc)]) + + mock_proc.send_signal.assert_not_called() + + def test_kills_process_on_timeout(self): + orch = _make_orchestrator() + mock_proc = MagicMock() + mock_proc.poll.return_value = None + mock_proc.wait.side_effect = [subprocess.TimeoutExpired(cmd="perfmon", timeout=30), None] + + orch._stop_perf_monitor([("node1", mock_proc)]) + + mock_proc.kill.assert_called_once() + + def test_handles_process_lookup_error(self): + orch = _make_orchestrator() + mock_proc = MagicMock() + mock_proc.poll.return_value = None + mock_proc.send_signal.side_effect = ProcessLookupError + + orch._stop_perf_monitor([("node1", mock_proc)]) # must not raise + + def test_stops_all_nodes(self): + orch = _make_orchestrator() + procs = [] + for node in ("node1", "node2", "node3"): + mock_proc = MagicMock() + mock_proc.poll.return_value = None + mock_proc.wait.return_value = 0 + procs.append((node, mock_proc)) + + orch._stop_perf_monitor(procs) + + for _, mock_proc in procs: + mock_proc.send_signal.assert_called_once_with(signal.SIGINT)