system-one-arcade / game /server.py
mlabonne's picture Aurelien-Lac's picture
system-one-arcade
ffc4ad0
Raw History Blame Contribute Delete
12.7 kB
#!/usr/bin/env python3
"""Camera games: the page sends webcam frames, and game states as text, the
model answers the game's questions about each in one pass, the page plays.
python game/server.py # NVIDIA GPU, Apple silicon (mps) or CPU
python game/server.py --compile # NVIDIA GPU: CUDA graphs, ~1 min to start
python game/server.py --no-token --host 0.0.0.0 --port 7860 # inside the Space's container
The model runs on its own code from the Hub (`trust_remote_code`), as its card describes.
POST /decide {"image": <JPEG data URL> | "state": <text or object>,
"questions": [{"text": ..., "options": {label: description}}]}
-> {"probs": [[p, ...], ...], "ms": <model time>}
GET /ws a WebSocket carrying the same requests, each with an "id" its answer repeats;
several can be in flight, and no request pays HTTP's headers or a proxy's checks
A question without options is yes/no and answers [p(yes), p(no)]; one with
"levels" (a list, lowest first) is a score and answers one p per level. One
question takes one forward pass over image or state and question; several share it.
"""
from __future__ import annotations
import argparse
import base64
import hashlib
import io
import json
import mimetypes
import secrets
import struct
import sys
import threading
import time
from concurrent.futures import ThreadPoolExecutor
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from PIL import Image
WEB = Path(__file__).parent / "web"
MAX_QUESTIONS = 32
MAX_BODY = 4 << 20 # a 512 px JPEG is ~60 KB
WS_GUID = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11" # RFC 6455
def question(q: dict) -> dict:
"""The page's question as the model's API takes it."""
if q.get("options"):
return {"type": "choice", "instructions": q["text"], "criteria": dict(q["options"])}
if q.get("levels"):
return {"type": "score", "instructions": q["text"], "criteria": list(q["levels"])}
return {"type": "noul", "instructions": q["text"]}
def frame(data_url: str, side: int) -> Image.Image:
image = Image.open(io.BytesIO(base64.b64decode(data_url.split(",")[-1]))).convert("RGB")
image.thumbnail((side, side))
return image
def decide(model, questions: list, image=None, state=None) -> tuple[list[list[float]], float]:
t0 = time.perf_counter()
probs = model.probabilities(state, questions, () if image is None else [image]) if questions else []
return probs, 1e3 * (time.perf_counter() - t0)
def parse(req: dict, side: int) -> tuple[list, dict]:
questions = [question(q) for q in req["questions"][:MAX_QUESTIONS]]
source = {"state": req["state"]} if "state" in req else {"image": frame(req["image"], side)}
return questions, source
class WebSocket:
"""The server side of RFC 6455, enough for the page: text messages, ping, close."""
def __init__(self, rfile, wfile):
self.rfile, self.wfile, self.lock = rfile, wfile, threading.Lock()
def _exact(self, n: int) -> bytes:
data = self.rfile.read(n)
if len(data) < n:
raise ConnectionError("closed")
return data
def _frame(self) -> tuple[int, bool, bytes]:
b0, b1 = self._exact(2)
n = b1 & 0x7F
if n == 126:
n = struct.unpack(">H", self._exact(2))[0]
elif n == 127:
n = struct.unpack(">Q", self._exact(8))[0]
if n > MAX_BODY:
raise ConnectionError("message too large")
mask = self._exact(4) if b1 & 0x80 else b""
data = self._exact(n)
if mask:
data = _unmask(data, mask)
return b0 & 0x0F, bool(b0 & 0x80), data
def receive(self) -> str | None:
"""The next text message, or None once the page closes the socket."""
parts: list[bytes] = []
while True:
op, fin, data = self._frame()
if op == 0x8:
self.send(data[:2], 0x8)
return None
if op == 0x9:
self.send(data, 0xA)
continue
if op in (0x1, 0x0):
parts.append(data)
if sum(map(len, parts)) > MAX_BODY:
raise ConnectionError("message too large")
if fin:
return b"".join(parts).decode()
def send(self, data: bytes | str, op: int = 0x1) -> None:
if isinstance(data, str):
data = data.encode()
n = len(data)
head = bytes([0x80 | op]) + (bytes([n]) if n < 126 else
bytes([126]) + struct.pack(">H", n) if n < 65536 else
bytes([127]) + struct.pack(">Q", n))
with self.lock:
self.wfile.write(head + data)
self.wfile.flush()
def _unmask(data: bytes, mask: bytes) -> bytes:
n = len(data)
key = int.from_bytes((mask * (n // 4 + 1))[:n], "big")
return (int.from_bytes(data, "big") ^ key).to_bytes(n, "big")
def make_handler(run, token: str | None, side: int):
"""`run(questions, image=..., state=...)` answers on the model's own thread. With `token`
None every request is served: in a Space, its visibility decides who can reach it."""
class Handler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1" # keep-alive: one connection for every frame
disable_nagle_algorithm = True # else headers and body wait ~40 ms for an ACK
def log_message(self, *args):
pass
def _send(self, code: int, body: bytes, kind: str) -> None:
try:
self.send_response(code)
self.send_header("Content-Type", kind)
self.send_header("Content-Length", str(len(body)))
self.send_header("Cache-Control", "no-store")
self.end_headers()
self.wfile.write(body)
except (BrokenPipeError, ConnectionResetError):
pass # the page closed or reloaded while its frame was answered
def do_GET(self):
if self.path.split("?")[0] == "/ws":
self._websocket()
return
name = self.path.split("?")[0].lstrip("/") or "index.html"
path = (WEB / name).resolve()
if WEB.resolve() not in path.parents or not path.is_file():
self._send(404, b"not found", "text/plain")
return
kind = mimetypes.guess_type(path.name)[0] or "application/octet-stream"
if kind in ("text/javascript", "application/javascript"):
kind = "text/javascript; charset=utf-8"
self._send(200, path.read_bytes(), kind)
def do_POST(self):
size = int(self.headers.get("Content-Length", 0))
if self.path.split("?")[0] != "/decide" or (token and self.headers.get("X-Token") != token) or size > MAX_BODY:
self._send(403, b"forbidden", "text/plain")
self.close_connection = True
return
try:
questions, source = parse(json.loads(self.rfile.read(size)), side)
except Exception: # noqa: BLE001 - anything malformed in the request
self._send(400, b"bad request", "text/plain")
return
probs, ms = run(questions, **source)
self._send(200, json.dumps({"probs": probs, "ms": ms}).encode(), "application/json")
def _websocket(self):
query = dict(p.split("=", 1) for p in self.path.partition("?")[2].split("&") if "=" in p)
key = self.headers.get("Sec-WebSocket-Key")
if "websocket" not in self.headers.get("Upgrade", "").lower() or not key or (token and query.get("token") != token):
self._send(403, b"forbidden", "text/plain")
self.close_connection = True
return
self.send_response(101)
self.send_header("Upgrade", "websocket")
self.send_header("Connection", "Upgrade")
self.send_header("Sec-WebSocket-Accept",
base64.b64encode(hashlib.sha1((key + WS_GUID).encode()).digest()).decode())
self.end_headers()
self.close_connection = True
ws = WebSocket(self.rfile, self.wfile)
# requests are decoded off the reading thread, so the next frame can arrive while one is answered
pool = ThreadPoolExecutor(max_workers=4)
def answer(text: str) -> None:
req = None
try:
req = json.loads(text)
questions, source = parse(req, side)
except Exception: # noqa: BLE001 - anything malformed in the request
out = {"error": "bad request"}
else:
try:
probs, ms = run(questions, **source)
out = {"probs": probs, "ms": ms}
except Exception as e: # noqa: BLE001 - the page must hear back, or its frame stays in flight
out = {"error": type(e).__name__}
out["id"] = req.get("id") if isinstance(req, dict) else None
try:
ws.send(json.dumps(out))
except OSError:
pass # the page went away while its frame was answered
try:
while (text := ws.receive()) is not None:
pool.submit(answer, text)
except (ConnectionError, OSError):
pass
finally:
pool.shutdown(wait=False, cancel_futures=True)
return Handler
class Server(ThreadingHTTPServer):
daemon_threads = True
def handle_error(self, request, client_address):
# a page closed or reloaded mid-request is routine, not an error to print
if not isinstance(sys.exc_info()[1], (BrokenPipeError, ConnectionResetError)):
super().handle_error(request, client_address)
def load(model_id: str, compile: bool):
"""The model's System One engine, on the best device here."""
import torch
from transformers import AutoModel
device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
model = AutoModel.from_pretrained(model_id, trust_remote_code=True, dtype=torch.bfloat16).to(device)
if compile:
model.compile(mode="reduce-overhead")
print(f"{model_id} on {device}", flush=True)
return model.engine
def main() -> int:
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--model", default="LiquidAI/d1-3B")
p.add_argument("--host", default="127.0.0.1")
p.add_argument("--port", type=int, default=8765)
p.add_argument("--side", type=int, default=512, help="frames are cut down to fit this square")
p.add_argument("--compile", action="store_true", help="CUDA graphs for one-question frames (NVIDIA GPU)")
p.add_argument("--no-token", action="store_true", help="serve without a token, for a Space that controls who can reach it")
args = p.parse_args()
model = load(args.model, args.compile)
# Every model call runs on this one thread: requests queue in order, and
# CUDA graphs are bound to the thread that captured them.
worker = ThreadPoolExecutor(max_workers=1)
def run(questions, **source):
return worker.submit(decide, model, questions, **source).result()
# The first calls of a shape pay kernel selection, or compilation with
# --compile: the games' frames (384 px, and the duel's half frames), two
# prompt lengths so that later ones reuse the dynamic graph.
print("warming up" + (" and compiling, about a minute" if args.compile else ""), flush=True)
short = question({"text": "Is there a person?"})
long_ = question({"text": "Is the person holding a cup in one hand?"})
for size in ((384, 216), (192, 216), (512, 288)):
warm = Image.new("RGB", size, (120, 120, 120))
for q in (short, long_):
run([q], image=warm)
run([short, long_], image=warm)
run([short, long_], state={"scene": "an empty room"})
token = None if args.no_token else secrets.token_urlsafe(12)
server = Server((args.host, args.port), make_handler(run, token, args.side))
print(f"ready: open http://localhost:{args.port}/" + (f"?token={token}" if token else ""), flush=True)
server.serve_forever()
return 0
if __name__ == "__main__":
raise SystemExit(main())