AAJerry's picture
Rebuild SAMS v4 bounded selector and strict 40-case gate
ceb51d3
Raw History Blame Contribute Delete
2.97 kB
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)