airfRANS-model-exploration/scripts/aggressive_oom_node_wrapper.py

338 lines
13 KiB
Python
Raw Normal View History

2026-07-28 18:04:50 +00:00
#!/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())