178 lines
6.5 KiB
Python
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() )
|