AAJerry's picture
Constrain v7 to prevalidated grounded wording choices
bc06d0b
Raw History Blame Contribute Delete
5.79 kB
from __future__ import annotations
import html
import json
import os
import subprocess
import sys
import threading
import time
from contextlib import asynccontextmanager
from pathlib import Path
import psutil
from fastapi import FastAPI, HTTPException
from fastapi.responses import FileResponse, HTMLResponse, JSONResponse, PlainTextResponse
from huggingface_hub import HfApi
from trainer.common import ARTIFACT_ROOT, LOG_ROOT, STATUS_ROOT, ensure_dirs, read_json
PIPELINE_PID = STATUS_ROOT / "pipeline.pid"
PIPELINE_LOG = LOG_ROOT / "pipeline.log"
def cost_watchdog() -> None:
maximum = int(os.environ.get("MAX_PAID_SECONDS", "0"))
token = os.environ.get("TRAINING_HF_TOKEN")
space_id = os.environ.get("SPACE_ID")
if maximum <= 0 or not token or not space_id:
return
time.sleep(maximum)
try:
HfApi(token=token).pause_space(space_id)
except Exception as exc:
print(f"Cost watchdog could not pause the Space: {exc}", flush=True)
def pipeline_running() -> bool:
try:
pid = int(PIPELINE_PID.read_text(encoding="utf-8").strip())
return psutil.pid_exists(pid)
except (FileNotFoundError, ValueError):
return False
def start_pipeline() -> None:
ensure_dirs()
if pipeline_running():
return
status = read_json(STATUS_ROOT / "status.json", {}) or {}
if status.get("phase") == "complete" and os.environ.get("FORCE_RETRAIN", "0") != "1":
return
log_handle = PIPELINE_LOG.open("a", encoding="utf-8", buffering=1)
process = subprocess.Popen(
[sys.executable, "-m", "trainer.run_pipeline"],
cwd=Path(__file__).resolve().parent,
stdout=log_handle,
stderr=subprocess.STDOUT,
start_new_session=True,
)
PIPELINE_PID.write_text(str(process.pid), encoding="utf-8")
def tail(path: Path, max_bytes: int = 30000) -> str:
if not path.exists():
return "No logs yet."
with path.open("rb") as handle:
handle.seek(0, 2)
size = handle.tell()
handle.seek(max(0, size - max_bytes))
return handle.read().decode("utf-8", errors="replace")
@asynccontextmanager
async def lifespan(_: FastAPI):
ensure_dirs()
threading.Thread(target=cost_watchdog, name="cost-watchdog", daemon=True).start()
if os.environ.get("AUTO_START", "1") == "1":
start_pipeline()
yield
app = FastAPI(title="SAMS Qwen v7 Approved-Wording Evaluator", lifespan=lifespan)
@app.get("/health")
def health() -> dict[str, object]:
return {"ok": True, "pipeline_running": pipeline_running()}
@app.get("/api/status")
def status() -> JSONResponse:
payload = read_json(STATUS_ROOT / "status.json", {}) or {"phase": "starting"}
payload["pipeline_running"] = pipeline_running()
payload["persistent_root"] = str(STATUS_ROOT.parent)
return JSONResponse(payload)
@app.get("/api/artifacts")
def artifact_index() -> JSONResponse:
items = []
for path in sorted(ARTIFACT_ROOT.rglob("*")):
if path.is_file():
relative = path.relative_to(ARTIFACT_ROOT).as_posix()
items.append({"path": relative, "size": path.stat().st_size, "url": f"/download/{relative}"})
return JSONResponse({"artifacts": items})
@app.get("/download/{artifact_path:path}")
def download_artifact(artifact_path: str) -> FileResponse:
requested = (ARTIFACT_ROOT / artifact_path).resolve()
root = ARTIFACT_ROOT.resolve()
if requested != root and root not in requested.parents:
raise HTTPException(status_code=400, detail="Invalid artifact path")
if not requested.is_file():
raise HTTPException(status_code=404, detail="Artifact not found")
return FileResponse(requested, filename=requested.name)
@app.get("/logs")
def logs() -> PlainTextResponse:
return PlainTextResponse(tail(PIPELINE_LOG))
@app.get("/")
def index() -> HTMLResponse:
status = read_json(STATUS_ROOT / "status.json", {}) or {"phase": "starting", "message": "Initializing"}
artifacts = sorted(str(path.relative_to(ARTIFACT_ROOT)) for path in ARTIFACT_ROOT.rglob("*") if path.is_file())
body = f"""
<!doctype html>
<html><head><meta charset="utf-8"><meta http-equiv="refresh" content="30">
<title>SAMS Qwen v7 Approved-Wording Evaluator</title>
<style>
body {{ font-family: system-ui, sans-serif; max-width: 1050px; margin: 2rem auto; padding: 0 1rem; background:#0b1020; color:#eef2ff; }}
.card {{ background:#151c33; border:1px solid #2c385f; border-radius:14px; padding:1rem 1.2rem; margin:1rem 0; }}
code, pre {{ background:#090d18; border-radius:8px; }} pre {{ padding:1rem; overflow:auto; max-height:34rem; white-space:pre-wrap; }}
.phase {{ color:#8dd8ff; font-size:1.4rem; font-weight:700; }} a {{ color:#9bd2ff; }}
</style></head><body>
<h1>SAMS Qwen3-1.7B approved-wording selector v7</h1>
<p>Source-grounded synthetic data requiring domain review. Safety decisions bypass this model.</p>
<div class="card"><div class="phase">{html.escape(str(status.get('phase', 'unknown')))}</div>
<p>{html.escape(str(status.get('message', '')))}</p>
<p>Updated: {html.escape(str(status.get('updated_at', 'not yet')))} - Process running: {pipeline_running()}</p></div>
<div class="card"><h2>Status</h2><pre>{html.escape(json.dumps(status, indent=2, sort_keys=True))}</pre></div>
<div class="card"><h2>Artifacts</h2><pre>{html.escape(chr(10).join(artifacts) or 'None yet')}</pre></div>
<div class="card"><h2>Recent pipeline log</h2><pre>{html.escape(tail(PIPELINE_LOG))}</pre></div>
</body></html>
"""
return HTMLResponse(body)
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=int(os.environ.get("PORT", "7860")))