Download server.py from robynsd/Laya-demo: direct link, hf CLI and curl.
- Browser
- Download file 5.25 kB
-
https://huggingface.co/spaces/robynsd/Laya-demo/resolve/d1ec01f84509a8ed3a4a47ef731d3bc99765578c/server.py
- Command line
-
hf download hf://spaces/robynsd/Laya-demo@d1ec01f84509a8ed3a4a47ef731d3bc99765578c/server.py
-
curl -L -o server.py https://huggingface.co/spaces/robynsd/Laya-demo/resolve/d1ec01f84509a8ed3a4a47ef731d3bc99765578c/server.py
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" | |
| 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), | |
| } | |
| 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() | |
| def bench(): | |
| return bench_state | |
| def health(): | |
| return {"backend": BACKEND, "models": list(MODELS.values())} | |
| def index(): | |
| return FileResponse(WEB / "index.html") | |
| 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") | |