ai-experiment/classify_runner.py

178 lines
6.5 KiB
Python

import sys
import json
import os
import time
import subprocess
import urllib.request
from pathlib import Path
from core import classify_room as ap
from config.classify import (
MODEL_ROOT, VARIANT, CLI, CHECKPOINT, BACKEND,
ALLOW_TRUNCATION, LAYA_URL, TIMEOUT, PRESETS, DEFAULT_PRESET,
)
def load_state(args):
if args.state_file is not None:
path = Path(args.state_file)
if not path.is_file():
print(f"State file not found: {path}", file=sys.stderr)
return None
raw = path.read_text(encoding="utf-8")
elif args.state is not None:
raw = args.state
else:
print("No state given. Pass it as an argument or with --state-file.", file=sys.stderr)
return None
raw = raw.strip()
if raw[:1] in ("{", "["): # JSON object/array state is forwarded as a structure
try:
return json.loads(raw)
except json.JSONDecodeError as e:
print(f"State looks like JSON but does not parse: {e}", file=sys.stderr)
return None
return raw
def load_questions(args):
if args.questions is not None:
path = Path(args.questions)
if not path.is_file():
print(f"Questions file not found: {path}", file=sys.stderr)
return None
try:
data = json.loads(path.read_text(encoding="utf-8"))
except json.JSONDecodeError as e:
print(f"Questions file is not valid JSON: {e}", file=sys.stderr)
return None
# Accept a bare questions map or a full laya.cpp request object.
data = data.get("questions", data) if isinstance(data, dict) else None
if not isinstance(data, dict) or not data:
print("Questions file must hold a non-empty questions object.", file=sys.stderr)
return None
return data
name = args.preset if args.preset is not None else DEFAULT_PRESET
if name not in PRESETS:
print(f"Unknown preset: {name}. Available: {', '.join(PRESETS)}", file=sys.stderr)
return None
return PRESETS[name]
def cli_predict(state, questions):
"""Spawn laya-cli for a single request. Weights are reloaded on every call."""
if not CLI.is_file():
raise FileNotFoundError(f"laya-cli not found: {CLI}. Build it first (see README).")
cmd = [str(CLI), "--model", str(MODEL_ROOT), "--variant", VARIANT, "--" + BACKEND]
if ALLOW_TRUNCATION:
cmd.append("--allow-truncation")
request = json.dumps({"state": state, "questions": questions}, ensure_ascii=False)
proc = subprocess.run(cmd, input=request, capture_output=True, text=True)
if proc.returncode != 0:
raise RuntimeError(proc.stderr.strip() or "laya-cli exited with an error")
out = proc.stdout.strip()
if not out:
raise RuntimeError("laya-cli produced no output")
return json.loads(out.splitlines()[-1]) # later lines win; "Ready:" goes to stderr
def server_predict(url, state, questions):
"""POST /v1/systemone so the checkpoint stays resident across runs."""
payload = json.dumps({"state": state, "questions": questions}, ensure_ascii=False).encode("utf-8")
req = urllib.request.Request(
url.rstrip("/") + "/v1/systemone",
data=payload,
headers={"Content-Type": "application/json"},
)
with urllib.request.urlopen(req, timeout=TIMEOUT) as resp:
return json.load(resp)
def predict(state, questions, url):
if url:
try:
return server_predict(url, state, questions), "server"
except Exception as e:
print(f"laya.cpp server at {url} unavailable ({e}). Falling back to the CLI.", file=sys.stderr)
return cli_predict(state, questions), "cli"
def answer_value(answer):
kind = answer.get("type")
if kind == "choice":
return str(answer.get("choice", ""))
if kind == "score":
score = float(answer.get("score", 0.0))
label = answer.get("legend", {}).get(str(int(round(score))))
return f"{score:.4f}" + (f" ({label})" if label else "")
return f"{float(answer.get('noul', 0.0)):.4f}"
def detail_line(answer):
probs = answer.get("probabilities")
if probs:
# Choice options read best by confidence; score levels stay in legend order.
items = probs.items() if answer.get("type") == "score" else sorted(
probs.items(), key=lambda kv: -float(kv[1]))
return " ".join(f"{k}={float(v):.4f}" for k, v in items)
act = answer.get("action", {}).get("act_probability")
return f"act_probability={float(act):.4f}" if act is not None else None
def classify_run(args):
if not CHECKPOINT.is_file():
print(f"Model not found: {CHECKPOINT}. Download it first (see README).", file=sys.stderr)
sys.exit(1)
state = load_state(args)
questions = load_questions(args)
if state is None or questions is None:
sys.exit(1)
url = args.url if args.url is not None else os.environ.get("LAYA_URL", LAYA_URL)
start = time.time()
try:
envelope, mode = predict(state, questions, url)
except Exception as e:
print(f"Inference failed: {e}", file=sys.stderr)
sys.exit(1)
elapsed = time.time() - start
if args.json:
print(json.dumps(envelope, ensure_ascii=False, indent=2))
return
results = envelope.get("results")
if results is not None and not results:
print("Inference failed: laya.cpp returned no results", file=sys.stderr)
sys.exit(1)
result = results[0] if results else envelope
answers = result.get("answers", {})
usage = result.get("usage", {})
width = max((len(q) for q in answers), default=0)
preview = state if isinstance(state, str) else json.dumps(state, ensure_ascii=False)
if len(preview) > 120:
preview = preview[:120] + "..."
print()
print(f"State : {preview}")
print(f"Questions : {len(answers)} ({mode})")
if envelope.get("backend"):
print(f"Backend : {envelope['backend']} ({envelope.get('device', '')})")
print(f"Elapsed : {elapsed:.3f} s")
if "elapsed_ms" in envelope:
print(f"Inference : {envelope['elapsed_ms'] / 1000:.3f} s")
if usage:
print(f"Tokens : {usage.get('input_tokens', 0)} in / {usage.get('output_tokens', 0)} out")
print()
for qid, answer in answers.items():
detail = detail_line(answer)
print(f"{qid.ljust(width)} : {answer_value(answer)}")
if detail:
print(f"{' ' * width} {detail}")
print()
if __name__ == "__main__":
classify_run( ap.parser.parse_args() )