Files
YG_FT/compute/agent/process_manager.py

282 lines
10 KiB
Python
Raw Normal View History

from __future__ import annotations
import os
import json
import contextlib
import hashlib
import signal
import subprocess
import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
TERMINAL_STATUSES = {"completed", "failed", "stopped"}
@dataclass
class ManagedProcess:
id: str
name: str
command: list[str]
work_dir: str
log_path: Path
output_dir: str
gpus: list[int]
process: subprocess.Popen[Any] | None
created_at: float
pid: int | None = None
status: str = "running"
progress: int = 5
artifacts: list[dict[str, Any]] = field(default_factory=list)
class ProcessManager:
def __init__(self, log_root: str) -> None:
self.log_root = Path(log_root)
self.log_root.mkdir(parents=True, exist_ok=True)
self.registry_path = self.log_root / "compute-jobs.json"
self.jobs: dict[str, ManagedProcess] = {}
self._load_registry()
def create_job(self, payload: dict[str, Any], command: list[str], work_dir: str) -> dict[str, Any]:
job_id = str(payload.get("id") or f"job_{int(time.time() * 1000)}")
if job_id in self.jobs and self.jobs[job_id].status not in TERMINAL_STATUSES:
raise ValueError(f"job {job_id} is already running")
output_dir = str(payload.get("output_dir") or f"/data/yg-ft/outputs/{payload.get('name', job_id)}")
Path(output_dir).mkdir(parents=True, exist_ok=True)
log_path = self.log_root / f"{job_id}.log"
env = os.environ.copy()
gpus = [int(item) for item in payload.get("gpus") or []]
locked = self.locked_gpus()
conflict = sorted(set(gpus).intersection(locked))
if conflict:
raise ValueError(f"gpu already locked: {conflict}")
if gpus:
env["CUDA_VISIBLE_DEVICES"] = ",".join(str(item) for item in gpus)
env.update({str(k): str(v) for k, v in payload.get("env", {}).items()})
cwd = work_dir if Path(work_dir).exists() else None
with log_path.open("ab") as log_file:
log_file.write(f"[INFO] starting job_id={job_id} command={' '.join(command)}\n".encode("utf-8"))
process = subprocess.Popen(
command,
cwd=cwd,
env=env,
stdout=log_file,
stderr=subprocess.STDOUT,
)
managed = ManagedProcess(
id=job_id,
name=str(payload.get("name") or job_id),
command=command,
work_dir=work_dir,
log_path=log_path,
output_dir=output_dir,
gpus=gpus,
process=process,
created_at=time.time(),
pid=process.pid,
progress=10,
)
self.jobs[job_id] = managed
data = self.serialize(managed)
self._save_registry()
return data
def get_job(self, job_id: str) -> dict[str, Any] | None:
job = self.jobs.get(job_id)
if not job:
return None
return self.serialize(job)
def list_jobs(self) -> list[dict[str, Any]]:
return [self.serialize(job) for job in self.jobs.values()]
def stop_job(self, job_id: str) -> dict[str, Any] | None:
job = self.jobs.get(job_id)
if not job:
return None
if job.status not in TERMINAL_STATUSES:
try:
if job.process is not None and os.name == "nt":
job.process.terminate()
elif job.pid is not None:
os.kill(job.pid, signal.SIGTERM)
if job.process is not None:
job.process.wait(timeout=10)
except Exception:
if job.process is not None:
job.process.kill()
elif job.pid is not None:
with contextlib.suppress(Exception):
os.kill(job.pid, signal.SIGKILL)
job.status = "stopped"
job.progress = min(job.progress, 99)
data = self.serialize(job)
self._save_registry()
return data
def logs(self, job_id: str) -> str:
job = self.jobs.get(job_id)
if not job or not job.log_path.exists():
return ""
return job.log_path.read_text(encoding="utf-8", errors="replace")
def serialize(self, job: ManagedProcess) -> dict[str, Any]:
code = job.process.poll() if job.process is not None else None
checkpoints = self._collect_checkpoints(job.output_dir)
if job.status not in TERMINAL_STATUSES:
if job.process is None and job.pid is not None and not self._pid_alive(job.pid):
job.status = "failed"
job.progress = min(job.progress, 99)
code = -1
elif code is None:
job.status = "running"
elapsed = max(0, int(time.time() - job.created_at))
job.progress = min(95, max(job.progress, 10 + elapsed // 6))
elif code == 0:
job.status = "completed"
job.progress = 100
job.artifacts = self._collect_artifacts(job.output_dir)
else:
job.status = "failed"
job.progress = min(job.progress, 99)
self._save_registry()
return {
"id": job.id,
"name": job.name,
"status": job.status,
"progress": job.progress,
"pid": job.pid,
"gpus": job.gpus,
"created_at": job.created_at,
"command": job.command,
"work_dir": job.work_dir,
"output_dir": job.output_dir,
"log_file": str(job.log_path),
"artifacts": job.artifacts,
"checkpoints": checkpoints,
"return_code": code,
}
def locked_gpus(self) -> set[int]:
locked: set[int] = set()
for job in self.jobs.values():
status = self.serialize(job)["status"]
if status in {"queued", "running"}:
locked.update(job.gpus)
return locked
def _collect_artifacts(self, output_dir: str) -> list[dict[str, Any]]:
root = Path(output_dir)
if not root.exists():
return []
artifacts: list[dict[str, Any]] = []
for path in root.rglob("*"):
if path.is_file():
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
size = path.stat().st_size
artifacts.append(
{
"path": str(path),
"name": path.name,
"size": size,
"size_bytes": size,
"checksum_sha256": digest.hexdigest(),
}
)
return artifacts[:200]
def _collect_checkpoints(self, output_dir: str) -> list[dict[str, Any]]:
root = Path(output_dir)
if not root.exists():
return []
checkpoints: list[dict[str, Any]] = []
for path in root.glob("checkpoint-*"):
if not path.is_dir():
continue
step = 0
try:
step = int(path.name.rsplit("-", 1)[-1])
except ValueError:
step = 0
size_bytes = sum(item.stat().st_size for item in path.rglob("*") if item.is_file())
checkpoints.append(
{
"step": step,
"name": path.name,
"path": str(path),
"size_bytes": size_bytes,
"create_time": path.stat().st_mtime,
}
)
return sorted(checkpoints, key=lambda item: (int(item.get("step") or 0), str(item.get("name") or "")))
def _save_registry(self) -> None:
items = []
for job in self.jobs.values():
items.append(
{
"id": job.id,
"name": job.name,
"command": job.command,
"work_dir": job.work_dir,
"log_path": str(job.log_path),
"output_dir": job.output_dir,
"gpus": job.gpus,
"pid": job.pid,
"created_at": job.created_at,
"status": job.status,
"progress": job.progress,
"artifacts": job.artifacts,
}
)
self.registry_path.write_text(json.dumps(items, ensure_ascii=False, indent=2), encoding="utf-8")
def _load_registry(self) -> None:
if not self.registry_path.exists():
return
try:
items = json.loads(self.registry_path.read_text(encoding="utf-8"))
except json.JSONDecodeError:
return
for item in items if isinstance(items, list) else []:
if not isinstance(item, dict):
continue
pid = item.get("pid")
status = item.get("status", "failed")
if status not in TERMINAL_STATUSES and pid and not self._pid_alive(int(pid)):
status = "failed"
job = ManagedProcess(
id=str(item["id"]),
name=str(item.get("name") or item["id"]),
command=[str(part) for part in item.get("command") or []],
work_dir=str(item.get("work_dir") or ""),
log_path=Path(item.get("log_path") or self.log_root / f"{item['id']}.log"),
output_dir=str(item.get("output_dir") or ""),
gpus=[int(gpu) for gpu in item.get("gpus") or []],
process=None,
pid=int(pid) if pid else None,
created_at=float(item.get("created_at") or time.time()),
status=status,
progress=int(item.get("progress") or 0),
artifacts=item.get("artifacts") or [],
)
self.jobs[job.id] = job
def _pid_alive(self, pid: int) -> bool:
if pid <= 0:
return False
try:
os.kill(pid, 0)
return True
except OSError:
return False