#!/usr/bin/env python3 from __future__ import annotations import argparse import json import shutil import threading import time import traceback from pathlib import Path from typing import Any from airfrans_frontier.remote.artifacts import verify_artifacts from airfrans_frontier.sweep import collect_job_status, run_node, sample_gpu_utilization _BASE_ARTIFACTS = ( "config.toml", "job_manifest.json", "metrics.jsonl", "latest_metrics.json", "heartbeat.json", "run_manifest.json", "environment_manifest.json", "utilization.jsonl", ) _TERMINAL_ARTIFACTS = ( "final_metrics.json", "failure_report.json", "checkpoint_latest.pt", "checkpoint_best.pt", "checkpoint_final.pt", "split_manifest.json", "data_manifest.json", "normalization.json", "calibration_manifest.json", "evaluation_protocol.json", "hf_upload_manifest.json", "verification_report.json", "artifact_manifest.json", "checksums.txt", ) _PRESERVE_IN_CURRENT_RUN = {"startup_timeline.jsonl", "nvidia_smi.txt", "disk_telemetry.json"} def main() -> int: parser = argparse.ArgumentParser(description="Run one aggressive OOM sweep node and flatten the latest terminal job artifact for remote-run collection.") parser.add_argument("--jobs", default="artifacts/aggressive_oom_sweep/jobs.jsonl") parser.add_argument("--node", required=True) parser.add_argument("--artifact-dir", default="artifacts/current_run") parser.add_argument("--utilization", default="artifacts/aggressive_oom_sweep/utilization.jsonl") parser.add_argument("--max-jobs", type=int) parser.add_argument("--sample-interval-seconds", type=float, default=30.0) parser.add_argument("--stale-after-seconds", type=float, default=21600.0) parser.add_argument("--max-attempts", type=int, default=2) parser.add_argument("--require-success", action="store_true") args = parser.parse_args() jobs_path = Path(args.jobs) artifact_dir = Path(args.artifact_dir) utilization_path = Path(args.utilization) artifact_dir.mkdir(parents=True, exist_ok=True) utilization_path.parent.mkdir(parents=True, exist_ok=True) started_at = time.time() stop_sampling = threading.Event() sampler = threading.Thread( target=_sample_until_stopped, args=(stop_sampling, utilization_path, artifact_dir, args.sample_interval_seconds, started_at, args.node), daemon=True, ) _write_node_heartbeat(artifact_dir, node=args.node, phase="starting", started_at=started_at) _append_jsonl(artifact_dir / "metrics.jsonl", _node_metric(args.node, phase="starting", started_at=started_at)) sampler.start() exit_code = 0 summary: dict[str, Any] | None = None status: dict[str, Any] | None = None try: summary = run_node( jobs_path=jobs_path, node_name=args.node, max_jobs=args.max_jobs, stale_after_seconds=args.stale_after_seconds, max_attempts=args.max_attempts, ) status = collect_job_status(jobs_path=jobs_path) if args.require_success and int(summary.get("succeeded", 0)) <= 0: exit_code = 1 except Exception as exc: # noqa: BLE001 - terminal failure artifact is the contract here. exit_code = 1 status = _safe_collect(jobs_path) summary = { "node": args.node, "claimed": 0, "succeeded": 0, "failed": 1, "error_type": type(exc).__name__, "error_message": str(exc), "traceback_tail": traceback.format_exc()[-4000:], } finally: stop_sampling.set() sampler.join(timeout=max(1.0, min(10.0, args.sample_interval_seconds))) finished_at = time.time() status = status or _safe_collect(jobs_path) summary = summary or {"node": args.node, "claimed": 0, "succeeded": 0, "failed": 0} summary = dict(summary) summary.update( { "started_at": started_at, "finished_at": finished_at, "duration_seconds": finished_at - started_at, "jobs_path": str(jobs_path), "utilization_path": str(utilization_path), "status_counts": (status or {}).get("counts", {}), } ) selected = _select_terminal_record(status or {}) if selected is not None and selected.get("run_dir"): _flatten_job_artifacts(Path(str(selected["run_dir"])), artifact_dir) summary["flattened_job_id"] = selected.get("job_id") summary["flattened_run_dir"] = selected.get("run_dir") summary["flattened_status"] = selected.get("status") else: exit_code = 1 _write_minimal_failure_artifacts( artifact_dir, node=args.node, started_at=started_at, finished_at=finished_at, summary=summary, message="no terminal job artifacts were available to flatten", ) _write_sweep_sidecars( artifact_dir=artifact_dir, jobs_path=jobs_path, utilization_path=utilization_path, summary=summary, status=status or {}, node=args.node, started_at=started_at, finished_at=finished_at, ) try: verify_artifacts(artifact_dir) except Exception as exc: # noqa: BLE001 exit_code = 1 if not (artifact_dir / "failure_report.json").is_file(): for terminal_name in ("final_metrics.json", "checkpoint_latest.pt", "checkpoint_best.pt", "checkpoint_final.pt"): path = artifact_dir / terminal_name if path.exists(): path.unlink() _write_minimal_failure_artifacts( artifact_dir, node=args.node, started_at=started_at, finished_at=time.time(), summary=summary, message=f"flattened artifact verification failed: {exc}", ) verify_artifacts(artifact_dir) print(json.dumps({"exit_code": exit_code, "summary": summary, "status_counts": (status or {}).get("counts", {})}, indent=2, sort_keys=True)) return exit_code def _sample_until_stopped(stop: threading.Event, utilization_path: Path, artifact_dir: Path, interval: float, started_at: float, node: str) -> None: interval = max(1.0, interval) while not stop.is_set(): sample = sample_gpu_utilization() sample.update({"node": node, "elapsed_seconds": time.time() - started_at}) _append_jsonl(utilization_path, sample) _append_jsonl(artifact_dir / "node_utilization.jsonl", sample) _write_node_heartbeat(artifact_dir, node=node, phase="running", started_at=started_at, latest_utilization=sample) _append_jsonl(artifact_dir / "metrics.jsonl", _node_metric(node, phase="running", started_at=started_at, latest_utilization=sample)) stop.wait(interval) def _node_metric(node: str, *, phase: str, started_at: float, latest_utilization: dict[str, Any] | None = None) -> dict[str, Any]: payload: dict[str, Any] = { "timestamp": time.time(), "node": node, "phase": phase, "elapsed_seconds": time.time() - started_at, } if latest_utilization: payload.update( { "gpu_util_percent": latest_utilization.get("gpu_util_percent"), "memory_used_mb": latest_utilization.get("memory_used_mb"), } ) return payload def _write_node_heartbeat(artifact_dir: Path, *, node: str, phase: str, started_at: float, latest_utilization: dict[str, Any] | None = None) -> None: _write_json( artifact_dir / "heartbeat.json", { "node": node, "phase": phase, "timestamp": time.time(), "started_at": started_at, "elapsed_seconds": time.time() - started_at, "latest_utilization": latest_utilization, }, ) def _safe_collect(jobs_path: Path) -> dict[str, Any]: try: return collect_job_status(jobs_path=jobs_path) except Exception as exc: # noqa: BLE001 return {"jobs": [], "counts": {"collect_error": 1}, "total": 0, "error": str(exc)} def _select_terminal_record(status: dict[str, Any]) -> dict[str, Any] | None: jobs = [dict(item) for item in status.get("jobs", []) if isinstance(item, dict)] succeeded = [item for item in jobs if item.get("status") == "succeeded" and item.get("run_dir")] if succeeded: return max(succeeded, key=lambda item: _mtime(Path(str(item["run_dir"])))) failed = [item for item in jobs if item.get("status") == "failed" and item.get("run_dir")] if failed: return max(failed, key=lambda item: _mtime(Path(str(item["run_dir"])))) incomplete = [item for item in jobs if item.get("run_dir")] if incomplete: return max(incomplete, key=lambda item: _mtime(Path(str(item["run_dir"])))) return None def _mtime(path: Path) -> float: try: return path.stat().st_mtime except FileNotFoundError: return 0.0 def _flatten_job_artifacts(run_dir: Path, artifact_dir: Path) -> None: if not run_dir.is_dir(): raise FileNotFoundError(f"terminal run directory not found: {run_dir}") for path in artifact_dir.iterdir(): if path.is_file() and path.name not in _PRESERVE_IN_CURRENT_RUN: path.unlink() for name in (*_BASE_ARTIFACTS, *_TERMINAL_ARTIFACTS): source = run_dir / name if source.is_file(): shutil.copy2(source, artifact_dir / name) def _write_minimal_failure_artifacts(artifact_dir: Path, *, node: str, started_at: float, finished_at: float, summary: dict[str, Any], message: str) -> None: for name in ("final_metrics.json", "checkpoint_latest.pt", "checkpoint_best.pt", "checkpoint_final.pt"): path = artifact_dir / name if path.exists(): path.unlink() base = { "run_id": node, "run_name": node, "node": node, "started_at": started_at, "finished_at": finished_at, "artifact_dir": str(artifact_dir), "summary": summary, } _write_json(artifact_dir / "config.toml", {"note": "placeholder"}) if False else None if not (artifact_dir / "config.toml").is_file(): (artifact_dir / "config.toml").write_text("[run]\nname = \"sweep_node_failure\"\n") _write_json(artifact_dir / "job_manifest.json", {**base, "command": "aggressive_oom_node_wrapper"}) _write_json(artifact_dir / "run_manifest.json", base) _write_json(artifact_dir / "environment_manifest.json", {"node": node, "recorded_at": time.time()}) if not (artifact_dir / "metrics.jsonl").is_file(): _append_jsonl(artifact_dir / "metrics.jsonl", {"timestamp": finished_at, "phase": "failed", "node": node}) _write_json(artifact_dir / "latest_metrics.json", {"timestamp": finished_at, "phase": "failed", "node": node}) _write_node_heartbeat(artifact_dir, node=node, phase="failed", started_at=started_at) if not (artifact_dir / "utilization.jsonl").is_file(): _append_jsonl(artifact_dir / "utilization.jsonl", {"timestamp": finished_at, "gpu_util_percent": None, "memory_used_mb": None}) _write_json( artifact_dir / "failure_report.json", { "run_id": node, "phase": "sweep_node", "failure_phase": "sweep_node", "failure_category": "sweep_node_artifact_collection", "error_type": "RuntimeError", "error_message": message, "summary": summary, "timestamp": finished_at, }, ) def _write_sweep_sidecars(*, artifact_dir: Path, jobs_path: Path, utilization_path: Path, summary: dict[str, Any], status: dict[str, Any], node: str, started_at: float, finished_at: float) -> None: _write_json(artifact_dir / "node_summary.json", summary) _write_json(artifact_dir / "sweep_collect.json", status) if jobs_path.is_file(): shutil.copy2(jobs_path, artifact_dir / "sweep_jobs.jsonl") else: (artifact_dir / "sweep_jobs.jsonl").write_text("") attempts_path = jobs_path.parent / "attempts.jsonl" if attempts_path.is_file(): shutil.copy2(attempts_path, artifact_dir / "sweep_attempts.jsonl") else: (artifact_dir / "sweep_attempts.jsonl").write_text("") if utilization_path.is_file(): shutil.copy2(utilization_path, artifact_dir / "sweep_utilization.jsonl") shutil.copy2(utilization_path, artifact_dir / "utilization.jsonl") elif not (artifact_dir / "utilization.jsonl").is_file(): _append_jsonl(artifact_dir / "utilization.jsonl", {"timestamp": finished_at, "gpu_util_percent": None, "memory_used_mb": None}) _append_jsonl( artifact_dir / "node_metrics.jsonl", { "timestamp": finished_at, "node": node, "phase": "finished", "elapsed_seconds": finished_at - started_at, "summary": summary, "status_counts": status.get("counts", {}), }, ) def _write_json(path: Path, payload: Any) -> None: path.parent.mkdir(parents=True, exist_ok=True) path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n") def _append_jsonl(path: Path, payload: Any) -> None: path.parent.mkdir(parents=True, exist_ok=True) with path.open("a", encoding="utf-8") as handle: handle.write(json.dumps(payload, sort_keys=True) + "\n") if __name__ == "__main__": raise SystemExit(main())