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() )