from __future__ import annotations import json import os import sys import time from datetime import datetime, timezone from pathlib import Path from typing import Any APP_ROOT = Path(__file__).resolve().parents[1] PERSIST_ROOT = Path(os.environ.get("PERSIST_ROOT", "/data")) try: PERSIST_ROOT.mkdir(parents=True, exist_ok=True) except OSError: PERSIST_ROOT = APP_ROOT / "data" PERSIST_ROOT.mkdir(parents=True, exist_ok=True) if not os.access(PERSIST_ROOT, os.W_OK): raise PermissionError(f"Persistent training root is not writable: {PERSIST_ROOT}") RAW_ROOT = PERSIST_ROOT / "raw" WORK_ROOT = PERSIST_ROOT / "work" DATASET_ROOT = PERSIST_ROOT / "dataset" CHECKPOINT_ROOT = PERSIST_ROOT / "checkpoints" ARTIFACT_ROOT = PERSIST_ROOT / "artifacts" STATUS_ROOT = PERSIST_ROOT / "status" LOG_ROOT = PERSIST_ROOT / "logs" def ensure_dirs() -> None: for path in ( RAW_ROOT, WORK_ROOT, DATASET_ROOT, CHECKPOINT_ROOT, ARTIFACT_ROOT, STATUS_ROOT, LOG_ROOT, ): path.mkdir(parents=True, exist_ok=True) def utc_now() -> str: return datetime.now(timezone.utc).isoformat() def atomic_json(path: Path, payload: dict[str, Any]) -> None: path.parent.mkdir(parents=True, exist_ok=True) tmp = path.with_suffix(path.suffix + ".tmp") with tmp.open("w", encoding="utf-8", newline="\n") as handle: handle.write(json.dumps(payload, indent=2, sort_keys=True)) handle.write("\n") tmp.replace(path) def read_json(path: Path, default: Any = None) -> Any: try: return json.loads(path.read_text(encoding="utf-8")) except (FileNotFoundError, json.JSONDecodeError): return default def update_status(phase: str, message: str, **extra: Any) -> None: ensure_dirs() current = read_json(STATUS_ROOT / "status.json", {}) or {} current.update( { "phase": phase, "message": message, "updated_at": utc_now(), "pid": os.getpid(), **extra, } ) atomic_json(STATUS_ROOT / "status.json", current) print(f"[{utc_now()}] {phase}: {message}", flush=True) def log_event(event: str, **fields: Any) -> None: ensure_dirs() row = {"time": utc_now(), "event": event, **fields} with (LOG_ROOT / "events.jsonl").open("a", encoding="utf-8") as handle: handle.write(json.dumps(row, sort_keys=True) + "\n") def wait_for_gpu(timeout_seconds: int = 1800) -> None: import torch started = time.time() while not torch.cuda.is_available(): if time.time() - started > timeout_seconds: raise RuntimeError("CUDA GPU did not become available before timeout") update_status("waiting_for_gpu", "Waiting for the requested L40S runtime") time.sleep(30) def fail(message: str) -> None: update_status("failed", message) log_event("failed", message=message) print(message, file=sys.stderr, flush=True)