diff --git a/README.md b/README.md index 35bc8b9..af05fcb 100644 --- a/README.md +++ b/README.md @@ -57,7 +57,7 @@ LAYA can run as an external Python worker for text requests; its in-repository m | Model | Status | | --- | --- | -| LAYA | [External worker](recipe/laya/README.md); [CPU checkpoint reader](src/models/laya/README.md); model execution planned | +| LAYA | [External worker](recipe/laya/README.md); [Python worker on Apple Silicon (MPS) and CPU](recipe/laya/apple-silicon.md); [CPU checkpoint reader](src/models/laya/README.md); model execution planned | | Cua-S1 4B 0.2 (`text` adapter) | [Python worker](recipe/cua_s1/text.md); [native worker](recipe/cua_s1/native.md), CUDA, run on sm_89 | CUDA and Metal coverage will be documented per model as implementations are added and validated. diff --git a/recipe/README.md b/recipe/README.md index 3795fd9..6215704 100644 --- a/recipe/README.md +++ b/recipe/README.md @@ -2,6 +2,8 @@ - [Laya text worker](laya/README.md): start the external Python worker, connect the Rust frontend and compare direct and proxied responses. +- [Laya on Apple Silicon](laya/apple-silicon.md): serve Laya on the Mac GPU with the Laya + worker, put the frontend in front of it and run the benchmarks. - [Cua-S1 4B 0.2 text worker](cua_s1/text.md): download the pinned weights, start the worker and connect the Rust frontend. - [Cua-S1 4B 0.2 native text worker](cua_s1/native.md): build the CUDA library and diff --git a/recipe/laya/README.md b/recipe/laya/README.md index b4f0497..254c12c 100644 --- a/recipe/laya/README.md +++ b/recipe/laya/README.md @@ -3,7 +3,8 @@ This recipe runs the external Laya Python package behind the Rust frontend. It validates text decisions; image, audio and video inference are not covered. -Run all commands from the repository root. +Run all commands from the repository root. To serve on the GPU of an Apple Silicon Mac, see +[Laya on Apple Silicon](apple-silicon.md). ## Start the worker diff --git a/recipe/laya/apple-silicon.md b/recipe/laya/apple-silicon.md new file mode 100644 index 0000000..9fb097f --- /dev/null +++ b/recipe/laya/apple-silicon.md @@ -0,0 +1,192 @@ +# Laya on Apple Silicon + +This recipe serves Laya on the GPU of an Apple Silicon Mac (PyTorch MPS) with the worker in +[`src/frontend/laya_mps.py`](../../src/frontend/laya_mps.py), puts the Rust frontend in front of it and +runs the benchmarks. The model-side code is in [`src/models/laya/`](../../src/models/laya/). The +[Laya text worker](README.md) recipe covers plain laya-serve on the CPU. + +Validated on an M1 Pro (16 GB, 16-core GPU), macOS 26.1, Python 3.12, `laya[serve]==0.3.20`, +torch 2.14.0 and the `english` checkpoint (`convaiinnovations/laya` at `55cf4c4`), and by another +contributor on an M5 (10-core GPU, 32 GB, macOS 26.5.2). Other M-series Macs have not been tested. + +Run all commands from the repository root. + +## Install + +Use Python 3.12. If `python3.12` is not on your `PATH` and you have uv, replace the first command below +with `uv venv --python 3.12 --seed .venv` (`--seed` puts pip in the environment). + +```sh +python3.12 -m venv .venv +.venv/bin/python -m pip install -r recipe/laya/requirements-mps.txt +.venv/bin/python -c "import torch; print(torch.backends.mps.is_available())" +``` + +The last command must print `True`. The standard macOS arm64 wheel of torch includes MPS. + +## Start the worker + +```sh +PYTHONPATH=src .venv/bin/python -m frontend.laya_mps --device mps --model english --require-device --port 8000 +``` + +First startup downloads the checkpoint (846 MB; 97 s into an empty cache at 8.7 MB/s when measured). The +worker loads the model, runs a warmup +over short, long and multi-question requests, and only then listens on port 8000, so the first +request it accepts is already warm: on an M1 Pro the first request after ready took 70–81 ms, against +0.7–1.1 s from plain laya-serve. `--require-device` makes it exit instead of silently serving on the CPU +when the model cannot be placed on MPS; without it the worker logs a warning and serves from the CPU. +Of laya-serve's environment variables, `LAYA_API_KEY` (bearer authentication) still applies. Those its +launcher reads do not: device, model, host, port and log level are the flags above, and `LAYA_THREADS` and +`LAYA_AUTO_TASK` are not read. The worker warns at startup if any of them is set. + +Laya loads another checkpoint when a request names it (`"model": "multilingual"`) or its routing picks it +(a non-English state). The worker prepares that checkpoint the same way inside that first request; other +requests wait behind it, `/health` names it under `preparing` meanwhile and lists it afterwards. On the +M1 Pro that first request took about 5–10 s without the options and about 70 s with `--compile` (plus the +download the first time, 680 MB for `multilingual`). The frontend gives a backend 60 s, so with `--compile` +it answered that request with 504 while the worker finished preparing; the same request sent again then +took 35 ms. To avoid that, send one request for each further checkpoint straight to the worker after +startup. Each resident checkpoint needs its own memory (see Troubleshooting). + +With `--require-device`, a checkpoint that does not land on the requested device is unloaded again and +the request fails with 500. If Laya evicted another checkpoint to make room for it (it keeps two by +default), the worker loads that one again. + +Check what it is running on: + +```sh +curl -s http://127.0.0.1:8000/health +``` + +`device` must be `mps` and `device_mismatch` `false`. The response also names the checkpoint and +revision, the weight dtype (`torch.float32`; Laya upcasts the fp16 checkpoint on MPS), the autocast +dtype Laya uses for requests with at least `mps_amp_min_rows` questions, and the warmup time. The device +and dtypes are read on every call: if a request runs out of GPU memory, Laya moves the model to the CPU +and keeps serving, and `/health` then shows `device: cpu` and `device_mismatch: true`. Triggered on the +M1 Pro by lowering PyTorch's MPS memory limit: the request that ran out of memory still returned 200 +after about 30 s, and later 68-token requests took 140–270 ms from the CPU. + +### Faster: compile and fp16 weights + +```sh +PYTHONPATH=src .venv/bin/python -m frontend.laya_mps --device mps --model english --require-device \ + --compile --weights fp16 --port 8000 +``` + +`--compile` compiles the model during warmup: one-question requests run the whole model compiled, +requests with several questions run the encoder compiled and Laya's decision head as it is. +`--weights fp16` keeps the checkpoint's fp16 weights instead of Laya's fp32 upcast on MPS. + +On the M1 Pro, with both workers running and every request sent to each back to back, the two options +together lowered warm p50 against the worker without them by 37–38% for a 68-token one-question +request (about 57 → 35 ms in those runs), 17–20% at 198–484 tokens, 14% for three questions and 18% for +six. Answers stayed within 0.0031 of the fp32 worker's. A worker running on its own uses about 3 GB +with the options instead of 4.2 GB (2.8 GB against 3.5 GB in those paired runs, where the two workers +shared the machine), measured on the six benchmark inputs; see below for how it grows. The price is startup: the worker became ready after 35–39 s instead of 8–10 s, and +its first request after that took 62–78 ms. + +On an M5 the same paired comparison gave median ratios of 0.51–0.53 for one-question requests at +47–68 tokens, 0.30–0.33 at 198–484 tokens, 0.37 for three questions and 0.60 for six, most of it from +the fp16 weights, which on that GPU speed up every input even without compile. There the worker was +ready after 19 s instead of 3 s; its first request took 21–36 ms in 21 of 23 fresh starts and 327 and +409 ms in the other two, not yet explained (132–143 ms from plain laya-serve). + +### What the warm numbers leave out + +The latencies above are for requests sent back to back. Measured on the M1 Pro: + +- **Idle gaps.** A request that follows a pause is slower, with or without the options, because the GPU + has slowed down in the meantime. For a short one-question request (25 ms back to back with the options, + 40 ms without) it took about 50 ms after 0.2–1 s of idle and 105–115 ms after 2–5 s (60–68 ms and + 114–127 ms without the options). This is also why the first request after ready costs more than a warm + one. An agent that asks once every few seconds sees these numbers, not the back-to-back ones. A + heartbeat of one forward pass every 0.5 s held it at about 45 ms in a probe, for 8% GPU load; the worker + does not do this. +- **New input lengths.** The first request of a length the worker has not seen costs about 15 ms more + once with the options (6 ms without). It is not a recompile (`recompiled_after_ready` stays `false`). +- **Memory grows with the lengths seen.** With `--compile`, PyTorch keeps host memory for every input + length the compiled model has run, about 5 MB each (fp16 weights alone add little): the 3 GB above became + 3.3 GB after 100 new lengths and 5.3 GB after all 477, more than the 4.0 GB of a worker without the + options. `torch.mps.empty_cache()` releases most of it, and those lengths then pay their first-request + cost again. + +`/health` reports under `compile` how many graphs existed when the worker became ready and how many +exist now; `recompiled_after_ready: true` means a request shape was not covered by the warmup. +`active` is `false` once no model runs the compiled path any more, i.e. after a fallback to the CPU. + +Both options apply on the GPU only. After a fallback to the CPU the worker runs Laya's fp32 model +uncompiled, like a worker started without them. + +## Start the frontend + +In another terminal: + +```sh +cargo build --release --locked +OMNI_JEV_BIND=127.0.0.1:8080 \ +OMNI_JEV_BACKEND_URL=http://127.0.0.1:8000 \ + ./target/release/omni-jev +``` + +## Send a request + +```sh +curl http://127.0.0.1:8080/v1/systemone \ + -H 'Content-Type: application/json' \ + -d '{"model":"english","state":"Please refund the duplicate charge.","questions":{"refund":{"type":"noul","instructions":"Does the customer ask for a refund?"}}}' +``` + +The frontend forwards the worker's response unchanged; `compare_with_backend.py` from the +[Laya text worker](README.md#compare-responses) recipe checks that against this setup as well. + +## Test + +The tests need `pytest` and `httpx2` (Starlette's `TestClient`; `httpx` works with a deprecation warning): + +```sh +.venv/bin/python -m pip install pytest httpx2 +PYTHONPATH=src .venv/bin/python -m pytest tests/laya # unit tests, no model +LAYA_CONTRACT=1 PYTHONPATH=src .venv/bin/python -m pytest tests/laya # plus contract tests against a CPU worker +``` + +The contract tests start a real worker and check readiness, the three decision types, error responses, and +that its answers match Laya run directly in fp32 on the CPU. On an Apple Silicon Mac, run them against the +GPU as well, without and with the options: + +```sh +LAYA_CONTRACT=1 LAYA_CONTRACT_DEVICE=mps PYTHONPATH=src .venv/bin/python -m pytest tests/laya/test_contract.py +LAYA_CONTRACT=1 LAYA_CONTRACT_DEVICE=mps LAYA_CONTRACT_FLAGS="--compile --weights fp16" \ + PYTHONPATH=src .venv/bin/python -m pytest tests/laya/test_contract.py +``` + +## Benchmark + +Stop the worker and frontend first; the benchmark starts its own. The scripts are listed in +[`bench/`](bench/README.md). A first pass that checks everything runs: + +```sh +.venv/bin/python recipe/laya/bench/bench_inproc.py --device mps --config C2 --run feasibility +.venv/bin/python recipe/laya/bench/bench_http.py --config C3 --run feasibility --spawn .venv/bin/laya-serve +.venv/bin/python recipe/laya/bench/bench_http.py --config C4 --run feasibility \ + --url http://127.0.0.1:8080 --frontend target/release/omni-jev --spawn .venv/bin/laya-serve +.venv/bin/python recipe/laya/bench/paired.py --run feasibility --a "" --b "--compile --weights fp16" +.venv/bin/python recipe/laya/bench/report.py recipe/laya/bench/results/*_feasibility.jsonl --ref C2 +.venv/bin/python recipe/laya/bench/paired.py --summarize recipe/laya/bench/results/paired_feasibility.jsonl +``` + +Runs labelled anything other than `feasibility` refuse to start on battery power or when the +1-minute load average is above 2, so close other heavy applications and plug the Mac in first. + +## Troubleshooting + +- `device_mismatch: true` at startup, or with `--require-device` the worker exits with + `asked for mps, english is on cpu`: MPS is not available to this Python. Check the `torch.backends.mps.is_available()` line above (an x86_64 + Python under Rosetta, for example, has no MPS). +- `device_mismatch: true` on a worker that started on MPS: Laya fell back to the CPU after a GPU + out-of-memory error. It keeps answering, several times slower; free memory and restart the worker + to get back on the GPU. +- The worker process uses about 4 GB, or 3 GB with fp16 weights (Activity Monitor's Memory column, which + counts MPS allocations), with one checkpoint loaded; with `--compile` it grows towards 5 GB as it sees + more input lengths (see "What the warm numbers leave out"); a second one Laya loads later adds its own. On a 16 GB Mac, close other large applications before benchmarking. +- `Address already in use`: another worker or frontend still holds port 8000 or 8080. diff --git a/recipe/laya/bench/README.md b/recipe/laya/bench/README.md new file mode 100644 index 0000000..4f8d818 --- /dev/null +++ b/recipe/laya/bench/README.md @@ -0,0 +1,60 @@ +# Laya benchmark scripts + +Scripts behind the numbers in the [Apple Silicon recipe](../apple-silicon.md). Each run writes raw +JSONL to `results/` (kept out of the repository); `report.py` builds the tables from it. + +| file | purpose | +| --- | --- | +| `workloads.jsonl` | the fixed inputs: W1–W6 are timed, P* are for answer comparison only. Tokens per row with Laya's tokenizer: W1 68, W2 198, W3 484, W4 68/48/47, W5 40–68, W6 47 | +| `bench_inproc.py` | Laya in-process (no HTTP): load, warmup, first request, warm latency, memory | +| `bench_http.py` | a `/v1/systemone` server, optionally started by the script and optionally behind the frontend: time to ready, first request, warm latency, throughput | +| `paired.py` | two configurations compared request by request, both alive at once: two worker flag sets, or two running servers (e.g. a worker directly and through the frontend) | +| `profile_mps.py` | where a request's time goes on MPS | +| `report.py` | tables from the JSONL, including the run-to-run gate and the answer comparison against a reference config | +| `env.py` | shared: versions, checkpoint, hardware, power and load recorded with each run; memory footprint | + +## Run + +From the repository root, in the recipe's environment (`.venv`), with the frontend built: + +```sh +python recipe/laya/bench/bench_inproc.py --device cpu --config C1 --run m1 +python recipe/laya/bench/bench_inproc.py --device mps --config C2 --run m1 +python recipe/laya/bench/bench_http.py --config C3 --run m1 --spawn .venv/bin/laya-serve +python recipe/laya/bench/bench_http.py --config C4 --run m1 --url http://127.0.0.1:8080 \ + --frontend target/release/omni-jev --spawn .venv/bin/laya-serve +python recipe/laya/bench/bench_http.py --config C3w --run m1 \ + --spawn .venv/bin/python -m frontend.laya_mps --device {device} --model {model} --port {port} +python recipe/laya/bench/bench_http.py --config C3o --run m1 \ + --spawn .venv/bin/python -m frontend.laya_mps --device {device} --model {model} --compile --weights fp16 --port {port} +python recipe/laya/bench/report.py recipe/laya/bench/results/*_m[0-9].jsonl --ref C1 +``` + +Repeat with `--run m2` for a second measured run. Runs refuse to start on battery power or above a +1-minute load average of `--max-load` (default 2) unless labelled `--run feasibility`. Memory is the +process's physical footprint, which on Apple Silicon includes MPS allocations. + +Two configurations compared request by request, which holds up under background load better than +separate runs: + +```sh +python recipe/laya/bench/paired.py --run p1 --a "" --b "--compile --weights fp16" +python recipe/laya/bench/paired.py --run f1 --a-url http://127.0.0.1:8000 --b-url http://127.0.0.1:8080 +python recipe/laya/bench/paired.py --summarize recipe/laya/bench/results/paired_p1.jsonl +``` + +## Results + +The measured runs on an M1 Pro are published as assets of one release on the fork, +: + +| asset | contents | sha256 | +| --- | --- | --- | +| `laya-mps-reports-2026-10-01.tar.gz` | the tables: baseline report and parity, frontend overhead, paired fp16, paired all optimizations | `e857da5082da4983e07a91104b20f84fb1ffd7994c56123e0a832e0d7a870cea` | +| `laya-mps-results-2026-09-28.tar.gz` | raw JSONL of the baseline runs (C1–C4, C3w, C3s) | `611ed30707ac8c98875b5aa5382360b5a7d760da166d61c626eb07ebe1ee6404` | +| `laya-mps-paired-fp16-2026-09-30.tar.gz` | raw JSONL of the paired fp16 runs | `cc0d6f5bda6e3e0ee1f40c6966f84429e902a2f585b8e1ee33658a9be139326e` | +| `laya-mps-paired-all-2026-09-30.tar.gz` | raw JSONL of the paired all-optimizations runs | `25d1b7bc9d6dff7173f9b972ebd0204f8fb4e2089fe2d27924d48ba6e589780c` | + +Extract the raw JSONL into `results/` and run `report.py` or `paired.py --summarize` on it to rebuild +the tables. `C3s` in the baseline runs is an earlier compile mode that compiled one-question requests +only; `--compile` does the same for them and adds the encoder for several questions. diff --git a/recipe/laya/bench/bench_http.py b/recipe/laya/bench/bench_http.py new file mode 100644 index 0000000..9626157 --- /dev/null +++ b/recipe/laya/bench/bench_http.py @@ -0,0 +1,333 @@ +"""HTTP benchmark against a /v1/systemone worker (C3 = laya-serve, C4 = laya-serve behind the Rust frontend). + +With --spawn the script starts the worker itself and times process start to the first successful +/health, then the first request per workload after that. Without it, it attaches to --url. With +--frontend as well, the worker listens on --backend-port and the Rust frontend binary is started on +--url in front of it; readiness is then the frontend's /health, which proxies the worker's. Each +workload runs at every --concurrency level; each client thread keeps one keep-alive connection. + + python recipe/laya/bench/bench_http.py --config C3 --run m1 --spawn .venv/bin/laya-serve + python recipe/laya/bench/bench_http.py --config C4 --run m1 --url http://127.0.0.1:8080 \ + --frontend target/release/omni-jev --spawn .venv/bin/laya-serve + python recipe/laya/bench/bench_http.py --config C3o --run m1 \ + --spawn .venv/bin/python -m frontend.laya_mps --device {device} --model {model} \ + --compile --weights fp16 --port {port} + +The spawned command gets LAYA_HOST/LAYA_PORT/LAYA_DEVICE/LAYA_MODELS in its environment (what laya-serve +reads) and PYTHONPATH=src (for `-m frontend.laya_mps`). `{port}`, `{device}` and `{model}` in its arguments +are replaced by the port and by --device and --model, which is how `frontend.laya_mps` gets them: it takes +flags and does not read those variables. +""" + +import argparse +import http.client +import json +import os +import random +import subprocess +import sys +import threading +import time +from pathlib import Path +from urllib.parse import urlsplit + +HERE = Path(__file__).resolve().parent +REPO = HERE.parents[2] +sys.path.insert(0, str(HERE)) +from env import footprint_mb, header, noise_problems # noqa: E402 + +CHECKPOINT = "convaiinnovations/laya" # what laya-serve's "english" model resolves to (laya/router.py) + + +class Client: + """One keep-alive connection. Not thread-safe: one per thread.""" + + def __init__(self, url, token=None): + parts = urlsplit(url) + self.conn = http.client.HTTPConnection(parts.hostname, parts.port or 80, timeout=120) + self.headers = {"Content-Type": "application/json"} + if token: + self.headers["Authorization"] = f"Bearer {token}" + + def request(self, method, path, body=None, retry=False): + """Timed calls pass retry=False so a dropped connection shows up as an error, not a slow request. + Untimed calls retry once: uvicorn closes keep-alive connections idle for 5 s. Any failure closes + the connection, so the next call reconnects instead of failing on a half-finished exchange.""" + try: + started = time.perf_counter() + self.conn.request(method, path, body=body, headers=self.headers) + response = self.conn.getresponse() + data = response.read() + return (time.perf_counter() - started) * 1000, response.status, data + except (http.client.RemoteDisconnected, BrokenPipeError, ConnectionResetError): + self.conn.close() + if not retry: + raise + return self.request(method, path, body) + except BaseException: + self.conn.close() + raise + + def close(self): + self.conn.close() + + +def wait_ready(url, processes, timeout_s): + """Poll /health every 50 ms; return seconds from now until it answers 200, and the body.""" + started = time.perf_counter() + while time.perf_counter() - started < timeout_s: + for name, process in processes.items(): + if process.poll() is not None: + sys.exit(f"{name} exited with {process.returncode} before it was ready") + client = Client(url) + try: + _, status, body = client.request("GET", "/health") + if status == 200: + return time.perf_counter() - started, json.loads(body) + except (OSError, http.client.HTTPException): + pass + finally: + client.close() + time.sleep(0.05) + sys.exit(f"worker not ready after {timeout_s} s") + + +def fetch_answers(client, body): + """(answers, None) or (None, error record fields) for one untimed request.""" + try: + _, status, data = client.request("POST", "/v1/systemone", body, retry=True) + except (OSError, http.client.HTTPException) as exc: + return None, {"status": 0, "detail": repr(exc)} + if status != 200: + return None, {"status": status, "detail": data[:200].decode(errors="replace")} + try: + return json.loads(data)["answers"], None + except (ValueError, KeyError) as exc: + return None, {"status": status, "detail": f"no answers in response: {exc!r}"} + + +def body_for(workload, model): + return json.dumps({"model": model, "state": workload["state"], "questions": workload["questions"]}).encode() + + +def run_level(url, token, body, n, concurrency): + """n requests split over `concurrency` threads. Returns [(thread, ms, status)], elapsed seconds.""" + per_thread = [n // concurrency + (i < n % concurrency) for i in range(concurrency)] + results, lock = [], threading.Lock() + barrier = threading.Barrier(concurrency + 1, timeout=120) # a thread that fails to connect breaks it + + def worker(index, count): + client = Client(url, token) + client.request("POST", "/v1/systemone", body, retry=True) # connect outside the timed window + barrier.wait() + mine = [] + for _ in range(count): + try: + ms, status, _ = client.request("POST", "/v1/systemone", body) + except (OSError, http.client.HTTPException): + ms, status = None, 0 # counted as an error, excluded from latency + mine.append((index, ms, status)) + client.close() + with lock: + results.extend(mine) + + threads = [threading.Thread(target=worker, args=(i, c)) for i, c in enumerate(per_thread)] + for t in threads: + t.start() + barrier.wait() + started = time.perf_counter() + for t in threads: + t.join() + return results, time.perf_counter() - started + + +def main(): + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("--config", required=True, help="label, e.g. C3 or C4") + parser.add_argument("--run", required=True, help="feasibility, m1, m2, ...") + parser.add_argument("--url", default="http://127.0.0.1:8000") + parser.add_argument("--model", default="english", help="model name the worker serves") + parser.add_argument("--frontend", help="Rust frontend binary to start on --url in front of the spawned worker") + parser.add_argument("--backend-port", type=int, default=8000, help="worker port when --frontend is used") + parser.add_argument("--spawn", nargs=argparse.REMAINDER, help="start this worker command, then benchmark it") + parser.add_argument("--device", default="mps", help="device for a spawned worker: LAYA_DEVICE and {device}") + parser.add_argument("--ready-timeout", type=float, default=600) + parser.add_argument("--workloads", default=str(HERE / "workloads.jsonl")) + parser.add_argument("--only", nargs="*", help="bench workload ids to run (default: all)") + parser.add_argument("-n", type=int, default=300, help="timed requests per workload and concurrency level") + parser.add_argument("--discard", type=int, default=20, help="warmup requests per workload") + parser.add_argument("--concurrency", type=int, nargs="+", default=[1, 4]) + parser.add_argument("--seed", type=int, default=0, help="workload order seed") + parser.add_argument("--out", default=str(HERE / "results")) + parser.add_argument("--max-load", type=float, default=2.0, help="1-min load average allowed for measured runs") + args = parser.parse_args() + token = os.environ.get("OMNI_JEV_TEST_TOKEN") + if args.frontend and not args.spawn: + parser.error("--frontend needs --spawn: the frontend is started in front of a spawned worker") + if args.frontend and urlsplit(args.url).port == args.backend_port: + parser.error("--url and --backend-port must differ when --frontend is used") + + problems = noise_problems(args.max_load) + if problems and args.run != "feasibility": + sys.exit("refusing a measured run: " + "; ".join(problems)) + for problem in problems: + print(f"warning: {problem}", file=sys.stderr) + + with open(args.workloads) as f: + workloads = [json.loads(line) for line in f if line.strip()] + bench = [w for w in workloads if w["kind"] == "bench" and (not args.only or w["id"] in args.only)] + parity = [w for w in workloads if w["kind"] == "parity"] + + Path(args.out).mkdir(parents=True, exist_ok=True) + processes = {} + + def memory(): + mem = footprint_mb(processes["worker"].pid) if "worker" in processes else {} + if "frontend" in processes: + mem["frontend_footprint_mb"] = footprint_mb(processes["frontend"].pid).get("footprint_mb") + return mem + + out = Path(args.out) / f"http_{args.config}_{args.run}.jsonl" + common = {"config": args.config, "run": args.run} + try: + if args.spawn: + port = args.backend_port if args.frontend else urlsplit(args.url).port + env = { + **os.environ, + "LAYA_HOST": "127.0.0.1", + "LAYA_PORT": str(port), + "LAYA_DEVICE": args.device, + "LAYA_MODELS": args.model, + "LAYA_PRELOAD": "1", + "LAYA_LOG_LEVEL": "warning", + } + env["PYTHONPATH"] = str(REPO / "src") + (os.pathsep + env["PYTHONPATH"] if env.get("PYTHONPATH") else "") + placeholders = {"{port}": str(port), "{device}": args.device, "{model}": args.model} + command = list(args.spawn) + for placeholder, value in placeholders.items(): + command = [arg.replace(placeholder, value) for arg in command] + spawn_log = open(Path(args.out) / f"http_{args.config}_{args.run}.worker.log", "w") # noqa: SIM115 + processes["worker"] = subprocess.Popen( + command, env=env, stdout=spawn_log, stderr=subprocess.STDOUT, cwd=REPO + ) + if args.frontend: + parts = urlsplit(args.url) + env = { + **os.environ, + "OMNI_JEV_BIND": f"{parts.hostname}:{parts.port}", + "OMNI_JEV_BACKEND_URL": f"http://127.0.0.1:{args.backend_port}", + } + frontend_log = open(Path(args.out) / f"http_{args.config}_{args.run}.frontend.log", "w") # noqa: SIM115 + processes["frontend"] = subprocess.Popen( + [args.frontend], env=env, stdout=frontend_log, stderr=subprocess.STDOUT + ) + ready_s, health = wait_ready(args.url, processes, args.ready_timeout) + with open(out, "w") as f: + + def emit(record): + f.write(json.dumps({**common, **record}) + "\n") + + # /health's device is what the worker reports; laya-serve 0.3.20 echoes LAYA_DEVICE. + emit( + header( + CHECKPOINT, + n=args.n, + discard=args.discard, + seed=args.seed, + url=args.url, + spawned=args.spawn, + frontend=args.frontend, + health=health, + device_actual=health.get("device"), + # The worker reports these in /health; laya-serve does not, so its values are assumed. + mps_amp_min_rows=health.get("mps_amp_min_rows", int(os.environ.get("LAYA_MPS_AMP_MIN_ROWS", "5"))), + amp_dtype=health.get("autocast_dtype") + or ("torch.float16" if health.get("device") == "mps" else "torch.float32"), + weights_dtype=health.get("weights_dtype") or "torch.float32", + dtype_source=None if "weights_dtype" in health else "assumed: laya 0.3.20 defaults", + ) + ) + + client = Client(args.url, token) + started = time.perf_counter() + first_ms, routing = {}, None + for w in bench: + ms, status, data = client.request("POST", "/v1/systemone", body_for(w, args.model), retry=True) + if status != 200: + sys.exit(f"{w['id']}: status {status}: {data[:200]!r}") + first_ms[w["id"]] = ms + routing = json.loads(data).get("routing") + for _ in range(args.discard - 1): + client.request("POST", "/v1/systemone", body_for(w, args.model), retry=True) + warmup_s = time.perf_counter() - started + emit( + { + "type": "phase", + "process_to_ready_s": round(ready_s, 3) if "worker" in processes else None, + "warmup_s": round(warmup_s, 3), + "first_ms": {k: round(v, 2) for k, v in first_ms.items()}, + "routing": routing, + "noise": problems, + **memory(), + } + ) + + order = bench[:] + random.Random(args.seed).shuffle(order) + for w in order: + body = body_for(w, args.model) + answers, error = fetch_answers(client, body) + if error: # the parity section of report.py reports the workload as missing + emit({"type": "answers_error", "workload": w["id"], **error}) + else: + emit({"type": "answers", "workload": w["id"], "answers": answers}) + for concurrency in args.concurrency: + results, elapsed = run_level(args.url, token, body, args.n, concurrency) + ok = sum(s == 200 for _, _, s in results) + for i, (thread, ms, status) in enumerate(results): + emit( + { + "type": "req", + "workload": w["id"], + "concurrency": concurrency, + "i": i, + "thread": thread, + "wall_ms": None if ms is None else round(ms, 3), + "status": status, + "rows": len(w["questions"]), + } + ) + emit( + { + "type": "throughput", + "workload": w["id"], + "concurrency": concurrency, + "n": len(results), + "errors": len(results) - ok, + "elapsed_s": round(elapsed, 3), + "rps": round(ok / elapsed, 2), + } + ) + + for w in parity: + answers, error = fetch_answers(client, body_for(w, args.model)) + if error: + emit({"type": "answers_error", "workload": w["id"], **error}) + else: + emit({"type": "answers", "workload": w["id"], "answers": answers}) + _, _, health_end = client.request("GET", "/health", retry=True) + client.close() + emit({"type": "end", "health": json.loads(health_end), **memory()}) + finally: + for proc in processes.values(): + proc.terminate() + for proc in processes.values(): + try: + proc.wait(timeout=30) + except subprocess.TimeoutExpired: + proc.kill() + print(out) + + +if __name__ == "__main__": + main() diff --git a/recipe/laya/bench/bench_inproc.py b/recipe/laya/bench/bench_inproc.py new file mode 100644 index 0000000..46cd764 --- /dev/null +++ b/recipe/laya/bench/bench_inproc.py @@ -0,0 +1,154 @@ +"""In-process Laya benchmark (configs C1 = CPU, C2 = MPS). + +Phases are timed separately: import, load, warmup, then warm requests. Every request is one line of +JSONL; report.py turns the file into tables. Run from the repository root or this directory: + + python recipe/laya/bench/bench_inproc.py --device mps --config C2 --run m1 +""" + +import argparse +import json +import random +import sys +import time +import warnings +from pathlib import Path + +T_START = time.perf_counter() +warnings.filterwarnings("ignore") +import laya # noqa: E402 +import torch # noqa: E402 + +T_IMPORT = time.perf_counter() - T_START + +HERE = Path(__file__).resolve().parent +sys.path.insert(0, str(HERE)) +from env import footprint_mb, header, noise_problems # noqa: E402 + + +def load_workloads(path): + with open(path) as f: + return [json.loads(line) for line in f if line.strip()] + + +def sync(device): + if device.type == "mps": + torch.mps.synchronize() + + +def timed_call(agent, workload): + started = time.perf_counter() + result = agent.system_one(workload["state"], workload["questions"]) + sync(agent.device) + return (time.perf_counter() - started) * 1000, result + + +def memory(device): + mem = footprint_mb() + if device.type == "mps": + mem["mps_driver_mb"] = round(torch.mps.driver_allocated_memory() / 2**20) + return mem + + +def main(): + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("--device", required=True, choices=["cpu", "mps"]) + parser.add_argument("--config", required=True, help="label, e.g. C1 or C2") + parser.add_argument("--run", required=True, help="feasibility, m1, m2, ...") + parser.add_argument("--checkpoint", default="convaiinnovations/laya") + parser.add_argument("--workloads", default=str(HERE / "workloads.jsonl")) + parser.add_argument("--only", nargs="*", help="bench workload ids to run (default: all)") + parser.add_argument("-n", type=int, default=300, help="timed requests per workload") + parser.add_argument("--discard", type=int, default=20, help="warmup requests per workload") + parser.add_argument("--seed", type=int, default=0, help="workload order seed") + parser.add_argument("--out", default=str(HERE / "results")) + parser.add_argument("--max-load", type=float, default=2.0, help="1-min load average allowed for measured runs") + args = parser.parse_args() + + problems = noise_problems(args.max_load) + if problems and args.run != "feasibility": + sys.exit("refusing a measured run: " + "; ".join(problems)) + for problem in problems: + print(f"warning: {problem}", file=sys.stderr) + + workloads = load_workloads(args.workloads) + bench = [w for w in workloads if w["kind"] == "bench" and (not args.only or w["id"] in args.only)] + parity = [w for w in workloads if w["kind"] == "parity"] + + out = Path(args.out) / f"inproc_{args.config}_{args.device}_{args.run}.jsonl" + out.parent.mkdir(parents=True, exist_ok=True) + common = {"config": args.config, "device_requested": args.device, "run": args.run} + + with open(out, "w") as f: + + def emit(record): + f.write(json.dumps({**common, **record}) + "\n") + + started = time.perf_counter() + agent = laya.load(args.checkpoint, device=args.device) + load_s = time.perf_counter() - started + + emit( + header( + args.checkpoint, + n=args.n, + discard=args.discard, + seed=args.seed, + device_actual=str(agent.device), + weights_dtype=str(next(agent.model.parameters()).dtype), + amp_dtype=str(agent.dtype), + mps_amp_min_rows=getattr(agent, "mps_amp_min_rows", None), + ) + ) + if agent.device.type != args.device: + print(f"warning: asked for {args.device}, laya is on {agent.device}", file=sys.stderr) + + # Warmup: the first call per workload is kept apart, it is the first-shape cost. + started = time.perf_counter() + first_ms = {} + for w in bench: + first_ms[w["id"]], _ = timed_call(agent, w) + for _ in range(args.discard - 1): + timed_call(agent, w) + warmup_s = time.perf_counter() - started + emit( + { + "type": "phase", + "import_s": round(T_IMPORT, 3), + "load_s": round(load_s, 3), + "warmup_s": round(warmup_s, 3), + "first_ms": {k: round(v, 2) for k, v in first_ms.items()}, + "noise": problems, + **memory(agent.device), + } + ) + + order = bench[:] + random.Random(args.seed).shuffle(order) + for w in order: + rows = len(w["questions"]) + for i in range(args.n): + ms, result = timed_call(agent, w) + emit( + { + "type": "req", + "workload": w["id"], + "i": i, + "wall_ms": round(ms, 3), + "tokens": result["usage"]["input_tokens"], + "rows": rows, + } + ) + if i == 0: + emit({"type": "answers", "workload": w["id"], "answers": result["answers"]}) + + for w in parity: + _, result = timed_call(agent, w) + emit({"type": "answers", "workload": w["id"], "answers": result["answers"]}) + + emit({"type": "end", "total_s": round(time.perf_counter() - T_START, 1), **memory(agent.device)}) + print(out) + + +if __name__ == "__main__": + main() diff --git a/recipe/laya/bench/env.py b/recipe/laya/bench/env.py new file mode 100644 index 0000000..edee840 --- /dev/null +++ b/recipe/laya/bench/env.py @@ -0,0 +1,159 @@ +"""Run header: everything needed to tell whether two result files are comparable. + +Standard library plus whatever the benchmark already imported; every probe degrades to None +instead of failing the run. +""" + +import importlib.metadata +import os +import platform +import subprocess +import sys +from datetime import datetime, timezone +from pathlib import Path + +REPO = Path(__file__).resolve().parents[3] + + +def _run(*cmd): + try: + return subprocess.run(cmd, capture_output=True, text=True, timeout=10, check=True).stdout.strip() + except (OSError, subprocess.SubprocessError): + return None + + +def _version(package): + try: + return importlib.metadata.version(package) + except importlib.metadata.PackageNotFoundError: + return None + + +def _checkpoint_revision(repo_id, ref="main"): + """The commit the cached ref points at. Read directly: laya fetches only the files it needs, so + the snapshot is partial and snapshot_download(local_files_only=True) refuses it.""" + try: + from huggingface_hub.constants import HF_HUB_CACHE + + ref_file = Path(HF_HUB_CACHE) / f"models--{repo_id.replace('/', '--')}" / "refs" / ref + return ref_file.read_text().strip() + except (ImportError, OSError): + return None + + +def _power_source(): + out = _run("pmset", "-g", "batt") + if not out: + return None + first = out.splitlines()[0] + return first.split("'")[1] if "'" in first else first + + +def _gpu_cores(): + out = _run("system_profiler", "SPDisplaysDataType") + for line in (out or "").splitlines(): + if "Total Number of Cores" in line: + return int(line.split(":")[1]) + return None + + +def noise_problems(max_load): + """Reasons this machine is not fit for a measured run, empty when it is.""" + problems = [] + power = _power_source() + if power and power != "AC Power": + problems.append(f"on {power}") + load = os.getloadavg()[0] + if load > max_load: + top = _run("ps", "-Ao", "pcpu=,comm=", "-r") or "" + busiest = "; ".join(" ".join(line.split()[:1] + [line.split("/")[-1]]) for line in top.splitlines()[:3]) + problems.append(f"1-min load {load:.1f} > {max_load} (busiest: {busiest})") + return problems + + +def header(checkpoint, **extra): + status = _run("git", "-C", str(REPO), "status", "--porcelain", "--", ".", ":!recipe/laya/bench/results") + return { + "type": "env", + "utc": datetime.now(timezone.utc).isoformat(timespec="seconds"), + "omni_sha": _run("git", "-C", str(REPO), "rev-parse", "HEAD"), + "omni_dirty": bool(status), + "checkpoint": checkpoint, + "checkpoint_revision": _checkpoint_revision(checkpoint), + "laya": _version("laya"), + "torch": _version("torch"), + "transformers": _version("transformers"), + "python": platform.python_version(), + "os": f"macOS {platform.mac_ver()[0]}" if sys.platform == "darwin" else platform.platform(), + "chip": _run("sysctl", "-n", "machdep.cpu.brand_string"), + "cpu_perf_cores": _run("sysctl", "-n", "hw.perflevel0.physicalcpu"), + "cpu_eff_cores": _run("sysctl", "-n", "hw.perflevel1.physicalcpu"), + "gpu_cores": _gpu_cores(), + "mem_gb": round(int(_run("sysctl", "-n", "hw.memsize") or 0) / 2**30), + "power": _power_source(), + "loadavg_1m": round(os.getloadavg()[0], 2), + "argv": sys.argv, + **extra, + } + + +def footprint_mb(pid=None): + """Physical footprint of a process, now and its lifetime peak (MB), from proc_pid_rusage. + + This is Activity Monitor's "Memory" column. On Apple silicon it includes Metal allocations, so it + covers MPS tensors that RSS misses, and it reads the same way for this process and a worker's pid. + """ + import ctypes + + class RusageInfoV4(ctypes.Structure): # , rusage_info_v4 + _fields_ = [("ri_uuid", ctypes.c_uint8 * 16)] + [ + (name, ctypes.c_uint64) + for name in [ + "user_time", + "system_time", + "pkg_idle_wkups", + "interrupt_wkups", + "pageins", + "wired_size", + "resident_size", + "phys_footprint", + "proc_start_abstime", + "proc_exit_abstime", + "child_user_time", + "child_system_time", + "child_pkg_idle_wkups", + "child_interrupt_wkups", + "child_pageins", + "child_elapsed_abstime", + "diskio_bytesread", + "diskio_byteswritten", + "cpu_time_qos_default", + "cpu_time_qos_maintenance", + "cpu_time_qos_background", + "cpu_time_qos_utility", + "cpu_time_qos_legacy", + "cpu_time_qos_user_initiated", + "cpu_time_qos_user_interactive", + "billed_system_time", + "serviced_system_time", + "logical_writes", + "lifetime_max_phys_footprint", + "instructions", + "cycles", + "billed_energy", + "serviced_energy", + "interval_max_phys_footprint", + "runnable_time", + ] + ] + + if sys.platform != "darwin": + return {} + info = RusageInfoV4() + libc = ctypes.CDLL("/usr/lib/libSystem.B.dylib", use_errno=True) + if libc.proc_pid_rusage(pid or os.getpid(), 4, ctypes.byref(info)) != 0: # RUSAGE_INFO_V4 + return {} + return { + "footprint_mb": round(info.phys_footprint / 2**20), + "footprint_peak_mb": round(info.lifetime_max_phys_footprint / 2**20), + } diff --git a/recipe/laya/bench/paired.py b/recipe/laya/bench/paired.py new file mode 100644 index 0000000..ecc077b --- /dev/null +++ b/recipe/laya/bench/paired.py @@ -0,0 +1,259 @@ +"""Paired comparison of two worker configurations: both run at once, and every request goes to A and to B +back to back, alternating which goes first, so background load that shifts both cancels out. + + python recipe/laya/bench/paired.py --run p1 --a "" --b "--compile --weights fp16" + python recipe/laya/bench/paired.py --run f1 --a-url http://127.0.0.1:8000 --b-url http://127.0.0.1:8080 + python recipe/laya/bench/paired.py --summarize recipe/laya/bench/results/paired_p1.jsonl + +`--a`/`--b` are extra flags for `frontend.laya_mps`, which the script starts on MPS with the english +model; `--a-url`/`--b-url` compare two servers that are already running (e.g. a worker directly and +through the Rust frontend). The summary reports, per input, the median of the per-pair ratio B/A with a +95% bootstrap interval, and B's answers against A's. +""" + +import argparse +import http.client +import json +import os +import random +import shlex +import statistics +import subprocess +import sys +from pathlib import Path + +HERE = Path(__file__).resolve().parent +REPO = HERE.parents[2] +sys.path.insert(0, str(HERE)) +from bench_http import Client, body_for, fetch_answers, wait_ready # noqa: E402 +from env import footprint_mb, header, noise_problems # noqa: E402 + +CHECKPOINT = "convaiinnovations/laya" + + +def spawn(flags, port, python, model, log_path): + env = {**os.environ, "PYTHONPATH": str(REPO / "src")} + command = [ + python, + "-m", + "frontend.laya_mps", + "--device", + "mps", + "--model", + model, + "--port", + str(port), + "--log-level", + "warning", + *shlex.split(flags), + ] + log = open(log_path, "w") # noqa: SIM115 + return subprocess.Popen(command, env=env, stdout=log, stderr=subprocess.STDOUT, cwd=REPO) + + +def run(args): + problems = noise_problems(args.max_load) + if problems and args.run != "feasibility": + sys.exit("refusing a measured run: " + "; ".join(problems)) + with open(args.workloads) as f: + workloads = [json.loads(line) for line in f if line.strip()] + bench = [w for w in workloads if w["kind"] == "bench"] + parity = [w for w in workloads if w["kind"] == "parity"] + out_dir = Path(args.out) + out_dir.mkdir(parents=True, exist_ok=True) + if args.a_url or args.b_url: + if not (args.a_url and args.b_url) or args.a is not None or args.b is not None: + sys.exit("give both --a-url and --b-url, without --a/--b") + sides = {"A": args.a_url, "B": args.b_url} + procs = {} + urls = sides + else: + sides = {"A": (args.a or "", args.port_a), "B": (args.b or "", args.port_b)} + procs = {} + urls = {s: f"http://127.0.0.1:{port}" for s, (_, port) in sides.items()} + try: + if not (args.a_url or args.b_url): + for s, (flags, port) in sides.items(): + procs[s] = spawn(flags, port, args.python, args.model, out_dir / f"paired_{args.run}_{s}.log") + health = {s: wait_ready(urls[s], {s: procs[s]} if s in procs else {}, args.ready_timeout)[1] for s in sides} + with open(out_dir / f"paired_{args.run}.jsonl", "w") as f: + + def emit(record): + f.write(json.dumps({"run": args.run, **record}) + "\n") + + emit( + header( + CHECKPOINT, + a=args.a_url or f"frontend.laya_mps {args.a or ''}".strip(), + b=args.b_url or f"frontend.laya_mps {args.b or ''}".strip(), + n=args.n, + discard=args.discard, + seed=args.seed, + health=health, + noise=problems, + ) + ) + clients = {s: Client(urls[s]) for s in sides} + for w in bench + parity: + body = body_for(w, args.model) + for s in sides: + answers, error = fetch_answers(clients[s], body) + emit({"type": "answers", "side": s, "workload": w["id"], "answers": answers, "error": error}) + order = bench[:] + random.Random(args.seed).shuffle(order) + for w in order: + body = body_for(w, args.model) + for i in range(args.discard + args.n): + first, second = ("A", "B") if i % 2 == 0 else ("B", "A") + ms = {} + for s in (first, second): + try: + t, status, _ = clients[s].request("POST", "/v1/systemone", body) + except (OSError, http.client.HTTPException): + t, status = None, 0 + ms[s] = t if status == 200 else None + if i >= args.discard: + emit( + { + "type": "pair", + "workload": w["id"], + "i": i - args.discard, + "first": first, + "a_ms": ms["A"], + "b_ms": ms["B"], + "rows": len(w["questions"]), + } + ) + end_health = {s: json.loads(clients[s].request("GET", "/health", retry=True)[2]) for s in sides} + emit( + { + "type": "end", + "health": end_health, + "footprint_mb": {s: footprint_mb(procs[s].pid).get("footprint_mb") for s in procs}, + } + ) + finally: + for p in procs.values(): + p.terminate() + for p in procs.values(): + try: + p.wait(timeout=30) + except subprocess.TimeoutExpired: + p.kill() + print(out_dir / f"paired_{args.run}.jsonl") + + +def median_interval(ratios, seed=0, resamples=2000): + rng = random.Random(seed) + meds = sorted(statistics.median(rng.choices(ratios, k=len(ratios))) for _ in range(resamples)) + return statistics.median(ratios), meds[int(0.025 * resamples)], meds[int(0.975 * resamples) - 1] + + +def flat(answer): + return answer.get("probabilities", {"p": answer.get("noul")}) + + +def decision(answer): + if answer["type"] == "choice": + return answer["choice"] + if answer["type"] == "noul": + return answer["noul"] >= 0.5 + return max(answer["probabilities"], key=answer["probabilities"].get) + + +def margin(answer): + p = sorted(flat(answer).values(), reverse=True) + return abs(p[0] - 0.5) if len(p) == 1 else p[0] - p[1] + + +def summarize(paths): + for path in paths: + with open(path) as f: + records = [json.loads(line) for line in f if line.strip()] + env = next(r for r in records if r["type"] == "env") + print(f"## {env['run']}: A = `{env['a']}`, B = `{env['b']}`, load at start {env['loadavg_1m']}\n") + print("| input | pairs | A p50 ms | B p50 ms | median B/A | 95% interval |\n|---|---|---|---|---|---|") + pairs, failed = {}, 0 + for r in records: + if r["type"] == "pair" and r["a_ms"] and r["b_ms"]: + pairs.setdefault(r["workload"], []).append(r) + elif r["type"] == "pair": + failed += 1 + for wid in sorted(pairs): + ps = pairs[wid] + med, lo, hi = median_interval([p["b_ms"] / p["a_ms"] for p in ps]) + print( + f"| {wid} | {len(ps)} | {statistics.median(p['a_ms'] for p in ps):.1f} | " + f"{statistics.median(p['b_ms'] for p in ps):.1f} | {med:.3f} | {lo:.3f}–{hi:.3f} |" + ) + answers = {} + for r in records: + if r["type"] == "answers": + answers.setdefault(r["workload"], {})[r["side"]] = r + worst, flips, errors = 0.0, [], [] + for wid, sides in answers.items(): + if any(sides.get(s, {"error": "missing"}).get("error") for s in "AB"): + errors.append(wid) + continue + a_answers, b_answers = sides["A"]["answers"], sides["B"]["answers"] + for q in sorted(a_answers.keys() | b_answers.keys()): + a, b = a_answers.get(q), b_answers.get(q) + if a is None or b is None or a.get("type") != b.get("type"): + errors.append(f"{wid}/{q}") + continue + worst = max(worst, max(abs(flat(a)[k] - flat(b).get(k, 0.0)) for k in flat(a))) + if decision(a) != decision(b): + flips.append((wid, q, round(margin(a), 4))) + print(f"\nB vs A answers: max |Δp| {worst:.4f}, flips {flips}, errors {errors}; failed pairs: {failed}") + end = next((r for r in records if r["type"] == "end"), None) + if end is None: + print("**The run did not finish: no end record, so no device, recompile or memory check.**\n") + continue + compile_state = {s: h.get("compile", {}).get("recompiled_after_ready") for s, h in end["health"].items()} + devices = {s: h.get("device") for s, h in end["health"].items()} + off_gpu = [ + s + for s, h in end["health"].items() + if h.get("device_mismatch") + or (h.get("compile", {}).get("enabled") and not h["compile"].get("active", True)) + ] + print(f"recompiled after ready: {compile_state}; device at end: {devices}; footprint MB: {end['footprint_mb']}") + if off_gpu: + print( + f"**Side {', '.join(off_gpu)} left its device or compiled path during the run; the ratios above mix both.**" + ) + print() + + +def main(): + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("--summarize", nargs="+", metavar="JSONL") + parser.add_argument("--run") + parser.add_argument("--a", help="extra frontend.laya_mps flags for side A (default: none)") + parser.add_argument("--b", help="extra frontend.laya_mps flags for side B (default: --compile --weights fp16)") + parser.add_argument("--a-url", help="instead of starting workers: an already running server for side A") + parser.add_argument("--b-url", help="... and for side B") + parser.add_argument("--port-a", type=int, default=8000) + parser.add_argument("--port-b", type=int, default=8001) + parser.add_argument("--python", default=sys.executable) + parser.add_argument("--model", default="english") + parser.add_argument("--workloads", default=str(HERE / "workloads.jsonl")) + parser.add_argument("-n", type=int, default=300, help="timed pairs per input") + parser.add_argument("--discard", type=int, default=20) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--ready-timeout", type=float, default=900) + parser.add_argument("--max-load", type=float, default=2.0) + parser.add_argument("--out", default=str(HERE / "results")) + args = parser.parse_args() + if args.summarize: + summarize(args.summarize) + elif args.run: + if args.b is None and not args.b_url: + args.b = "--compile --weights fp16" + run(args) + else: + parser.error("give --run or --summarize") + + +if __name__ == "__main__": + main() diff --git a/recipe/laya/bench/profile_mps.py b/recipe/laya/bench/profile_mps.py new file mode 100644 index 0000000..5bf0bdd --- /dev/null +++ b/recipe/laya/bench/profile_mps.py @@ -0,0 +1,254 @@ +"""Where a Laya request's time goes on MPS. Wraps laya's stages on the loaded instance; laya +itself is not modified. + +Three measurements: + +1. sweep: one choice question, state length swept; fits wall = a + b * tokens. A large `a` relative + to a short request means fixed per-request cost (dispatch, Python, sync) dominates. +2. stages: per request, time in encode (tokenize and build sequences), collate, host dispatch of the + forward (the call returns once kernels are queued), waiting for the GPU after dispatch, copy back, + and decode; GPU execution time from MPS events. Every stage runs synchronously in order, so the + stages add up to the request; "other" is what the wrappers do not cover. +3. ops: torch.profiler CPU trace of the forward: operator calls per request and the top operators by + self CPU time, i.e. the host cost of issuing the forward. + + python recipe/laya/bench/profile_mps.py --run feasibility +""" + +import argparse +import json +import statistics +import sys +import time +import warnings +from collections import defaultdict +from pathlib import Path + +warnings.filterwarnings("ignore") +import laya # noqa: E402 +import laya.agent # noqa: E402 +import torch # noqa: E402 + +HERE = Path(__file__).resolve().parent +sys.path.insert(0, str(HERE)) +from env import header, noise_problems # noqa: E402 + +STAGES = ["encode", "collate", "dispatch", "gpu_wait", "copy_back", "decode", "other"] + + +class StageTimer: + """Times laya's request stages by wrapping them on one Agent instance.""" + + def __init__(self, agent): + self.agent = agent + self.current = None + self.use_events = agent.device.type == "mps" + self._install() + + def _add(self, stage, ms): + if self.current is not None: + self.current[stage] = self.current.get(stage, 0.0) + ms + + def _timed(self, stage, fn): + def wrapper(*args, **kwargs): + started = time.perf_counter() + try: + return fn(*args, **kwargs) + finally: + self._add(stage, (time.perf_counter() - started) * 1000) + + return wrapper + + def _install(self): + agent = self.agent + agent._encode_state = self._timed("encode", agent._encode_state) + agent._decode_answers = self._timed("decode", agent._decode_answers) + laya.agent.collate_items = self._timed("collate", laya.agent.collate_items) + infer = agent._infer + + def forward(b): + # Replaces Agent._forward: same result, with dispatch, GPU wait and copy back split apart. + start_event = end_event = None + if self.use_events: + start_event = torch.mps.Event(enable_timing=True) + end_event = torch.mps.Event(enable_timing=True) + start_event.record() + t0 = time.perf_counter() + logits, act = infer(b) + if self.use_events: + end_event.record() + t1 = time.perf_counter() + if agent.device.type == "mps": + torch.mps.synchronize() + t2 = time.perf_counter() + out = logits.float().cpu().numpy(), torch.softmax(act.float(), -1).cpu().numpy() + t3 = time.perf_counter() + self._add("dispatch", (t1 - t0) * 1000) + self._add("gpu_wait", (t2 - t1) * 1000) + self._add("copy_back", (t3 - t2) * 1000) + if self.use_events: + self._add("gpu_exec", start_event.elapsed_time(end_event)) + return out + + agent._forward = forward + + def request(self, state, questions): + self.current = {} + started = time.perf_counter() + result = self.agent.system_one(state, questions) + if self.agent.device.type == "mps": + torch.mps.synchronize() + stages, self.current = self.current, None + stages["wall"] = (time.perf_counter() - started) * 1000 + stages["other"] = stages["wall"] - sum(stages.get(s, 0.0) for s in STAGES if s != "other") + return stages, result + + +def fit(xs, ys): + """Least squares y = a + b x, with R^2.""" + mx, my = statistics.fmean(xs), statistics.fmean(ys) + sxx = sum((x - mx) ** 2 for x in xs) + b = sum((x - mx) * (y - my) for x, y in zip(xs, ys)) / sxx + a = my - b * mx + ss_res = sum((y - (a + b * x)) ** 2 for x, y in zip(xs, ys)) + ss_tot = sum((y - my) ** 2 for y in ys) + return a, b, 1 - ss_res / ss_tot if ss_tot else 1.0 + + +def main(): + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("--run", required=True) + parser.add_argument("--device", default="mps", choices=["mps", "cpu"]) + parser.add_argument("--checkpoint", default="convaiinnovations/laya") + parser.add_argument("--workloads", default=str(HERE / "workloads.jsonl")) + parser.add_argument("--stage-workloads", nargs="+", default=["W1", "W3", "W5"]) + parser.add_argument("-n", type=int, default=100, help="timed requests per point") + parser.add_argument("--discard", type=int, default=10) + parser.add_argument("--sweep-words", type=int, nargs="+", default=[1, 8, 24, 56, 120, 250, 380]) + parser.add_argument("--ops-requests", type=int, default=20) + parser.add_argument("--out", default=str(HERE / "results")) + parser.add_argument("--max-load", type=float, default=2.0) + args = parser.parse_args() + + problems = noise_problems(args.max_load) + if problems and args.run != "feasibility": + sys.exit("refusing a measured run: " + "; ".join(problems)) + for problem in problems: + print(f"warning: {problem}", file=sys.stderr) + + with open(args.workloads) as f: + workloads = {w["id"]: w for w in (json.loads(line) for line in f if line.strip())} + route = workloads["W1"]["questions"] + + agent = laya.load(args.checkpoint, device=args.device) + timer = StageTimer(agent) + out = Path(args.out) / f"profile_{args.device}_{args.run}.jsonl" + out.parent.mkdir(parents=True, exist_ok=True) + common = {"config": f"profile-{args.device}", "run": args.run} + + with open(out, "w") as f: + + def emit(record): + f.write(json.dumps({**common, **record}) + "\n") + + emit( + header( + args.checkpoint, + device_actual=str(agent.device), + noise=problems, + weights_dtype=str(next(agent.model.parameters()).dtype), + amp_dtype=str(agent.dtype), + mps_amp_min_rows=getattr(agent, "mps_amp_min_rows", None), + ) + ) + + print("## Length sweep (1 choice question, 5 options)\n") + print("| words | tokens | p50 ms |\n|---|---|---|") + points = [] + for words in args.sweep_words: + state = " ".join(["delivery"] * words) + for _ in range(args.discard): + timer.request(state, route) + walls, tokens = [], None + for _ in range(args.n): + stages, result = timer.request(state, route) + walls.append(stages["wall"]) + tokens = result["usage"]["input_tokens"] + p50 = statistics.median(walls) + points.append((tokens, p50)) + emit( + { + "type": "sweep", + "words": words, + "tokens": tokens, + "p50_ms": round(p50, 3), + "wall_ms": [round(w, 3) for w in walls], + } + ) + print(f"| {words} | {tokens} | {p50:.1f} |") + a, b, r2 = fit([t for t, _ in points], [p for _, p in points]) + emit({"type": "fit", "a_ms": round(a, 3), "b_ms_per_token": round(b, 5), "r2": round(r2, 4)}) + print(f"\nwall ≈ {a:.1f} ms + {b:.3f} ms/token × tokens (R² {r2:.3f})") + + print("\n## Stages (median ms per request)\n") + columns = ["wall", *STAGES, "gpu_exec"] + print("| workload | tokens | rows | " + " | ".join(columns) + " |\n|" + "---|" * (len(columns) + 3)) + for wid in args.stage_workloads: + w = workloads[wid] + for _ in range(args.discard): + timer.request(w["state"], w["questions"]) + per_stage, tokens = defaultdict(list), None + for _ in range(args.n): + stages, result = timer.request(w["state"], w["questions"]) + tokens = result["usage"]["input_tokens"] + for k, v in stages.items(): + per_stage[k].append(v) + medians = {k: statistics.median(v) for k, v in per_stage.items()} + emit( + { + "type": "stages", + "workload": wid, + "tokens": tokens, + "rows": len(w["questions"]), + "median_ms": {k: round(v, 3) for k, v in medians.items()}, + "samples": {k: [round(x, 3) for x in v] for k, v in per_stage.items()}, + } + ) + cells = " | ".join(f"{medians[c]:.1f}" if c in medians else "" for c in columns) + print(f"| {wid} | {tokens} | {len(w['questions'])} | {cells} |") + + print("\n## Host operators for W1 (torch.profiler, CPU)\n") + w = workloads["W1"] + with torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CPU]) as prof: + for _ in range(args.ops_requests): + timer.request(w["state"], w["questions"]) + events = [e for e in prof.key_averages() if e.key.startswith("aten::")] + calls = sum(e.count for e in events) / args.ops_requests + top = sorted(events, key=lambda e: e.self_cpu_time_total, reverse=True)[:12] + emit( + { + "type": "ops", + "workload": "W1", + "requests": args.ops_requests, + "aten_calls_per_request": calls, + "top": [ + { + "op": e.key, + "calls_per_request": e.count / args.ops_requests, + "self_cpu_ms_per_request": e.self_cpu_time_total / 1000 / args.ops_requests, + } + for e in top + ], + } + ) + print(f"aten calls per request: {calls:.0f}\n") + print("| op | calls/request | self CPU ms/request |\n|---|---|---|") + for e in top: + print( + f"| {e.key} | {e.count / args.ops_requests:.0f} | {e.self_cpu_time_total / 1000 / args.ops_requests:.2f} |" + ) + print(f"\n{out}") + + +if __name__ == "__main__": + main() diff --git a/recipe/laya/bench/report.py b/recipe/laya/bench/report.py new file mode 100644 index 0000000..97a4b42 --- /dev/null +++ b/recipe/laya/bench/report.py @@ -0,0 +1,295 @@ +"""Turn benchmark JSONL into markdown tables. The only place numbers are computed from raw data. + + python recipe/laya/bench/report.py recipe/laya/bench/results/*.jsonl + +The parity section compares every run's answers with the reference config's (`--ref`, default C1): the +decision must match (choice: option; score: most likely level; noul: side of 0.5) and the largest |Δp| +must stay within 1e-3 for fp32 and 1e-2 where the run used fp16 (autocast or fp16 weights). +Percentiles are nearest-rank. The run-to-run gate compares p50 across measured runs (every run whose +label is not "feasibility") of the same config, workload and concurrency: (max - min) / min <= 10%. +In-process results have concurrency 1. +""" + +import argparse +import json +import math +import statistics +import sys +from collections import defaultdict + +GATE = 0.10 +PHASES = ["import_s", "load_s", "process_to_ready_s", "warmup_s"] # whichever a result file has + + +def percentile(sorted_values, p): + return sorted_values[max(0, math.ceil(p * len(sorted_values)) - 1)] + + +def read(paths): + records = [] + for path in paths: + with open(path) as f: + rows = [json.loads(line) for line in f if line.strip()] + if any("config" not in r for r in rows): # e.g. paired.py results: summarize those with paired.py + print(f"skipping {path}: not a bench_inproc/bench_http result", file=sys.stderr) + continue + records.extend(rows) + return records + + +def table(headers, rows): + lines = ["| " + " | ".join(headers) + " |", "|" + "---|" * len(headers)] + lines += ["| " + " | ".join("" if c is None else str(c) for c in row) + " |" for row in rows] + return "\n".join(lines) + + +def short(sha, dirty=False): + return (sha or "?")[:7] + ("+dirty" if dirty else "") + + +def environment(envs): + rows = [] + for k, e in sorted(envs.items()): + dtype = f"{e['weights_dtype']} (amp {e['amp_dtype']}, >= {e['mps_amp_min_rows']} rows)" + if e.get("dtype_source"): + dtype += f" [{e['dtype_source']}]" + rows.append( + [ + *k, + e["device_actual"], + dtype, + e["chip"], + e["os"], + e["power"], + e["loadavg_1m"], + f"laya {e['laya']} / torch {e['torch']}", + short(e["checkpoint_revision"]), + short(e["omni_sha"], e["omni_dirty"]), + ] + ) + headers = [ + "config", + "run", + "device", + "weights (autocast)", + "chip", + "os", + "power", + "load 1m", + "versions", + "ckpt", + "omni", + ] + return table(headers, rows) + + +def phase_table(phases): + present = [p for p in PHASES if any(ph.get(p) is not None for ph in phases.values())] + workload_ids = sorted({w for ph in phases.values() for w in ph["first_ms"]}) + rows = [ + [*k, *(ph.get(p) for p in present), *(ph["first_ms"].get(w) for w in workload_ids)] + for k, ph in sorted(phases.items()) + ] + headers = ["config", "run", *(p.removesuffix("_s") for p in present), *(f"first {w}" for w in workload_ids)] + return table(headers, rows) + + +def latency(records): + samples = defaultdict(list) + for r in records: + if r["type"] == "req" and r.get("status", 200) == 200: + samples[(r["config"], r["workload"], r.get("concurrency", 1), r["run"])].append(r["wall_ms"]) + rows, p50s = [], defaultdict(dict) + for (config, workload, concurrency, run), values in sorted(samples.items()): + values.sort() + mean = statistics.fmean(values) + cv = statistics.stdev(values) / mean if len(values) > 1 else 0.0 + p50 = percentile(values, 0.50) + if run != "feasibility": + p50s[(config, workload, concurrency)][run] = p50 + rows.append( + [ + config, + workload, + concurrency, + run, + len(values), + round(p50, 2), + round(percentile(values, 0.95), 2), + round(mean, 2), + f"{cv:.1%}", + ] + ) + return table(["config", "workload", "conc", "run", "n", "p50", "p95", "mean", "CV"], rows), p50s + + +def gate(p50s): + rows = [] + for (config, workload, concurrency), by_run in sorted(p50s.items()): + if len(by_run) < 2: + rows.append([config, workload, concurrency, len(by_run), "", "needs 2 measured runs"]) + continue + spread = (max(by_run.values()) - min(by_run.values())) / min(by_run.values()) + rows.append([config, workload, concurrency, len(by_run), f"{spread:.1%}", "PASS" if spread <= GATE else "FAIL"]) + return table(["config", "workload", "conc", "runs", "spread", "gate"], rows) + + +def throughput(records): + rows = [ + [r["config"], r["workload"], r["concurrency"], r["run"], r["n"], r["errors"], r["elapsed_s"], r["rps"]] + for r in records + if r["type"] == "throughput" + ] + return table(["config", "workload", "conc", "run", "n", "errors", "elapsed s", "successful req/s"], sorted(rows)) + + +def memory(phases, ends): + """Physical footprint (Activity Monitor's "Memory", includes Metal allocations) of the process running + Laya: the benchmark itself in-process, the worker over HTTP. Peak is the process lifetime maximum.""" + rows = [] + for k in sorted(ends): + ph, e = phases.get(k, {}), ends[k] + rows.append( + [ + *k, + ph.get("footprint_mb"), + e.get("footprint_mb"), + e.get("footprint_peak_mb"), + ph.get("mps_driver_mb"), + e.get("mps_driver_mb"), + e.get("frontend_footprint_mb"), + ] + ) + headers = [ + "config", + "run", + "footprint after warmup", + "footprint end", + "footprint peak", + "MPS driver after warmup", + "MPS driver end", + "frontend footprint", + ] + return table(headers, rows) + + +TOLERANCE = {"fp32": 1e-3, "fp16": 1e-2} + + +def read_answers(records): + """Per (config, run): the env record, the answers per workload, the probes that failed, and which runs are + benchmark runs (bench_inproc/bench_http write a phase record before any answers; profile_mps never answers).""" + envs, answers, errors, benchmarks = {}, {}, {}, set() + for r in records: + key = (r["config"], r["run"]) + if r["type"] == "env": + envs[key] = r + elif r["type"] == "phase": + benchmarks.add(key) + elif r["type"] == "answers": + answers.setdefault(key, {})[r["workload"]] = r["answers"] + elif r["type"] == "answers_error": + errors.setdefault(key, {})[r["workload"]] = f"status {r.get('status')}: {r.get('detail')}" + return envs, answers, errors, benchmarks + + +def outcome(answer): + """(decision, {outcome: probability}) for one question's answer.""" + kind = answer["type"] + if kind == "noul": + p = answer["noul"] + return p >= 0.5, {"yes": p} + probs = answer["probabilities"] + if kind == "choice": + return answer["choice"], probs + return max(probs, key=probs.get), probs # score: most likely level + + +def margin(probs): + if len(probs) == 1: # noul: distance from the 0.5 boundary + return abs(next(iter(probs.values())) - 0.5) + top = sorted(probs.values(), reverse=True) + return top[0] - top[1] + + +def path(env, rows): + autocast = ( + env.get("device_actual") == "mps" + and env.get("amp_dtype") == "torch.float16" + and rows >= (env.get("mps_amp_min_rows") or 10**9) + ) + return "fp16" if autocast or env.get("weights_dtype") == "torch.float16" else "fp32" + + +def parity(records, ref): + """Answers of every run against the reference config, with the tolerances declared in advance.""" + envs, answers, errors, benchmarks = read_answers(records) + refs = sorted(k for k in answers if k[0] == ref) + if not refs: + return f"no answers for reference config {ref}" + ref_key = refs[0] + reference = answers[ref_key] + lines = [ + f"Reference: {ref_key[0]} / {ref_key[1]}. Tolerances: fp32 {TOLERANCE['fp32']}, fp16 {TOLERANCE['fp16']}.", + "", + "| config | run | workload | question | type | path | decision ref → run | max abs Δp | ref margin | result |", + "|---|---|---|---|---|---|---|---|---|---|", + ] + failed = total = 0 + for key in sorted(benchmarks | answers.keys() | errors.keys()): # a benchmark run without answers still counts + if key == ref_key: + continue + env = envs.get(key, {}) + for workload, questions in sorted(reference.items()): + got = answers.get(key, {}).get(workload) + if got is None: + why = errors.get(key, {}).get(workload, "missing") + lines.append(f"| {key[0]} | {key[1]} | {workload} | | | | {why} | | | FAIL |") + failed += 1 + total += 1 + continue + precision = path(env, len(questions)) + for qid, ref_answer in sorted(questions.items()): + total += 1 + if qid not in got: + lines.append(f"| {key[0]} | {key[1]} | {workload} | {qid} | | | missing | | | FAIL |") + failed += 1 + continue + ref_decision, ref_probs = outcome(ref_answer) + decision, probs = outcome(got[qid]) + delta = max(abs(ref_probs.get(o, 0.0) - probs.get(o, 0.0)) for o in ref_probs.keys() | probs.keys()) + ok = decision == ref_decision and delta <= TOLERANCE[precision] + failed += not ok + flip = f"{ref_decision} → {decision}" if decision != ref_decision else f"{decision}" + lines.append( + f"| {key[0]} | {key[1]} | {workload} | {qid} | {ref_answer['type']} | {precision} | {flip} " + f"| {delta:.4f} | {margin(ref_probs):.4f} | {'PASS' if ok else 'FAIL'} |" + ) + lines += ["", f"{total - failed}/{total} questions within tolerance."] + return "\n".join(lines) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("files", nargs="+") + parser.add_argument("--ref", default="C1", help="reference config for the parity section (default C1)") + args = parser.parse_args() + records = read(args.files) + + def by_run(kind): + return {(r["config"], r["run"]): r for r in records if r["type"] == kind} + + envs, phases, ends = by_run("env"), by_run("phase"), by_run("end") + latency_table, p50s = latency(records) + print("## Environment\n\n" + environment(envs)) + print("\n## Phases (s) and first request per workload (ms)\n\n" + phase_table(phases)) + print("\n## Warm latency (ms)\n\n" + latency_table) + print(f"\n## Run-to-run gate (p50 spread across measured runs <= {GATE:.0%})\n\n" + gate(p50s)) + if any(r["type"] == "throughput" for r in records): + print("\n## Throughput\n\n" + throughput(records)) + print("\n## Memory (MB)\n\n" + memory(phases, ends)) + print(f"\n## Parity against {args.ref}\n\n" + parity(records, args.ref)) + + +if __name__ == "__main__": + main() diff --git a/recipe/laya/bench/results/.gitignore b/recipe/laya/bench/results/.gitignore new file mode 100644 index 0000000..2ef39dd --- /dev/null +++ b/recipe/laya/bench/results/.gitignore @@ -0,0 +1,3 @@ +# Benchmark output stays out of the repository; the reports and raw data are release assets (see ../README.md). +* +!.gitignore diff --git a/recipe/laya/bench/workloads.jsonl b/recipe/laya/bench/workloads.jsonl new file mode 100644 index 0000000..12f444b --- /dev/null +++ b/recipe/laya/bench/workloads.jsonl @@ -0,0 +1,11 @@ +{"id": "W1", "kind": "bench", "target_tokens_per_row": 68, "state": "My package never arrived and tracking has not updated in ten days.", "questions": {"route": {"type": "choice", "instructions": "Route the ticket to the queue that owns it.", "criteria": {"logistics": "Shipping, delivery and tracking", "payment": "Charges, invoices and refunds", "returns": "Returns and exchanges", "account": "Login, password and profile", "human": "Anything else"}}}} +{"id": "W2", "kind": "bench", "target_tokens_per_row": 200, "state": "Hello, I ordered a pair of running shoes three weeks ago and paid for express delivery. The confirmation email said the parcel would arrive within two business days, but the tracking page has shown 'label created' ever since. I contacted the courier and they told me they never received the parcel from your warehouse. In the meantime I was charged twice on my credit card, once for the original amount and once for a slightly different amount that I do not recognise. I need the shoes for a race next weekend, so please either ship them today with a tracking number that actually works or cancel the order and refund both charges. I have been a customer for years and this is the first time something like this has happened.", "questions": {"route": {"type": "choice", "instructions": "Route the ticket to the queue that owns it.", "criteria": {"logistics": "Shipping, delivery and tracking", "payment": "Charges, invoices and refunds", "returns": "Returns and exchanges", "account": "Login, password and profile", "human": "Anything else"}}}} +{"id": "W3", "kind": "bench", "target_tokens_per_row": 480, "state": "Hello, I ordered a pair of running shoes three weeks ago and paid for express delivery. The confirmation email said the parcel would arrive within two business days, but the tracking page has shown 'label created' ever since. I contacted the courier and they told me they never received the parcel from your warehouse. In the meantime I was charged twice on my credit card, once for the original amount and once for a slightly different amount that I do not recognise. I need the shoes for a race next weekend, so please either ship them today with a tracking number that actually works or cancel the order and refund both charges. I have been a customer for years and this is the first time something like this has happened. Hello, I ordered a pair of running shoes three weeks ago and paid for express delivery. The confirmation email said the parcel would arrive within two business days, but the tracking page has shown 'label created' ever since. I contacted the courier and they told me they never received the parcel from your warehouse. In the meantime I was charged twice on my credit card, once for the original amount and once for a slightly different amount that I do not recognise. I need the shoes for a race next weekend, so please either ship them today with a tracking number that actually works or cancel the order and refund both charges. I have been a customer for years and this is the first time something like this has happened. Hello, I ordered a pair of running shoes three weeks ago and paid for express delivery. The confirmation email said the parcel would arrive within two business days, but the tracking page has shown 'label created' ever since. I contacted the courier and they told me they never received the parcel from your warehouse. In the meantime I was charged twice on my credit card, once for the original amount and once for a slightly different amount that I do not recognise. I need the shoes for a race next weekend, so please either ship them today with a tracking number that actually works or cancel the order and refund both charges. I have been a customer for years and this is the first time something like this has happened.", "questions": {"route": {"type": "choice", "instructions": "Route the ticket to the queue that owns it.", "criteria": {"logistics": "Shipping, delivery and tracking", "payment": "Charges, invoices and refunds", "returns": "Returns and exchanges", "account": "Login, password and profile", "human": "Anything else"}}}} +{"id": "W4", "kind": "bench", "target_tokens_per_row": 54, "state": "My package never arrived and tracking has not updated in ten days.", "questions": {"route": {"type": "choice", "instructions": "Route the ticket to the queue that owns it.", "criteria": {"logistics": "Shipping, delivery and tracking", "payment": "Charges, invoices and refunds", "returns": "Returns and exchanges", "account": "Login, password and profile", "human": "Anything else"}}, "urgency": {"type": "score", "instructions": "How urgent is the ticket?", "criteria": ["Not urgent", "Needs attention soon", "Needs attention immediately"]}, "refund": {"type": "noul", "instructions": "Does the customer ask for a refund?"}}} +{"id": "W5", "kind": "bench", "target_tokens_per_row": 49, "state": "My package never arrived and tracking has not updated in ten days.", "questions": {"route": {"type": "choice", "instructions": "Route the ticket to the queue that owns it.", "criteria": {"logistics": "Shipping, delivery and tracking", "payment": "Charges, invoices and refunds", "returns": "Returns and exchanges", "account": "Login, password and profile", "human": "Anything else"}}, "urgency": {"type": "score", "instructions": "How urgent is the ticket?", "criteria": ["Not urgent", "Needs attention soon", "Needs attention immediately"]}, "refund": {"type": "noul", "instructions": "Does the customer ask for a refund?"}, "angry": {"type": "noul", "instructions": "Is the customer angry?"}, "cancel": {"type": "noul", "instructions": "Does the customer want to cancel the order?"}, "lang": {"type": "choice", "instructions": "Which language is the ticket written in?", "criteria": {"en": "English", "de": "German", "fr": "French"}}}} +{"id": "W6", "kind": "bench", "target_tokens_per_row": 47, "state": "My package never arrived and tracking has not updated in ten days.", "questions": {"refund": {"type": "noul", "instructions": "Does the customer ask for a refund?"}}} +{"id": "P2", "kind": "parity", "target_tokens_per_row": null, "state": "My package never arrived and tracking has not updated in ten days.", "questions": {"route": {"type": "choice", "instructions": "Route the ticket to the queue that owns it.", "criteria": {"logistics": "The logistics team", "payment": "The payment team"}}}} +{"id": "P5", "kind": "parity", "target_tokens_per_row": null, "state": "My package never arrived and tracking has not updated in ten days.", "questions": {"route": {"type": "choice", "instructions": "Route the ticket to the queue that owns it.", "criteria": {"logistics": "The logistics team", "payment": "The payment team", "returns": "The returns team", "account": "The account team", "human": "The human team"}}}} +{"id": "P10", "kind": "parity", "target_tokens_per_row": null, "state": "My package never arrived and tracking has not updated in ten days.", "questions": {"route": {"type": "choice", "instructions": "Route the ticket to the queue that owns it.", "criteria": {"logistics": "The logistics team", "payment": "The payment team", "returns": "The returns team", "account": "The account team", "human": "The human team", "billing": "The billing team", "technical": "The technical team", "sales": "The sales team", "legal": "The legal team", "security": "The security team"}}}} +{"id": "P2L", "kind": "parity", "target_tokens_per_row": null, "state": "Hello, I ordered a pair of running shoes three weeks ago and paid for express delivery. The confirmation email said the parcel would arrive within two business days, but the tracking page has shown 'label created' ever since. I contacted the courier and they told me they never received the parcel from your warehouse. In the meantime I was charged twice on my credit card, once for the original amount and once for a slightly different amount that I do not recognise. I need the shoes for a race next weekend, so please either ship them today with a tracking number that actually works or cancel the order and refund both charges. I have been a customer for years and this is the first time something like this has happened.", "questions": {"route": {"type": "choice", "instructions": "Route the ticket to the queue that owns it.", "criteria": {"logistics": "The logistics team", "payment": "The payment team"}}, "urgency": {"type": "score", "instructions": "How urgent is the ticket?", "criteria": ["Not urgent", "Needs attention soon", "Needs attention immediately"]}, "refund": {"type": "noul", "instructions": "Does the customer ask for a refund?"}}} +{"id": "P10L", "kind": "parity", "target_tokens_per_row": null, "state": "Hello, I ordered a pair of running shoes three weeks ago and paid for express delivery. The confirmation email said the parcel would arrive within two business days, but the tracking page has shown 'label created' ever since. I contacted the courier and they told me they never received the parcel from your warehouse. In the meantime I was charged twice on my credit card, once for the original amount and once for a slightly different amount that I do not recognise. I need the shoes for a race next weekend, so please either ship them today with a tracking number that actually works or cancel the order and refund both charges. I have been a customer for years and this is the first time something like this has happened.", "questions": {"route": {"type": "choice", "instructions": "Route the ticket to the queue that owns it.", "criteria": {"logistics": "The logistics team", "payment": "The payment team", "returns": "The returns team", "account": "The account team", "human": "The human team", "billing": "The billing team", "technical": "The technical team", "sales": "The sales team", "legal": "The legal team", "security": "The security team"}}, "urgency": {"type": "score", "instructions": "How urgent is the ticket?", "criteria": ["Not urgent", "Needs attention soon", "Needs attention immediately"]}, "refund": {"type": "noul", "instructions": "Does the customer ask for a refund?"}}} diff --git a/recipe/laya/requirements-mps.txt b/recipe/laya/requirements-mps.txt new file mode 100644 index 0000000..dca8e7a --- /dev/null +++ b/recipe/laya/requirements-mps.txt @@ -0,0 +1,10 @@ +# The environment the Apple Silicon recipe was validated in (M1 Pro and M5, Python 3.12). +# laya[serve] pulls in torch, transformers, fastapi and uvicorn; the pins below are the versions it resolved to. +laya[serve]==0.3.20 +torch==2.14.0 +transformers==5.17.0 +tokenizers==0.23.2 +safetensors==0.8.0 +huggingface-hub==1.33.0 +fastapi==0.141.1 +uvicorn==0.54.0 diff --git a/src/frontend/laya_mps.py b/src/frontend/laya_mps.py new file mode 100644 index 0000000..c29eda4 --- /dev/null +++ b/src/frontend/laya_mps.py @@ -0,0 +1,241 @@ +"""HTTP worker for Laya on Apple Silicon (PyTorch MPS) and CPU: `GET /health` and `POST /v1/systemone`. + +PYTHONPATH=src python -m frontend.laya_mps --device mps --model english [--compile] [--weights fp16] + +It is laya-serve (`laya[serve]==0.3.20`) with its request handling unchanged and three changes: + +- It binds only after every loaded model has run a warmup over short, long and multi-question requests, + so a reachable worker is a warm one. laya-serve answers /health before any forward pass. A checkpoint + laya loads later (a request that names another model, or routes to it) gets the same options and + warmup inside that first request. +- /health describes the loaded models as they are now: device, weight and autocast dtypes, checkpoint + and the revision the weights were loaded from, and `device_mismatch` when a model is not on the + requested device. laya-serve reports the configured device, and laya moves a model to the CPU on a + GPU out-of-memory error and keeps serving. `--require-device` exits at startup instead of serving + from another device. +- `--compile` and `--weights fp16` make the GPU path faster (models/laya/optimize.py). Both apply on + the GPU only; on the CPU, including after a fallback, the worker runs laya's fp32 model uncompiled. + +Of laya-serve's environment variables, the ones its app reads still apply, notably LAYA_API_KEY for bearer +authentication. The ones its launcher reads do not, because the flags above replace it: LAYA_DEVICE, +LAYA_MODELS, LAYA_PRELOAD, LAYA_HOST, LAYA_PORT, LAYA_LOG_LEVEL, LAYA_THREADS and LAYA_AUTO_TASK. The +worker warns at startup about any of those that are set. +""" + +from __future__ import annotations + +import argparse +import logging +import os +import sys +from typing import Any + +from models.laya import engine, optimize + +log = logging.getLogger("laya-worker") + +# Read by laya-serve's launcher (laya.serve.build_router and main), which this worker does not run. +UNREAD_LAYA_SERVE_VARIABLES = ( + "LAYA_DEVICE", + "LAYA_MODELS", + "LAYA_PRELOAD", + "LAYA_HOST", + "LAYA_PORT", + "LAYA_LOG_LEVEL", + "LAYA_THREADS", + "LAYA_AUTO_TASK", +) + + +class Lifecycle: + """laya Router hooks that hand checkpoint loads and evictions to one worker app.""" + + def __init__(self, on_load, on_evict): + self.on_load = on_load + self.on_evict = on_evict + + +def build_app( + router: Any, + model: str, + requested: str | None, + *, + require_device: bool = False, + compile: bool = False, + fp16: bool = False, + graph_counter=optimize.compiled_graphs, + revisions: dict[str, str] | None = None, +): + """Prepare every loaded model (options, warmup, device check), then return laya's app with /health + replaced. A checkpoint laya loads later, for a request that names or routes to it, is prepared the same + way before it answers. `model` is the one summarised at the top of /health, and the one loaded if + nothing is preloaded. Raises if preparing fails, so the caller never binds a worker that cannot answer.""" + from laya.serve import create_app + + resident: dict[str, tuple[Any, dict[str, Any]]] = {} # prepared checkpoints: name -> (agent, warmup result) + preparing: set[str] = set() + graphs_at_ready = None + + def apply_options(name: str, agent: Any) -> None: + if (fp16 or compile) and not optimize.apply(agent, fp16=fp16, compile=compile): + log.warning("%s is on the CPU: --compile and --weights fp16 apply on the GPU only", name) + + def make_ready(name: str, agent: Any) -> str | None: + """Warm the checkpoint up and start describing it. Returns where it is if not on the requested device.""" + nonlocal graphs_at_ready + warmed = engine.warmup(router, name) + warmed["revision"] = engine.loaded_revision(revisions, warmed["routing"]) # fixed here: see engine + resident[name] = (agent, warmed) + autocast_rows = getattr(agent, "mps_amp_min_rows", None) + if str(agent.device).startswith("mps") and autocast_rows and autocast_rows > engine.WARMUP_MAX_ROWS: + log.warning( + "%s: laya autocasts from %d questions but the warmup stops at %d; the first request that " + "large is not warm", + name, + autocast_rows, + engine.WARMUP_MAX_ROWS, + ) + if compile: + graphs_at_ready = graph_counter() + described = engine.describe(agent, requested, warmed["routing"], warmed["revision"]) + return f"{name} is on {described['device']}" if described["device_mismatch"] else None + + def check_device(misplaced: list[str]) -> None: + if misplaced: + message = f"asked for {requested}, {', '.join(misplaced)}" + if require_device: + raise RuntimeError(message) + log.warning(message) + + evicted: list[str] = [] # what laya dropped to make room for the checkpoint it is loading + + def on_evict(ctx: Any) -> None: + resident.pop(ctx.model, None) + evicted.append(ctx.model) + + def on_load(ctx: Any) -> None: + """A checkpoint loaded while serving is prepared like the ones loaded at startup, or not kept at all; + in that case the checkpoints laya evicted for it are loaded again.""" + nonlocal graphs_at_ready + made_room = [name for name in evicted if name != ctx.model] + evicted.clear() + preparing.add(ctx.model) + try: + apply_options(ctx.model, ctx.agent) + check_device(list(filter(None, [make_ready(ctx.model, ctx.agent)]))) + except Exception: + log.exception("%s could not be prepared and is unloaded", ctx.model) + router.unload(ctx.model) + evicted.clear() + for name in made_room: + try: + router.load(name) + except Exception: # noqa: BLE001 -- the request fails for the first reason either way + log.exception("%s was evicted for %s and could not be loaded again", name, ctx.model) + if compile: + graphs_at_ready = graph_counter() # graphs the unloaded checkpoint compiled are not recompiles + raise + finally: + preparing.discard(ctx.model) + + names = list(router.loaded) or [model] # never load a model the worker was not asked to serve + startup = {name: router.load(name) for name in names} + for name, agent in startup.items(): + apply_options(name, agent) + check_device(list(filter(None, [make_ready(name, agent) for name, agent in startup.items()]))) + for hook in [h for h in getattr(router, "hooks", ()) if isinstance(h, Lifecycle)]: + router.remove_hook(hook) # an app built earlier on this router + router.add_hook(Lifecycle(on_load, on_evict)) + + def current() -> dict[str, Any]: + """The agents as they are now, not as they were at startup (see the module docstring).""" + models = { + name: { + **engine.describe(agent, requested, warmed["routing"], warmed["revision"]), + "warmup_ms": warmed["warmup_ms"], + } + for name, (agent, warmed) in list(resident.items()) + } + return { + **models.get(model, next(iter(models.values()), {})), + "device_mismatch": any(m["device_mismatch"] for m in models.values()), + "warmup_ms": round(sum(m["warmup_ms"] for m in models.values()), 1), + "models": models, + } + + app = create_app(router) + app.router.routes[:] = [r for r in app.router.routes if getattr(r, "path", None) != "/health"] + + @app.get("/health") + def health() -> dict[str, Any]: + compiled: dict[str, Any] = {"enabled": compile} + if compile: + now = graph_counter() + compiled.update( + active=any(optimize.compile_active(agent) for agent, _ in list(resident.values())), + graphs_at_ready=graphs_at_ready, + graphs_now=now, + recompiled_after_ready=now > graphs_at_ready and not preparing, + ) + return { + "status": "ok", + "ready": True, + "loaded": router.loaded, + "preparing": sorted(preparing), + **current(), + "compile": compiled, + } + + return app + + +def make_router(device: str | None, model: str) -> Any: + from laya.router import Router + + router = Router(device=device) + router.preload([model]) + return router + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("--device", default=None, help="torch device for laya: mps, cpu (default: laya's choice)") + parser.add_argument("--model", default="english", help="laya checkpoint to serve: english, multilingual, ...") + parser.add_argument("--compile", action="store_true", help="torch.compile the GPU path during warmup") + parser.add_argument("--weights", default="fp32", choices=["fp32", "fp16"], help="weight precision on the GPU") + parser.add_argument("--require-device", action="store_true", help="exit if a model is not on --device") + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--port", type=int, default=8000) + parser.add_argument( + "--log-level", + default="info", + choices=["critical", "error", "warning", "info", "debug"], + help="for the worker's own log and uvicorn's", + ) + args = parser.parse_args() + + import uvicorn + + logging.basicConfig(level=args.log_level.upper(), format="%(name)s: %(message)s") + unread = [name for name in UNREAD_LAYA_SERVE_VARIABLES if os.environ.get(name)] + if unread: + log.warning("%s: read by laya-serve's launcher, not by this worker; use the flags", ", ".join(unread)) + revisions = engine.record_snapshot_revisions() + try: + app = build_app( + make_router(args.device, args.model), + args.model, + args.device, + require_device=args.require_device, + compile=args.compile, + fp16=args.weights == "fp16", + revisions=revisions, + ) + except Exception as exc: # noqa: BLE001 -- any failure before binding means not ready, ever + log.exception("startup failed") + sys.exit(f"laya-worker: not starting: {exc}") + uvicorn.run(app, host=args.host, port=args.port, log_level=args.log_level) + + +if __name__ == "__main__": + main() diff --git a/src/models/laya/README.md b/src/models/laya/README.md index 858da19..3cf9ec0 100644 --- a/src/models/laya/README.md +++ b/src/models/laya/README.md @@ -22,3 +22,34 @@ cargo test --release --locked -p omni-laya --test weights -- --ignored ``` These two CPU tests check all 206 tensor names and shapes, 618 conversion hashes, and the legacy temperature buffer. The normal CI job skips them because it does not download the full checkpoint. + +## Python worker + +The Python worker serves LAYA through laya-serve on CPU and Apple Silicon (PyTorch MPS, +validated on an M1 Pro and, by another contributor, an M5). No native CUDA or Metal backend yet. + +- [`src/frontend/laya_mps.py`](../../frontend/laya_mps.py): the HTTP worker. laya-serve (`laya[serve]==0.3.20`) + with its request handling unchanged, started as `PYTHONPATH=src python -m frontend.laya_mps --device mps`. +- `engine.py`: what the worker runs before readiness (a warmup of every loaded model over short, long and + multi-question requests) and what `/health` reports about a loaded model, read on every call: device, + weight and autocast dtypes, checkpoint and the revision the weights were loaded from, `device_mismatch`. +- `optimize.py`: the two GPU options. `--compile` compiles one-question requests end to end and, for several + questions, only the encoder (Laya's decision head is slower compiled on MPS). `--weights fp16` keeps the + checkpoint's fp16 weights instead of Laya's fp32 upcast (`act_head` stays fp32). Both apply on the GPU + only: on the CPU, including after Laya falls back to it on a GPU out-of-memory error, the worker runs + Laya's fp32 model uncompiled. +- Tests: [`tests/laya/`](../../../tests/laya/). The unit tests use a fake router; `LAYA_CONTRACT=1` adds + contract tests against a real worker on the CPU. + +What the worker changes against laya-serve, with measurements, is in the +[Apple Silicon recipe](../../../recipe/laya/apple-silicon.md): it binds only after the warmup (laya-serve +answers `/health` before any forward pass, so its first request took 0.7–1.1 s against 70–81 ms), and +`/health` tells the device the model is actually on (laya-serve reports the configured one; Laya falls +back to the CPU with only a printed warning). `--require-device` exits at startup if a model is not on +the requested device. + +```sh +PYTHONPATH=src python -m frontend.laya_mps --device mps --model english +PYTHONPATH=src python -m pytest tests/laya # unit tests, no model +LAYA_CONTRACT=1 PYTHONPATH=src python -m pytest tests/laya # plus contract tests on CPU, loads the checkpoint +``` diff --git a/src/models/laya/engine.py b/src/models/laya/engine.py new file mode 100644 index 0000000..2ec8cd5 --- /dev/null +++ b/src/models/laya/engine.py @@ -0,0 +1,107 @@ +"""Laya as a served model: what to run before readiness and what to report about a loaded model. + +laya's `Router` and `Agent` do the loading and the forward passes. This module adds the warmup every +loaded model runs before the worker binds, the description of a loaded model that `/health` returns, and +the record of which checkpoint revision laya actually loaded. The HTTP worker is `frontend/laya_mps.py`; +the GPU optimizations are `optimize.py` next to this file. +""" + +import time +from pathlib import Path +from typing import Any + +# (words of state, questions). Short, mid-length and near-window states, then several questions; the last +# shape has enough rows to reach laya's fp16 autocast on MPS (`agent.mps_amp_min_rows`). +_CHOICE = { + "type": "choice", + "instructions": "Which team should handle this?", + "criteria": {"billing": "Charges and refunds", "technical": "Software problems", "other": "Anything else"}, +} +_SCORE = {"type": "score", "instructions": "How urgent is it?", "criteria": ["Low", "Medium", "High"]} +_NOUL = {"type": "noul", "instructions": "Does the customer ask for a refund?"} +WARMUP_SHAPES = [ + (10, {"q": _CHOICE}), + (150, {"q": _CHOICE}), + (400, {"q": _CHOICE}), + (10, {"a": _CHOICE, "b": _SCORE, "c": _NOUL}), + (10, {f"q{i}": q for i, q in enumerate([_CHOICE, _SCORE, _NOUL, _CHOICE, _SCORE, _NOUL])}), +] +WARMUP_REPEATS = 2 +WARMUP_MAX_ROWS = max(len(questions) for _, questions in WARMUP_SHAPES) + + +def warmup(router: Any, model: str, shapes=WARMUP_SHAPES, repeats: int = WARMUP_REPEATS) -> dict[str, Any]: + """Any failure propagates: a worker that cannot answer must not bind.""" + started = time.perf_counter() + routing = None + for words, questions in shapes: + state = " ".join(["refund"] * words) + for _ in range(repeats): + result = router.predict(state, questions, model=model) + routing = result.get("routing") or routing + return {"warmup_ms": round((time.perf_counter() - started) * 1000, 1), "routing": routing} + + +def _checkpoint_name(repo_id: str, allow_patterns: Any) -> str: + """laya's name for what one download fetched: "", or "/" for a bundled checkpoint. + laya restricts each download to one checkpoint's files, which all sit under its subfolder if it has one.""" + patterns = [allow_patterns] if isinstance(allow_patterns, str) else list(allow_patterns or [""]) + folders = {pattern.split("/")[0] if "/" in pattern else "" for pattern in patterns} + subfolder = folders.pop() if len(folders) == 1 else "" + return f"{repo_id}/{subfolder}" if subfolder else repo_id + + +def record_snapshot_revisions() -> dict[str, str]: + """Record the commit each Hugging Face checkpoint was last downloaded at, keyed by laya's name for it + (`routing["repo"]`). + + laya calls huggingface_hub.snapshot_download while loading and keeps only the repo id; the + returned path (.../snapshots//...) is the only place the loaded revision appears. Call this + before the router loads anything. A checkpoint loaded from a local path records nothing. An entry + changes when the same checkpoint is downloaded again, so read it with `loaded_revision` right after a + checkpoint has loaded and keep that value. + """ + import huggingface_hub + + revisions: dict[str, str] = {} + original = huggingface_hub.snapshot_download + + def recording(repo_id, *args, **kwargs): + path = original(repo_id, *args, **kwargs) + parts = Path(path).parts + if "snapshots" in parts[:-1]: + name = _checkpoint_name(repo_id, kwargs.get("allow_patterns")) + revisions[name] = parts[parts.index("snapshots") + 1] + return path + + huggingface_hub.snapshot_download = recording + return revisions + + +def loaded_revision(revisions: dict[str, str] | None, routing: dict[str, Any] | None) -> str | None: + """The commit of the checkpoint that has just loaded, if it was downloaded.""" + return (revisions or {}).get((routing or {}).get("repo")) + + +def describe( + agent: Any, requested: str | None, routing: dict[str, Any] | None, revision: str | None = None +) -> dict[str, Any]: + """What /health reports about one loaded agent, read from the agent as it is now.""" + device = str(getattr(agent, "device", "unknown")) + model = getattr(agent, "model", None) + try: + weights = str(next(model.parameters()).dtype) if model is not None else None + except (AttributeError, StopIteration, TypeError): + weights = None + repo = (routing or {}).get("repo") + requested_type = requested.split(":")[0] if requested else None + return { + "device": device, + "requested_device": requested or "auto", + "device_mismatch": bool(requested_type) and device.split(":")[0] != requested_type, + "weights_dtype": weights, + "autocast_dtype": str(getattr(agent, "dtype", None)), + "mps_amp_min_rows": getattr(agent, "mps_amp_min_rows", None), + "checkpoint": repo, + "revision": revision, + } diff --git a/src/models/laya/optimize.py b/src/models/laya/optimize.py new file mode 100644 index 0000000..2b6d369 --- /dev/null +++ b/src/models/laya/optimize.py @@ -0,0 +1,111 @@ +"""What the Laya worker changes about the model itself to make it faster on the GPU. + +- fp16 weights: keep the checkpoint's own precision instead of laya's fp32 upcast. +- Compile: torch.compile for the batches where it pays off on MPS. + +Both go through one wrapper module that replaces `agent.model`, so that on the CPU the worker always runs +laya's own fp32 model. frontend/laya_mps.py decides when to apply them and reports the result in /health. +""" + +import functools +from typing import Any + + +@functools.cache +def _served_class() -> type: + """The module class the worker puts in place of laya's model (built on first use: torch is imported late).""" + import torch + + class Served(torch.nn.Module): + def __init__(self, eager): + super().__init__() + self.eager = eager + self.fp16 = False + self.paths = None # (whole model compiled, encoder-only compiled); a tuple is not a submodule + + def forward(self, input_ids, *args, **kwargs): + if input_ids.device.type == "cpu": + if self.fp16: + self.eager.float() + self.fp16 = False + return self.eager(input_ids, *args, **kwargs) + if self.paths is None: + return self.eager(input_ids, *args, **kwargs) + whole, encoder_only = self.paths + return (whole if input_ids.shape[0] == 1 else encoder_only)(input_ids, *args, **kwargs) + + return Served + + +def _served(agent: Any) -> Any: + """Wrap `agent.model` once and return the wrapper. + + On the GPU it runs the compiled paths when there are any, otherwise laya's model. On the CPU it always + runs laya's model in fp32: laya moves the model to the CPU when a request runs out of GPU memory, and + there fp16 weights are slower than fp32 and the compiled graphs would have to recompile first. So after + a fallback the worker behaves like plain laya. + """ + if not isinstance(agent.model, _served_class()): + agent.model = _served_class()(agent.model) + return agent.model + + +def use_fp16_weights(agent: Any) -> None: + """Keep the weights in fp16, the checkpoint's own precision, so the conversion is exact. laya 0.3.20 + upcasts them to fp32 on MPS and CPU. `act_head` stays fp32 because laya feeds it `.float()` features. + For the GPU only: see _served for what happens on the CPU.""" + served = _served(agent) + served.eager.half() + act_head = getattr(served.eager, "act_head", None) + if act_head is not None: + act_head.float() + served.fp16 = True + + +def compile_agent(agent: Any) -> None: + """Compile the model for the batches where it pays off on MPS, sharing the same parameters. + + A batch of one row (one question) runs the whole model compiled. A batch of several rows runs only + the encoder compiled and laya's decision head eagerly: the head is two nn.TransformerEncoderLayer + with a key padding mask, which lose PyTorch's fused fast path when compiled and get slower on + padded multi-row batches. Does nothing for an agent that is already compiled. + """ + import copy + + import torch + + served = _served(agent) + if served.paths is not None: + return + eager = served.eager + whole = torch.compile(eager, dynamic=True) + encoder_only = copy.copy(eager) # same parameters and submodules ... + encoder_only._modules = dict(eager._modules) # ... except the encoder slot + encoder_only._modules["encoder"] = torch.compile(eager.encoder, dynamic=True) + served.paths = (whole, encoder_only) + + +def _on_cpu(agent: Any) -> bool: + return str(getattr(agent, "device", "")).startswith("cpu") + + +def apply(agent: Any, *, fp16: bool, compile: bool) -> bool: + """Does nothing and returns False for a model on the CPU.""" + if _on_cpu(agent): + return False + if fp16: + use_fp16_weights(agent) + if compile: + compile_agent(agent) + return True + + +def compile_active(agent: Any) -> bool: + """Whether requests to this agent run the compiled paths now. False on the CPU, also after a fallback.""" + return getattr(getattr(agent, "model", None), "paths", None) is not None and not _on_cpu(agent) + + +def compiled_graphs() -> int: + from torch._dynamo.utils import counters + + return int(counters["stats"]["unique_graphs"]) diff --git a/tests/laya/test_bench.py b/tests/laya/test_bench.py new file mode 100644 index 0000000..f0e9374 --- /dev/null +++ b/tests/laya/test_bench.py @@ -0,0 +1,155 @@ +"""Checks on the benchmark inputs, the run header and the documented commands in recipe/laya. No model.""" + +import json +import subprocess +import sys +from pathlib import Path + +REPO = Path(__file__).resolve().parents[2] +BENCH = REPO / "recipe/laya/bench" +sys.path.insert(0, str(BENCH)) + +import env as bench_env # noqa: E402 + + +def rows(path): + return [json.loads(line) for line in Path(path).read_text().splitlines()] + + +def test_header_records_what_makes_two_runs_comparable(monkeypatch): + record = bench_env.header("some/repo", extra=1) + assert record["type"] == "env" and record["extra"] == 1 and record["checkpoint"] == "some/repo" + head = subprocess.run(["git", "-C", str(REPO), "rev-parse", "HEAD"], capture_output=True, text=True).stdout.strip() + assert record["omni_sha"] == head and isinstance(record["omni_dirty"], bool) + import laya # noqa: F401 + from importlib.metadata import version + + assert (record["laya"], record["torch"], record["transformers"]) == tuple( + version(package) for package in ("laya", "torch", "transformers") + ) + assert record["argv"] == sys.argv and record["python"] == ".".join(map(str, sys.version_info[:3])) + assert record["loadavg_1m"] >= 0 + if sys.platform == "darwin": # the machine probes use macOS tools; elsewhere they are None (test below) + assert record["power"] and record["chip"] and record["mem_gb"] > 0 + assert record["utc"].endswith("+00:00") + + +def test_the_fixed_inputs_span_question_types_lengths_and_option_counts(): + workloads = rows(BENCH / "workloads.jsonl") + assert len({w["id"] for w in workloads}) == len(workloads) + questions = [q for w in workloads for q in w["questions"].values()] + assert {q["type"] for q in questions} == {"choice", "score", "noul"} + assert {len(q["criteria"]) for q in questions if q["type"] == "choice"} >= {2, 5, 10} + bench = [w for w in workloads if w["kind"] == "bench"] + words = sorted(len(w["state"].split()) for w in bench) + assert ( + words[0] < 20 and 100 < max(w for w in words if w < 200) and words[-1] > 300 + ) # short, medium, near the window + assert {len(w["questions"]) for w in bench} >= {1, 3, 6} # below and above laya's autocast threshold of 5 rows + parity = [w for w in workloads if w["kind"] == "parity"] + assert {q["type"] for w in parity for q in w["questions"].values()} == {"choice", "score", "noul"} + + +DOCUMENTS = ["recipe/laya/apple-silicon.md", "recipe/laya/bench/README.md", "src/models/laya/README.md"] +SCRIPTS = { + "frontend.laya_mps": "src/frontend/laya_mps.py", + **{f"recipe/laya/bench/{name}.py": f"recipe/laya/bench/{name}.py" + for name in ("bench_http", "bench_inproc", "paired", "profile_mps", "report")}, +} # fmt: skip + + +def documented_commands(): + import re + + for document in DOCUMENTS: + text = (REPO / document).read_text() + for block in re.findall(r"```sh\n(.*?)```", text, flags=re.S): + for command in block.replace("\\\n", " ").splitlines(): + for name, source in SCRIPTS.items(): + if name in command: + yield document, command, name, source + + +def test_documented_commands_use_flags_and_files_that_exist(): + import re + + commands = list(documented_commands()) + assert len(commands) >= 15 + for document, command, name, source in commands: + assert (REPO / source).exists(), f"{document}: {source}" + options = set(re.findall(r'add_argument\(\s*"(--?[a-z][a-z-]*)"', (REPO / source).read_text())) + ours = command.split("--spawn")[0] if name.startswith("recipe") else command.split(name, 1)[1] + if name == "recipe/laya/bench/paired.py": + ours = re.sub(r'"[^"]*"', "", ours) # --a/--b carry flags of the worker, checked below + for worker_flags in re.findall(r'--[ab] "([^"]*)"', command): + assert set(re.findall(r"--[a-z-]+", worker_flags)) <= set( + re.findall(r'add_argument\(\s*"(--[a-z-]+)"', (REPO / SCRIPTS["frontend.laya_mps"]).read_text()) + ), f"{document}: {command}" + used = set(re.findall(r"(? 0 + + +def test_the_checkpoint_stays_loaded_across_requests(worker): + port, _ = worker + for _ in range(5): + assert decide(port, {"q": NOUL})[0] == 200 + health = json.loads(call(port, "GET", "/health")[1]) + assert health["loaded"] == ["english"] and health["warmup_ms"] == startup["first_health"]["warmup_ms"] + + +def test_first_request_after_ready_is_not_a_cold_start(worker): + port, (status, _, first_ms) = worker + assert status == 200 + warm = statistics.median(decide(port, {"q": CHOICE})[2] for _ in range(20)) + # Without the warmup the first request costs several hundred ms more than a warm one. A few tens of ms + # remain on MPS: the GPU has been idle since the warmup, and any request after a pause pays that. + assert first_ms <= warm + 100, f"first {first_ms:.0f} ms, warm p50 {warm:.0f} ms" + + +@pytest.mark.parametrize( + ("questions", "kinds"), + [ + ({"q": CHOICE}, {"q": "choice"}), + ({"q": SCORE}, {"q": "score"}), + ({"q": NOUL}, {"q": "noul"}), + ({"a": CHOICE, "b": SCORE, "c": NOUL}, {"a": "choice", "b": "score", "c": "noul"}), + ], + ids=["choice", "score", "noul", "combined"], +) +def test_decisions(worker, questions, kinds): + port, _ = worker + status, body, _ = decide(port, questions) + assert status == 200 + result = json.loads(body) + assert set(result["answers"]) == set(kinds) + assert result["usage"]["input_tokens"] > 0 + for qid, kind in kinds.items(): + answer = result["answers"][qid] + assert answer["type"] == kind + if kind == "noul": + assert 0.0 <= answer["noul"] <= 1.0 + else: + assert sum(answer["probabilities"].values()) == pytest.approx(1.0, abs=1e-3) + if kind == "choice": + assert answer["choice"] in CHOICE["criteria"] + + +SIX = {f"q{i}": q for i, q in enumerate([CHOICE, SCORE, NOUL, CHOICE, SCORE, NOUL])} + + +@pytest.mark.parametrize( + "questions", + [{"q": CHOICE}, {"q": SCORE}, {"q": NOUL}, {"a": CHOICE, "b": SCORE, "c": NOUL}, SIX], + ids=["choice", "score", "noul", "combined", "six-questions"], +) +def test_answers_match_laya_itself(worker, reference, questions): + port, _ = worker + status, body, _ = decide(port, questions) + assert status == 200 + served = json.loads(body) + expected = reference.system_one(STATE, questions) + assert set(expected) <= set(served) # the worker adds `routing`, it drops nothing + assert served["usage"] == expected["usage"] + reduced = "fp16" in FLAGS or (DEVICE == "mps" and len(questions) >= startup["first_health"]["mps_amp_min_rows"]) + tolerance = 1e-2 if reduced else 1e-3 + for qid, want in expected["answers"].items(): + got = served["answers"][qid] + assert set(got) == set(want) + if want["type"] == "noul": + assert got["noul"] == pytest.approx(want["noul"], abs=tolerance) + continue + assert set(got["probabilities"]) == set(want["probabilities"]) + for option, p in want["probabilities"].items(): + assert got["probabilities"][option] == pytest.approx(p, abs=tolerance) + if want["type"] == "choice": + assert got["choice"] == want["choice"] == max(got["probabilities"], key=got["probabilities"].get) + + +def test_same_request_same_answer(worker): + port, _ = worker + first, second = (json.loads(decide(port, {"a": CHOICE, "b": NOUL})[1])["answers"] for _ in range(2)) + assert first == second + + +@pytest.mark.parametrize( + ("body", "raw", "token", "expected"), + [ + (b"{not json", True, TOKEN, 400), + ({"model": "english", "state": STATE}, False, TOKEN, 400), + ( + {"model": "english", "state": STATE, "questions": {"q": {"type": "bogus", "instructions": "?"}}}, + False, + TOKEN, + 422, + ), + (b"x" * (2 * 1024 * 1024 + 1), True, TOKEN, 413), + ({"model": "english", "state": STATE, "questions": {"q": NOUL}}, False, "wrong", 401), + ({"model": "english", "state": STATE, "questions": {"q": NOUL}}, False, None, 401), + ], + ids=["malformed-json", "no-questions", "bad-question", "too-large", "wrong-token", "no-token"], +) +def test_errors(worker, body, raw, token, expected): + port, _ = worker + status, payload, _ = call(port, "POST", "/v1/systemone", body, token=token, raw=raw) + assert status == expected, payload[:200] diff --git a/tests/laya/test_worker.py b/tests/laya/test_worker.py new file mode 100644 index 0000000..b8f81f7 --- /dev/null +++ b/tests/laya/test_worker.py @@ -0,0 +1,772 @@ +"""Unit tests for the Laya worker. A fake Router stands in for laya's; no model is loaded. + +python -m pytest src/models/laya/tests +""" + +import sys +from pathlib import Path +from types import SimpleNamespace + +import pytest +from fastapi.testclient import TestClient + +from frontend import laya_mps as worker +from models.laya import engine, optimize + +ANSWER = {"type": "noul", "noul": 0.9, "confidence": 0.9} + + +def on_gpu(rows): + """Stands in for input_ids on the GPU: the wrapper only reads its device type and number of rows.""" + return SimpleNamespace(device=SimpleNamespace(type="mps"), shape=(rows, 7)) + + +class FakeAgent: + def __init__(self, device="mps", dtype="torch.float16"): + self.device = device + self.dtype = dtype + self.mps_amp_min_rows = 5 + + +class FakeRouter: + """The part of laya.router.Router the worker and laya.serve.create_app use.""" + + def __init__(self, agent=None, fail_on_call=None, agents=None): + self.agent = agent or FakeAgent() + self.agents = agents if agents is not None else {"english": self.agent} + self.calls = [] + self.loads = [] + self.hooks = [] + self.fail_on_call = fail_on_call + + @property + def loaded(self): + return list(self.agents) + + def load(self, name): + self.loads.append(name) + return self.agents.setdefault(name, self.agent) + + def add_hook(self, hook): + self.hooks.append(hook) + + def remove_hook(self, hook): + self.hooks.remove(hook) + + def unload(self, name): + self.agents.pop(name, None) + for hook in self.hooks: + hook.on_evict(SimpleNamespace(model=name)) + + def load_while_serving(self, name, agent): + """What laya's Router.load does for a checkpoint that is not resident yet.""" + self.agents[name] = agent + for hook in self.hooks: + hook.on_load(SimpleNamespace(model=name, agent=agent)) + + def predict(self, state, questions, model=None): + self.calls.append((state, questions, model)) + if self.fail_on_call is not None and len(self.calls) == self.fail_on_call: + raise RuntimeError("MPS backend out of memory") + return { + "model": "laya-rl-agent", + "answers": {qid: ANSWER for qid in questions}, + "usage": {"input_tokens": 10, "output_tokens": 0}, + "routing": {"model": model, "repo": "convaiinnovations/laya"}, + } + + +def test_warmup_covers_short_long_and_fp16_multi_question_shapes(): + router = FakeRouter() + engine.warmup(router, "english") + words = {len(state.split()) for state, _, _ in router.calls} + rows = {len(questions) for _, questions, _ in router.calls} + assert min(words) <= 20 and max(words) >= 400 + assert max(rows) >= router.agent.mps_amp_min_rows + assert {q["type"] for _, questions, _ in router.calls for q in questions.values()} == {"choice", "score", "noul"} + assert len(router.calls) == len(engine.WARMUP_SHAPES) * engine.WARMUP_REPEATS + assert all(model == "english" for _, _, model in router.calls) + + +def test_warmup_runs_before_the_app_exists(): + router = FakeRouter() + worker.build_app(router, "english", "mps") + assert len(router.calls) == len(engine.WARMUP_SHAPES) * engine.WARMUP_REPEATS + + +def test_warmup_failure_raises_and_no_app_is_built(): + with pytest.raises(RuntimeError, match="out of memory"): + worker.build_app(FakeRouter(fail_on_call=3), "english", "mps") + + +def test_health_reports_the_agent_device_not_the_requested_one(): + router = FakeRouter(FakeAgent(device="cpu", dtype="torch.float32")) + health = TestClient(worker.build_app(router, "english", "mps")).get("/health").json() + assert health["device"] == "cpu" + assert health["requested_device"] == "mps" + assert health["device_mismatch"] is True + assert health["ready"] is True + + +def test_health_on_the_requested_device(): + health = TestClient(worker.build_app(FakeRouter(), "english", "mps")).get("/health").json() + assert health["device"] == "mps" + assert health["device_mismatch"] is False + assert health["autocast_dtype"] == "torch.float16" + assert health["checkpoint"] == "convaiinnovations/laya" + assert health["warmup_ms"] >= 0 + + +def test_device_index_is_not_a_mismatch(): + router = FakeRouter(FakeAgent(device="cuda:0")) + assert TestClient(worker.build_app(router, "english", "cuda")).get("/health").json()["device_mismatch"] is False + + +def test_auto_device_is_never_a_mismatch(): + router = FakeRouter(FakeAgent(device="cpu")) + health = TestClient(worker.build_app(router, "english", None)).get("/health").json() + assert health["requested_device"] == "auto" + assert health["device_mismatch"] is False + + +def test_require_device_refuses_to_serve_on_another_device(): + router = FakeRouter(FakeAgent(device="cpu")) + with pytest.raises(RuntimeError, match="asked for mps, english is on cpu"): + worker.build_app(router, "english", "mps", require_device=True) + + +def test_only_one_health_route_remains(): + app = worker.build_app(FakeRouter(), "english", "mps") + assert [r.path for r in app.router.routes if getattr(r, "path", None) == "/health"] == ["/health"] + + +def test_decisions_still_go_through_laya_serve(): + router = FakeRouter() + client = TestClient(worker.build_app(router, "english", "mps")) + before = len(router.calls) + response = client.post( + "/v1/systemone", + json={"model": "english", "state": "refund me", "questions": {"r": {"type": "noul", "instructions": "?"}}}, + ) + assert response.status_code == 200 + assert response.json()["answers"]["r"]["noul"] == 0.9 + assert len(router.calls) == before + 1 + + +def test_main_exits_non_zero_when_warmup_fails(monkeypatch, caplog): + monkeypatch.setattr(worker, "make_router", lambda device, model: FakeRouter(fail_on_call=1)) + monkeypatch.setattr(sys, "argv", ["laya_mps", "--device", "mps"]) + monkeypatch.setattr("uvicorn.run", lambda *a, **k: pytest.fail("must not bind")) + with caplog.at_level("ERROR", logger="laya-worker"), pytest.raises(SystemExit, match="not starting"): + worker.main() + assert "startup failed" in caplog.text and "Traceback" in caplog.text and "out of memory" in caplog.text + + +def test_compile_wraps_the_model_before_warmup(monkeypatch): + router = FakeRouter() + order = [] + monkeypatch.setattr(optimize, "compile_agent", lambda agent: order.append((agent, len(router.calls)))) + worker.build_app(router, "english", "mps", compile=True, graph_counter=lambda: 3) + assert order == [(router.agent, 0)] + + +def test_health_reports_compile_off_by_default(): + health = TestClient(worker.build_app(FakeRouter(), "english", "mps")).get("/health").json() + assert health["compile"] == {"enabled": False} + + +def test_health_flags_graphs_compiled_after_ready(monkeypatch): + monkeypatch.setattr(optimize, "compile_agent", lambda agent: None) + graphs = iter([4, 4, 5]) # at readiness, first /health, second /health after a new shape compiled + client = TestClient( + worker.build_app(FakeRouter(), "english", "mps", compile=True, graph_counter=lambda: next(graphs)) + ) + first = client.get("/health").json()["compile"] + assert first == { + "enabled": True, + "active": False, # compile_agent is stubbed out here + "graphs_at_ready": 4, + "graphs_now": 4, + "recompiled_after_ready": False, + } + assert client.get("/health").json()["compile"]["recompiled_after_ready"] is True + + +def test_compile_failure_means_no_app(monkeypatch): + def broken(agent): + raise RuntimeError("inductor: unsupported op on mps") + + monkeypatch.setattr(optimize, "compile_agent", broken) + with pytest.raises(RuntimeError, match="unsupported op"): + worker.build_app(FakeRouter(), "english", "mps", compile=True) + + +def test_compiled_paths_by_batch_rows(monkeypatch): + import torch + + class Encoder(torch.nn.Module): + def forward(self, input_ids): + return "eager encoder" + + class Model(torch.nn.Module): + def __init__(self): + super().__init__() + self.encoder = Encoder() + self.head = torch.nn.Linear(2, 2) + + def forward(self, input_ids): + return ("head", self.encoder(input_ids)) + + class Stub(torch.nn.Module): + def __init__(self, kind): + super().__init__() + self.kind = kind + + def forward(self, *args): + return self.kind + + def fake_compile(module, dynamic): + assert dynamic is True + return Stub("compiled encoder" if isinstance(module, Encoder) else "whole model compiled") + + monkeypatch.setattr(torch, "compile", fake_compile) + agent = FakeAgent() + model = Model() + agent.model = model + optimize.compile_agent(agent) + assert agent.model(on_gpu(1)) == "whole model compiled" + assert agent.model(on_gpu(3)) == ("head", "compiled encoder") + assert agent.model(torch.zeros(1, 7)) == ("head", "eager encoder") + assert model.encoder(torch.zeros(1, 7)) == "eager encoder" # the original model is left as it was + assert len(list(agent.model.parameters())) == len(list(model.parameters())) # one set of weights + + +def test_every_loaded_model_is_warmed_and_described(): + agents = {"english": FakeAgent(), "multilingual": FakeAgent(device="cpu", dtype="torch.float32")} + router = FakeRouter(agents=agents) + health = TestClient(worker.build_app(router, "english", "mps")).get("/health").json() + per_model = len(engine.WARMUP_SHAPES) * engine.WARMUP_REPEATS + assert [m for _, _, m in router.calls].count("multilingual") == per_model + assert [m for _, _, m in router.calls].count("english") == per_model + assert set(health["models"]) == {"english", "multilingual"} + assert health["device"] == "mps" # the top level summarises --model + assert health["models"]["multilingual"]["device"] == "cpu" + assert health["device_mismatch"] is True + + +def test_a_model_that_is_not_preloaded_is_not_loaded_for_warmup(): + router = FakeRouter(agents={"multilingual": FakeAgent()}) + health = TestClient(worker.build_app(router, "english", "mps")).get("/health").json() + assert "english" not in router.loads + assert {m for _, _, m in router.calls} == {"multilingual"} + assert set(health["models"]) == {"multilingual"} + + +def test_nothing_preloaded_warms_the_worker_model(): + router = FakeRouter(agents={}) + worker.build_app(router, "english", "mps") + assert {m for _, _, m in router.calls} == {"english"} + + +def test_revision_comes_from_the_loaded_snapshot_not_a_guess(): + revisions = {"convaiinnovations/laya": "55cf4c4"} + health = TestClient(worker.build_app(FakeRouter(), "english", "mps", revisions=revisions)).get("/health") + assert health.json()["revision"] == "55cf4c4" + unknown = TestClient(worker.build_app(FakeRouter(), "english", "mps")).get("/health").json() + assert unknown["revision"] is None + + +def test_record_snapshot_revisions_reads_the_downloaded_path(monkeypatch): + import huggingface_hub + + paths = { + "convaiinnovations/laya": "/cache/models--convaiinnovations--laya/snapshots/55cf4c4abc/multilingual", + "/local/checkpoint": "/local/checkpoint", + } + monkeypatch.setattr(huggingface_hub, "snapshot_download", lambda repo_id, **kwargs: paths[repo_id]) + revisions = engine.record_snapshot_revisions() + assert ( + huggingface_hub.snapshot_download("convaiinnovations/laya", allow_patterns=["*"]) + == paths["convaiinnovations/laya"] + ) + huggingface_hub.snapshot_download("/local/checkpoint") + assert revisions == {"convaiinnovations/laya": "55cf4c4abc"} + huggingface_hub.snapshot_download( + "convaiinnovations/laya", allow_patterns=["multilingual/model.safetensors", "multilingual/tokenizer/*"] + ) + huggingface_hub.snapshot_download("convaiinnovations/laya", allow_patterns=["model.safetensors", "tokenizer/*"]) + assert set(revisions) == {"convaiinnovations/laya", "convaiinnovations/laya/multilingual"} + + +def test_fp16_weights_keep_act_head_in_fp32(): + import torch + + class Model(torch.nn.Module): + def __init__(self): + super().__init__() + self.encoder = torch.nn.Linear(4, 4) + self.act_head = torch.nn.Linear(4, 2) + + agent = FakeAgent() + agent.model = Model() + optimize.use_fp16_weights(agent) + assert agent.model.eager.encoder.weight.dtype == torch.float16 + assert agent.model.eager.act_head.weight.dtype == torch.float32 + + +def test_fp16_weights_are_applied_to_every_loaded_model_before_warmup(monkeypatch): + order = [] + agents = {"english": FakeAgent(), "multilingual": FakeAgent()} + router = FakeRouter(agents=agents) + monkeypatch.setattr(optimize, "use_fp16_weights", lambda agent: order.append((agent, len(router.calls)))) + worker.build_app(router, "english", "mps", fp16=True) + assert order == [(agents["english"], 0), (agents["multilingual"], 0)] + + +def test_options_are_not_applied_to_a_model_on_the_cpu(monkeypatch, caplog): + applied = [] + monkeypatch.setattr(optimize, "use_fp16_weights", lambda agent: applied.append("fp16")) + monkeypatch.setattr(optimize, "compile_agent", lambda agent: applied.append("compile")) + with caplog.at_level("WARNING", logger="laya-worker"): + worker.build_app( + FakeRouter(FakeAgent(device="cpu")), "english", "cpu", fp16=True, compile=True, graph_counter=lambda: 0 + ) + assert applied == [] + assert "apply on the GPU only" in caplog.text + caplog.clear() + with caplog.at_level("WARNING", logger="laya-worker"): + worker.build_app(FakeRouter(), "english", "mps", fp16=True, compile=True, graph_counter=lambda: 0) + assert applied == ["fp16", "compile"] + assert "GPU only" not in caplog.text + + +def test_after_a_fallback_to_cpu_the_model_runs_fp32_and_uncompiled(monkeypatch): + import torch + + class Model(torch.nn.Module): + def __init__(self): + super().__init__() + self.encoder = torch.nn.Linear(4, 4) + self.act_head = torch.nn.Linear(4, 2) + + def forward(self, input_ids): + return ("eager", self.encoder.weight.dtype) + + monkeypatch.setattr(torch, "compile", lambda module, dynamic: lambda *a, **k: "compiled") + agent = FakeAgent() + agent.model = Model() + optimize.use_fp16_weights(agent) + optimize.compile_agent(agent) + assert agent.model(on_gpu(1)) == "compiled" + assert next(agent.model.parameters()).dtype == torch.float16 + assert agent.model(torch.zeros(1, 7)) == ("eager", torch.float32) + assert {p.dtype for p in agent.model.parameters()} == {torch.float32} + assert agent.model(torch.zeros(3, 7)) == ("eager", torch.float32) + + +def test_health_follows_a_fallback_to_cpu_after_startup(): + agent = FakeAgent(device="mps") + client = TestClient(worker.build_app(FakeRouter(agent), "english", "mps", require_device=True)) + assert client.get("/health").json()["device_mismatch"] is False + agent.device = "cpu" + agent.dtype = "torch.float32" + health = client.get("/health").json() + assert health["device"] == "cpu" + assert health["device_mismatch"] is True + assert health["models"]["english"]["device"] == "cpu" + assert health["autocast_dtype"] == "torch.float32" + + +def test_health_compile_active_follows_a_fallback_to_cpu(monkeypatch): + import torch + + compiles = [] + monkeypatch.setattr(torch, "compile", lambda module, dynamic: compiles.append(module) or module) + agent = FakeAgent() + agent.model = torch.nn.Sequential() + agent.model.encoder = torch.nn.Identity() + client = TestClient(worker.build_app(FakeRouter(agent), "english", "mps", compile=True, graph_counter=lambda: 2)) + assert client.get("/health").json()["compile"]["active"] is True + optimize.compile_agent(agent) # a second name for the same agent must not compile again + assert len(compiles) == 2 + agent.device = "cpu" + assert client.get("/health").json()["compile"]["active"] is False + + +def test_warning_when_the_warmup_does_not_reach_layas_autocast_rows(caplog): + agent = FakeAgent() + agent.mps_amp_min_rows = engine.WARMUP_MAX_ROWS + 1 + with caplog.at_level("WARNING", logger="laya-worker"): + worker.build_app(FakeRouter(agent), "english", "mps") + assert "the warmup stops at" in caplog.text + caplog.clear() + with caplog.at_level("WARNING", logger="laya-worker"): + worker.build_app(FakeRouter(), "english", "mps") + assert "the warmup stops at" not in caplog.text + + +def test_log_level_applies_to_the_workers_own_log(monkeypatch): + seen = {} + monkeypatch.setattr(worker, "make_router", lambda device, model: FakeRouter()) + monkeypatch.setattr(worker.logging, "basicConfig", lambda **kw: seen.update(kw)) + monkeypatch.setattr("uvicorn.run", lambda app, **kw: seen.update(uvicorn=kw)) + monkeypatch.setattr(sys, "argv", ["laya_mps", "--device", "mps", "--log-level", "warning"]) + worker.main() + assert (seen["level"], seen["uvicorn"]["log_level"]) == ("WARNING", "warning") + assert (seen["uvicorn"]["host"], seen["uvicorn"]["port"]) == ("127.0.0.1", 8000) # local only by default + + +def test_a_checkpoint_loaded_while_serving_is_prepared_and_described(monkeypatch): + applied = [] + monkeypatch.setattr(optimize, "use_fp16_weights", lambda agent: applied.append(agent)) + router = FakeRouter() + client = TestClient(worker.build_app(router, "english", "mps", fp16=True)) + before = len(router.calls) + late = FakeAgent(device="cpu", dtype="torch.float32") + router.load_while_serving("multilingual", late) + assert applied == [router.agent] # the late one is on the CPU, where the options do not apply + assert [m for _, _, m in router.calls[before:]] == ["multilingual"] * ( + len(engine.WARMUP_SHAPES) * engine.WARMUP_REPEATS + ) + health = client.get("/health").json() + assert set(health["models"]) == {"english", "multilingual"} + assert health["models"]["multilingual"]["device"] == "cpu" + assert health["device_mismatch"] is True + assert health["preparing"] == [] + router.unload("multilingual") + health = client.get("/health").json() + assert set(health["models"]) == {"english"} + assert health["device_mismatch"] is False + + +def test_a_late_checkpoint_that_cannot_be_prepared_is_not_kept(): + router = FakeRouter() + client = TestClient(worker.build_app(router, "english", "mps", require_device=True)) + with pytest.raises(RuntimeError, match="multilingual is on cpu"): + router.load_while_serving("multilingual", FakeAgent(device="cpu")) + assert router.loaded == ["english"] + assert set(client.get("/health").json()["models"]) == {"english"} + router.fail_on_call = len(router.calls) + 1 + with pytest.raises(RuntimeError, match="out of memory"): + router.load_while_serving("multilingual", FakeAgent()) + assert router.loaded == ["english"] + + +def test_compile_baseline_moves_when_a_late_checkpoint_compiles(monkeypatch): + monkeypatch.setattr(optimize, "compile_agent", lambda agent: None) + graphs = iter([3, 7, 7]) # startup, after the late checkpoint compiled, /health + router = FakeRouter() + client = TestClient(worker.build_app(router, "english", "mps", compile=True, graph_counter=lambda: next(graphs))) + router.load_while_serving("multilingual", FakeAgent()) + compiled = client.get("/health").json()["compile"] + assert (compiled["graphs_at_ready"], compiled["recompiled_after_ready"]) == (7, False) + + +def test_startup_error_names_every_checkpoint_off_the_requested_device(): + agents = {"english": FakeAgent(device="cpu"), "multilingual": FakeAgent(device="cpu")} + with pytest.raises(RuntimeError, match="asked for mps, english is on cpu, multilingual is on cpu"): + worker.build_app(FakeRouter(agents=agents), "english", "mps", require_device=True) + + +def test_health_names_a_checkpoint_while_it_is_being_prepared(): + router = FakeRouter() + client = TestClient(worker.build_app(router, "english", "mps")) + seen = [] + predict = router.predict + + def predict_and_look(state, questions, model=None): + seen.append(client.get("/health").json()["preparing"]) + return predict(state, questions, model=model) + + router.predict = predict_and_look + router.load_while_serving("multilingual", FakeAgent()) + assert seen[0] == ["multilingual"] + assert client.get("/health").json()["preparing"] == [] + + +class StubAgent: + """Stands in for laya.agent.Agent under laya's real Router: no weights, fixed answers.""" + + devices: dict = {} # subfolder -> the device that checkpoint lands on + + def __init__(self, repo, device=None, token=None, subfolder=None): + self.device = self.devices.get(subfolder, device) + self.dtype = "torch.float16" + self.mps_amp_min_rows = 5 + + def system_one(self, state, questions, lang=None, **_): + return {"model": "stub", "answers": {qid: ANSWER for qid in questions}, "usage": {}} + + +def test_a_failed_late_load_gives_back_the_checkpoint_laya_evicted_for_it(monkeypatch): + import laya.agent + from laya.router import Router + + monkeypatch.setattr(laya.agent, "Agent", StubAgent) + monkeypatch.setattr(StubAgent, "devices", {"typed-decisions": "cpu"}) + router = Router(device="mps") + router.preload(["english", "multilingual"]) # laya keeps two checkpoints by default + client = TestClient(worker.build_app(router, "english", "mps", require_device=True)) + with pytest.raises(RuntimeError, match="typed-decisions is on cpu"): + router.predict("refund me", {"r": {"type": "noul", "instructions": "?"}}, model="typed-decisions") + assert sorted(router.loaded) == ["english", "multilingual"] + health = client.get("/health").json() + assert set(health["models"]) == {"english", "multilingual"} + assert health["preparing"] == [] + + +def test_building_a_second_app_on_a_router_replaces_the_first_apps_hooks(): + router = FakeRouter() + worker.build_app(router, "english", "mps") + worker.build_app(router, "english", "mps") + assert len(router.hooks) == 1 + before = len(router.calls) + router.load_while_serving("multilingual", FakeAgent()) + assert len(router.calls) - before == len(engine.WARMUP_SHAPES) * engine.WARMUP_REPEATS + + +def test_main_warns_about_laya_serve_variables_it_does_not_read(monkeypatch, caplog): + monkeypatch.setattr(worker, "make_router", lambda device, model: FakeRouter()) + monkeypatch.setattr(worker.logging, "basicConfig", lambda **kw: None) + monkeypatch.setattr("uvicorn.run", lambda app, **kw: None) + monkeypatch.setattr(sys, "argv", ["laya_mps", "--device", "mps"]) + monkeypatch.setenv("LAYA_THREADS", "4") + monkeypatch.setenv("LAYA_API_KEY", "k") + with caplog.at_level("WARNING", logger="laya-worker"): + worker.main() + assert "LAYA_THREADS" in caplog.text + assert "LAYA_API_KEY" not in caplog.text + + +def test_graphs_compiled_by_a_late_checkpoint_that_fails_are_not_reported_as_recompiles(monkeypatch): + monkeypatch.setattr(optimize, "compile_agent", lambda agent: None) + router = FakeRouter() + graphs = {"n": 0} + predict = router.predict + + def predict_and_compile(state, questions, model=None): + graphs["n"] += 1 + return predict(state, questions, model=model) + + router.predict = predict_and_compile + client = TestClient(worker.build_app(router, "english", "mps", compile=True, graph_counter=lambda: graphs["n"])) + router.fail_on_call = len(router.calls) + 3 + with pytest.raises(RuntimeError, match="out of memory"): + router.load_while_serving("multilingual", FakeAgent()) + assert client.get("/health").json()["compile"]["recompiled_after_ready"] is False + + +class Repository: + """What the downloading stub agents see: the repository's current commit, and where checkpoints land.""" + + commit = 0 + device: dict = {} # subfolder -> device + broken: set = set() # subfolders whose forward raises + + +class DownloadingAgent: + """laya.agent.Agent without weights. It downloads the way laya does, so the worker's recording sees it.""" + + def __init__(self, repo, device=None, token=None, subfolder=None): + import huggingface_hub + + prefix = f"{subfolder}/" if subfolder else "" + path = huggingface_hub.snapshot_download( + repo, token=token, allow_patterns=[prefix + "model.safetensors", prefix + "tokenizer/*"] + ) + self.loaded_commit = Path(path).name + self.subfolder = subfolder + self.device = Repository.device.get(subfolder, device) + self.dtype = "torch.float16" + self.mps_amp_min_rows = 5 + + def system_one(self, state, questions, lang=None, **_): + if self.subfolder in Repository.broken: + raise RuntimeError("MPS backend out of memory") + return {"model": "stub", "answers": {qid: ANSWER for qid in questions}, "usage": {}} + + +@pytest.fixture +def laya_router(monkeypatch): + """laya's real Router over downloading stub agents, and the revisions the worker records for them.""" + import huggingface_hub + import laya.agent + from laya.router import Router + + monkeypatch.setattr(Repository, "commit", 0) + monkeypatch.setattr(Repository, "device", {}) + monkeypatch.setattr(Repository, "broken", set()) + monkeypatch.setattr( + huggingface_hub, "snapshot_download", lambda repo, **kw: f"/hf/snapshots/commit-{Repository.commit}" + ) + monkeypatch.setattr(laya.agent, "Agent", DownloadingAgent) + return Router(device="mps"), engine.record_snapshot_revisions() + + +def assert_health_matches(client, router): + health = client.get("/health").json() + agents = dict(router._agents) + assert set(health["models"]) == set(router.loaded) == set(agents) + assert health["preparing"] == [] + for name, agent in agents.items(): + assert health["models"][name]["device"] == str(agent.device) + assert health["models"][name]["revision"] == agent.loaded_commit + assert health["device_mismatch"] == any(str(agent.device) != "mps" for agent in agents.values()) + + +def test_checkpoints_of_one_repository_keep_the_revision_of_their_own_load(laya_router): + router, revisions = laya_router + router.load("english") + Repository.commit = 1 # the repository moves on between the two loads + router.load("multilingual") + client = TestClient(worker.build_app(router, "english", "mps", revisions=revisions)) + assert_health_matches(client, router) + Repository.commit = 2 + router.predict("refund me", {"r": {"type": "noul", "instructions": "?"}}, model="typed-decisions") + assert_health_matches(client, router) + + +@pytest.mark.parametrize("require_device", [False, True]) +@pytest.mark.parametrize("seed", range(8)) +def test_health_matches_the_router_after_any_sequence_of_loads(laya_router, seed, require_device): + import random + + router, revisions = laya_router + names = ["english", "multilingual", "typed-decisions"] + subfolder = {"english": None, "multilingual": "multilingual", "typed-decisions": "typed-decisions"} + rng = random.Random(seed) + router.preload(rng.sample(names, rng.choice([1, 2]))) + client = TestClient( + worker.build_app(router, router.loaded[0], "mps", require_device=require_device, revisions=revisions) + ) + assert_health_matches(client, router) + for _ in range(12): + Repository.commit += rng.random() < 0.3 + Repository.device = {subfolder[n]: "cpu" for n in names if rng.random() < 0.25} + Repository.broken = {subfolder[n] for n in names if rng.random() < 0.15} + try: + router.predict("refund me", {"r": {"type": "noul", "instructions": "?"}}, model=rng.choice(names)) + except RuntimeError: + pass # a failed late load or forward: the request fails, /health must still match the router + Repository.broken = set() + assert_health_matches(client, router) + + +def test_health_answers_with_no_resident_checkpoint(laya_router): + router, revisions = laya_router + router.preload(["english"]) + client = TestClient(worker.build_app(router, "english", "mps", revisions=revisions)) + router.unload() + assert client.get("/health").json()["models"] == {} + + +def test_a_checkpoint_loaded_while_serving_gets_the_gpu_options_before_its_warmup(monkeypatch): + applied = [] + router = FakeRouter() + monkeypatch.setattr(optimize, "use_fp16_weights", lambda agent: applied.append(("fp16", len(router.calls)))) + monkeypatch.setattr(optimize, "compile_agent", lambda agent: applied.append(("compile", len(router.calls)))) + worker.build_app(router, "english", "mps", fp16=True, compile=True, graph_counter=lambda: 0) + warmed_at_startup = len(router.calls) + router.load_while_serving("multilingual", FakeAgent()) + assert applied[2:] == [("fp16", warmed_at_startup), ("compile", warmed_at_startup)] + + +def test_a_checkpoint_off_the_requested_device_is_logged_when_it_is_allowed(caplog): + with caplog.at_level("WARNING", logger="laya-worker"): + worker.build_app(FakeRouter(FakeAgent(device="cpu")), "english", "mps") + assert "asked for mps, english is on cpu" in caplog.text + + +def test_no_warning_when_the_warmup_just_reaches_layas_autocast_rows(caplog): + agent = FakeAgent() + agent.mps_amp_min_rows = engine.WARMUP_MAX_ROWS + with caplog.at_level("WARNING", logger="laya-worker"): + worker.build_app(FakeRouter(agent), "english", "mps") + assert "the warmup stops at" not in caplog.text + + +def test_failed_late_loads_are_logged_with_their_cause(laya_router, caplog): + router, revisions = laya_router + router.preload(["english", "multilingual"]) + worker.build_app(router, "english", "mps", require_device=True, revisions=revisions) + Repository.device = {"typed-decisions": "cpu", None: "cpu"} # english cannot come back either + with caplog.at_level("ERROR", logger="laya-worker"), pytest.raises(RuntimeError): + router.predict("refund me", {"r": {"type": "noul", "instructions": "?"}}, model="typed-decisions") + assert "typed-decisions could not be prepared and is unloaded" in caplog.text + assert "english was evicted for typed-decisions and could not be loaded again" in caplog.text + assert "asked for mps, typed-decisions is on cpu" in caplog.text # the traceback of the cause + assert router.loaded == ["multilingual"] + + +def test_a_failed_late_load_reloads_only_what_was_evicted_for_it(laya_router, monkeypatch): + router, revisions = laya_router + router.max_loaded = 1 + router.load("english") + worker.build_app(router, "english", "mps", require_device=True, revisions=revisions) + router.predict("refund me", {"r": {"type": "noul", "instructions": "?"}}, model="multilingual") # evicts english + built = [] + original = DownloadingAgent.__init__ + monkeypatch.setattr( + DownloadingAgent, + "__init__", + lambda self, repo, **kw: built.append(kw.get("subfolder")) or original(self, repo, **kw), + ) + Repository.device = {"typed-decisions": "cpu"} + with pytest.raises(RuntimeError): + router.predict("refund me", {"r": {"type": "noul", "instructions": "?"}}, model="typed-decisions") + assert built == ["typed-decisions", "multilingual"] and router.loaded == ["multilingual"] + + +def test_describe_reads_the_weight_dtype_from_the_model_as_it_is(): + import torch + + agent = FakeAgent() + assert engine.describe(agent, "mps", None)["weights_dtype"] is None # no model to read + agent.model = torch.nn.Linear(2, 2) + assert engine.describe(agent, "mps", None)["weights_dtype"] == "torch.float32" + agent.model.half() + assert engine.describe(agent, "mps", None)["weights_dtype"] == "torch.float16" + + +def test_warmup_time_is_reported_in_milliseconds(monkeypatch): + clock = iter([10.0, 10.25]) + monkeypatch.setattr(engine.time, "perf_counter", lambda: next(clock)) + assert engine.warmup(FakeRouter(), "english")["warmup_ms"] == 250.0 + + +def test_apply_says_whether_the_options_were_applied(monkeypatch): + monkeypatch.setattr(optimize, "use_fp16_weights", lambda agent: None) + assert optimize.apply(FakeAgent(device="cpu"), fp16=True, compile=False) is False + assert optimize.apply(FakeAgent(device="mps"), fp16=True, compile=False) is True + + +def test_the_checkpoint_is_built_once_and_serves_every_request(laya_router, monkeypatch): + router, revisions = laya_router + router.preload(["english"]) + built = [] + original = DownloadingAgent.__init__ + monkeypatch.setattr( + DownloadingAgent, "__init__", lambda self, repo, **kw: built.append(repo) or original(self, repo, **kw) + ) + client = TestClient(worker.build_app(router, "english", "mps", revisions=revisions)) + body = {"model": "english", "state": "refund me", "questions": {"r": {"type": "noul", "instructions": "?"}}} + for _ in range(3): + response = client.post("/v1/systemone", json=body) + assert response.status_code == 200 and response.json()["answers"]["r"] == ANSWER + assert built == [] and router.loaded == ["english"] + + +def test_main_binds_only_after_every_warmup_request(monkeypatch): + router = FakeRouter() + bound_after = [] + monkeypatch.setattr(worker, "make_router", lambda device, model: router) + monkeypatch.setattr(worker.logging, "basicConfig", lambda **kw: None) + monkeypatch.setattr("uvicorn.run", lambda app, **kw: bound_after.append(len(router.calls))) + monkeypatch.setattr(sys, "argv", ["laya_mps", "--device", "mps"]) + worker.main() + assert bound_after == [len(engine.WARMUP_SHAPES) * engine.WARMUP_REPEATS] + + +def test_the_warmup_asks_every_question_type(): + kinds = {q["type"] for _, questions in engine.WARMUP_SHAPES for q in questions.values()} + assert kinds == {"choice", "score", "noul"}