File size: 8,003 Bytes
c1264c3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
#!/usr/bin/env python3
"""HTTP server implementing the jev-adapter /v1/systemone contract with the joint schema head (DESIGN.md 5.1).

  python3 serve_head.py --snapshot <base or merged dir> --adapter <run>/final/adapter --head <run>/final/head \
      --port 30171 [--temperature-file calib.json]

Endpoints: POST /v1/systemone, GET /v1/models, /model_info, /server_info (also /get_model_info, /get_server_info),
/health. The engine snapshot fields the benchmark runner checks are reported (disable_radix_cache=true,
mm_preprocess_cache_size_mb=0, speculative_algorithm=null): every request is a full prefill, nothing is cached.
Dynamic batching: requests waiting in the queue are forwarded together while the padded token count stays within
--max-batch-tokens (image requests run alone). Stdlib HTTP server; binds 127.0.0.1 by default.
"""
import argparse
import json
import queue
import sys
import threading
import time
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path

HERE = Path(__file__).resolve().parent
sys.path.insert(0, str(HERE))

from scoring import RequestError, load_scorer  # noqa: E402


class Batcher:
    def __init__(self, scorer, max_batch_tokens, max_batch_requests, wait_ms):
        self.scorer, self.max_tokens, self.max_requests, self.wait = scorer, max_batch_tokens, max_batch_requests, wait_ms / 1000
        self.q = queue.Queue()
        threading.Thread(target=self.loop, daemon=True).start()

    def submit(self, body):
        started = time.perf_counter()
        req, opts, encs = self.scorer.plan(body)        # validation + encoding in the caller's thread
        job = {"req": req, "opts": opts, "encs": encs, "started": started, "done": threading.Event()}
        self.q.put(job)
        job["done"].wait()
        if "error" in job:
            raise job["error"]
        return job["response"]

    def loop(self):
        while True:
            jobs = [self.q.get()]
            image = bool(jobs[0]["req"]["images"])
            deadline = time.perf_counter() + self.wait
            while not image and len(jobs) < self.max_requests:
                try:
                    nxt = self.q.get(timeout=max(0.0, deadline - time.perf_counter()))
                except queue.Empty:
                    break
                encs = [e for j in jobs + [nxt] for e in j["encs"]]
                cost = len(encs) * max(e["n_tokens"] for e in encs)
                if nxt["req"]["images"] or cost > self.max_tokens:
                    self.q.put(nxt)   # goes to the next batch
                    break
                jobs.append(nxt)
            try:
                encs = [e for j in jobs for e in j["encs"]]
                outs, ps, hs = self.scorer.run(encs)
                k = 0
                for j in jobs:
                    n = len(j["encs"])
                    j["response"] = self.scorer.respond(j["req"], j["opts"], j["encs"], outs[k:k + n], ps, hs, j["started"])
                    j["response"]["metadata"]["batch_requests"] = len(jobs)
                    k += n
            except Exception as e:  # noqa: BLE001
                for j in jobs:
                    j["error"] = e
            for j in jobs:
                j["done"].set()


def make_handler(scorer, batcher, info):
    model_info = {"model_path": info["model_path"], "served_model_name": scorer.served_model, "is_generation": False,
                  "has_image_understanding": True, "model_type": "schema-head", "architectures": ["SchemaHeadModel"]}
    server_info = {"model_path": info["model_path"], "served_model_name": scorer.served_model,
                   "disable_radix_cache": True, "mm_preprocess_cache_size_mb": 0, "enable_prefix_mm_cache": False,
                   "enable_mm_global_cache": False, "speculative_algorithm": None, "context_length": 32768,
                   "dtype": info["dtype"], "attention_backend": "torch-sdpa", "version": "schema-head-serve-v1"}

    class Handler(BaseHTTPRequestHandler):
        protocol_version = "HTTP/1.1"

        def log_message(self, fmt, *args):
            pass

        def _send(self, code, obj):
            data = json.dumps(obj, ensure_ascii=False, allow_nan=False).encode()
            self.send_response(code)
            self.send_header("Content-Type", "application/json")
            self.send_header("Content-Length", str(len(data)))
            self.end_headers()
            self.wfile.write(data)

        def do_GET(self):
            path = self.path.split("?")[0].rstrip("/")
            if path == "/health":
                return self._send(200, {"status": "ok"})
            if path == "/v1/models":
                return self._send(200, {"object": "list", "data": [{"id": scorer.served_model, "object": "model",
                                                                     "owned_by": "schema-head", "label_scheme": "schema-v1"}]})
            if path in ("/model_info", "/get_model_info"):
                return self._send(200, model_info)
            if path in ("/server_info", "/get_server_info"):
                return self._send(200, server_info)
            return self._send(404, {"error": {"code": "not_found", "message": path}})

        def do_POST(self):
            if self.path.split("?")[0].rstrip("/") != "/v1/systemone":
                return self._send(404, {"error": {"code": "not_found", "message": self.path}})
            try:
                body = json.loads(self.rfile.read(int(self.headers.get("Content-Length", "0"))))
                return self._send(200, batcher.submit(body))
            except RequestError as e:
                return self._send(e.status, {"error": {"code": e.code, "message": str(e), "field": e.field}})
            except json.JSONDecodeError as e:
                return self._send(400, {"error": {"code": "invalid_json", "message": str(e)}})
            except Exception as e:  # noqa: BLE001
                return self._send(500, {"error": {"code": "scoring_failed", "message": f"{type(e).__name__}: {e}"}})
    return Handler


def main():
    ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("--snapshot", type=Path)
    ap.add_argument("--tokenizer", type=Path)
    ap.add_argument("--tiny", help="JSON tiny config (tests only)")
    ap.add_argument("--adapter", type=Path)
    ap.add_argument("--head", type=Path, required=True)
    ap.add_argument("--served-model-name", default="standardthinking/standard-schema-8b")
    ap.add_argument("--temperature-file", type=Path)
    ap.add_argument("--host", default="127.0.0.1")
    ap.add_argument("--port", type=int, default=30171)
    ap.add_argument("--device", default="auto")
    ap.add_argument("--max-batch-tokens", type=int, default=32768)
    ap.add_argument("--max-batch-requests", type=int, default=8)
    ap.add_argument("--batch-wait-ms", type=float, default=5.0)
    ap.add_argument("--memory-cap-bytes", type=float, default=0, help="torch allocator cap on GPU (0 = none)")
    args = ap.parse_args()
    tiny = json.loads(args.tiny) if args.tiny else None
    scorer, info = load_scorer(args.snapshot, args.head, args.adapter, tiny, args.tokenizer, args.served_model_name,
                               args.temperature_file, args.device, max_concurrency=args.max_batch_requests,
                               memory_cap_bytes=args.memory_cap_bytes)
    batcher = Batcher(scorer, args.max_batch_tokens, args.max_batch_requests, args.batch_wait_ms)
    meta = {"model_path": str(args.snapshot or "tiny"), "dtype": "bfloat16" if scorer.device.type == "cuda" else "float32"}
    server = ThreadingHTTPServer((args.host, args.port), make_handler(scorer, batcher, meta))
    print(json.dumps({"serving": f"http://{args.host}:{args.port}", "model": args.served_model_name,
                      "head_params": info["head_params"], "device": str(scorer.device)}), flush=True)
    server.serve_forever()


if __name__ == "__main__":
    main()