import os import platform import time from collections import defaultdict from pathlib import Path from fastapi import FastAPI, HTTPException from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import FileResponse, Response from pydantic import BaseModel from questions import CASES WEB = Path(__file__).parent / "web" # Three engines run the same Laya weights and return the same answers: # - "mlx": laya-mlx, Apple Silicon only, fast local runs; # - "torch": the original laya package (PyTorch), for Linux servers; # - "onnx": the same checkpoint exported to ONNX at startup (onnx_backend.py), CPU only. UPSTREAM = {"typed": "convaiinnovations/laya-typed-decisions", "english": "convaiinnovations/laya"} CHECKPOINTS = { "mlx": {"typed": "aac6fef/laya-typed-decisions-mlx", "english": "aac6fef/laya-mlx"}, "torch": UPSTREAM, "onnx": UPSTREAM, } APPLE_SILICON = platform.system() == "Darwin" and platform.machine() == "arm64" BACKEND = os.environ.get("LAYA_BACKEND", "mlx" if APPLE_SILICON else "torch") # Comma-separated aliases; "typed" alone halves memory (about 2.7 GB with PyTorch). # Questions routed to a model that is not loaded fall back to "typed". LOADED = os.environ.get("LAYA_MODELS", "typed,english").split(",") MODELS = {alias: CHECKPOINTS[BACKEND][alias] for alias in LOADED} if BACKEND == "mlx": import laya_mlx as laya else: import laya app = FastAPI() # The page may be served from another origin (e.g. Vercel) than the API. origins = [o for o in os.environ.get("ALLOWED_ORIGINS", "").split(",") if o] if origins: app.add_middleware(CORSMiddleware, allow_origins=origins, allow_methods=["GET", "POST"], allow_headers=["Content-Type"]) # Optional device override ("cpu" to mimic a GPU-less server when testing locally). device = {"device": os.environ["LAYA_DEVICE"]} if os.environ.get("LAYA_DEVICE") else {} def load_agent(backend, model_id): if backend == "onnx": import onnx_backend return onnx_backend.load(model_id) return laya.load(model_id, **device) agents = {alias: load_agent(BACKEND, model_id) for alias, model_id in MODELS.items()} for agent in agents.values(): agent.predict("warmup", {"q": {"type": "noul", "instructions": "Is this a test?"}}) class Request(BaseModel): case: str text: str def get_case(name): if name not in CASES: raise HTTPException(404, f"unknown case {name}") return CASES[name] def model_for(case, key): alias = case["models"].get(key, "typed") return alias if alias in agents else "typed" @app.post("/predict") def predict(req: Request): case = get_case(req.case) by_model = defaultdict(dict) for key, question in case["questions"]().items(): by_model[model_for(case, key)][key] = question answers, input_tokens, output_tokens = {}, 0, 0 start = time.perf_counter() for alias, questions in by_model.items(): result = agents[alias].predict(req.text, questions) answers.update(result["answers"]) input_tokens += result["usage"]["input_tokens"] output_tokens += result["usage"]["output_tokens"] ms = (time.perf_counter() - start) * 1000 return { **case["summary"](answers), "ms": round(ms, 1), "input_tokens": input_tokens, "output_tokens": output_tokens, "models": sorted(MODELS[alias] for alias in by_model), } @app.get("/cases/{name}") def explain(name: str): case = get_case(name) return {**case["explain"](), "routing": {key: MODELS[model_for(case, key)] for key in case["questions"]()}} # Temporary benchmark (LAYA_BENCH=1): times PyTorch and ONNX on this machine, in the # background after startup since it takes about a minute. Results: GET /bench and the logs. bench_state = {"status": "disabled"} def run_bench(): import onnx_backend bench_state["status"] = "running" try: model_id = UPSTREAM["typed"] torch_agent = agents["typed"] if BACKEND == "torch" else load_agent("torch", model_id) engines = {"torch": torch_agent, "onnx": load_agent("onnx", model_id)} bench_state.update(status="done", result=onnx_backend.compare(engines, CASES)) except Exception as error: bench_state.update(status="failed", error=f"{type(error).__name__}: {error}") print(f"bench: {bench_state}", flush=True) if os.environ.get("LAYA_BENCH") == "1": import threading threading.Thread(target=run_bench, daemon=True).start() @app.get("/bench") def bench(): return bench_state @app.get("/health") def health(): return {"backend": BACKEND, "models": list(MODELS.values())} @app.get("/") def index(): return FileResponse(WEB / "index.html") @app.get("/config.js") def config(): # Same-origin API when the page is served by this server; web/config.js # holds the API URL for the static deployment instead. LAYA_DEBOUNCE (ms) lets a # slow CPU server ask the page to wait longer after the last keystroke. script = 'window.LAYA_API = "";' if os.environ.get("LAYA_DEBOUNCE", "").isdigit(): script += f' window.LAYA_DEBOUNCE = {os.environ["LAYA_DEBOUNCE"]};' return Response(script, media_type="text/javascript")