Laya-demo / server.py
robynsd's picture
Temporarily compare PyTorch and ONNX on the Space's hardware
c966bc6
Raw History Blame
5.25 kB
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")