From f0cccd8e8441881d4dda68909630893e9a64f24b Mon Sep 17 00:00:00 2001 From: Tianyao Wu Date: Sun, 27 Sep 2026 20:24:10 +0800 Subject: [PATCH 01/10] cua_s1: add a text worker that loads Qwen3.5-4B through Transformers Add a /v1/systemone worker for the Cua-S1 4B 0.2 text adapter in src/models/cua_s1/text/, next to the multimodal worker proposed in #12. It loads the base model and the PEFT adapter directly through Transformers and PEFT, follows the contract in src/models/cua_s1/README.md, and answers choice questions only. Add the fixed input set and tests in tests/cua_s1/ (contract and HTTP tests that need no weights, and tokenizer checks), and recipe/cua_s1/text.md with setup, launch, a parity check against upstream FourBModel and a latency script. Ignore the recipe's weights/ and .venv/ with the same .gitignore lines as #12. Part of #10. Signed-off-by: Tianyao Wu --- .gitignore | 7 + recipe/README.md | 2 +- recipe/cua_s1/bench_text.py | 122 ++++++++ recipe/cua_s1/compare_text_with_upstream.py | 204 +++++++++++++ recipe/cua_s1/requirements-text.txt | 13 + recipe/cua_s1/text.md | 99 ++++++ src/models/cua_s1/README.md | 16 +- src/models/cua_s1/text/THIRD_PARTY_NOTICES.md | 23 ++ src/models/cua_s1/text/adapter.py | 51 ++++ src/models/cua_s1/text/contract.py | 273 +++++++++++++++++ src/models/cua_s1/text/engine.py | 84 +++++ src/models/cua_s1/text/server.py | 234 ++++++++++++++ tests/cua_s1/data/text_inputs.json | 289 ++++++++++++++++++ tests/cua_s1/test_text_adapter.py | 36 +++ tests/cua_s1/test_text_contract.py | 216 +++++++++++++ tests/cua_s1/test_text_server.py | 167 ++++++++++ tests/cua_s1/test_text_tokenizer.py | 45 +++ 17 files changed, 1876 insertions(+), 5 deletions(-) create mode 100644 recipe/cua_s1/bench_text.py create mode 100644 recipe/cua_s1/compare_text_with_upstream.py create mode 100644 recipe/cua_s1/requirements-text.txt create mode 100644 recipe/cua_s1/text.md create mode 100644 src/models/cua_s1/text/THIRD_PARTY_NOTICES.md create mode 100644 src/models/cua_s1/text/adapter.py create mode 100644 src/models/cua_s1/text/contract.py create mode 100644 src/models/cua_s1/text/engine.py create mode 100644 src/models/cua_s1/text/server.py create mode 100644 tests/cua_s1/data/text_inputs.json create mode 100644 tests/cua_s1/test_text_adapter.py create mode 100644 tests/cua_s1/test_text_contract.py create mode 100644 tests/cua_s1/test_text_server.py create mode 100644 tests/cua_s1/test_text_tokenizer.py diff --git a/.gitignore b/.gitignore index ad679558..a2db6446 100644 --- a/.gitignore +++ b/.gitignore @@ -19,3 +19,10 @@ target # and can be added to the global gitignore or merged into this file. For a more nuclear # option (not recommended) you can uncomment the following to ignore the entire idea folder. #.idea/ + +# Python workers and local model artifacts +__pycache__/ +.pytest_cache/ +.ruff_cache/ +.venv/ +weights/ diff --git a/recipe/README.md b/recipe/README.md index 1824bc5c..7c13f9cf 100644 --- a/recipe/README.md +++ b/recipe/README.md @@ -4,4 +4,4 @@ Top-level home for model setup instructions, launch commands, configuration exam Recipes use the frontend, model engines, and GPU backends. Reusable implementation code belongs in those components rather than in recipes. -Status: layout only; no runnable recipes yet. Add commands and supported configurations once they can be validated against an implementation. +Status: [`cua_s1/text.md`](cua_s1/text.md) runs the Cua-S1 4B 0.2 text worker. Add commands and supported configurations once they can be validated against an implementation. diff --git a/recipe/cua_s1/bench_text.py b/recipe/cua_s1/bench_text.py new file mode 100644 index 00000000..f5ee9c88 --- /dev/null +++ b/recipe/cua_s1/bench_text.py @@ -0,0 +1,122 @@ +"""Send the fixed input set to a running worker, directly and through the frontend. + +For each case it checks that the frontend returns the same status, content +type and body bytes as the worker, then measures warm end-to-end latency on +both paths. Warmup requests are sent first and reported separately. Requests +are sequential (concurrency 1). + + python recipe/cua_s1/bench_text.py --direct http://127.0.0.1:8000 \ + --frontend http://127.0.0.1:8080 --warmup 3 --repeat 20 --out results.json +""" + +from __future__ import annotations + +import argparse +import json +import statistics +import time +import urllib.error +import urllib.request +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[2] +# Ignore http_proxy and friends: the worker and the frontend are local. +OPENER = urllib.request.build_opener(urllib.request.ProxyHandler({})) + + +def post(url: str, body: bytes, token: str | None) -> tuple[int, str, bytes, float]: + headers = {"Content-Type": "application/json"} + if token: + headers["Authorization"] = f"Bearer {token}" + request = urllib.request.Request(url + "/v1/systemone", data=body, headers=headers) + started = time.perf_counter() + try: + with OPENER.open(request, timeout=120) as response: + data = response.read() + status, ctype = response.status, response.headers.get("content-type", "") + except urllib.error.HTTPError as error: + data, status, ctype = ( + error.read(), + error.code, + error.headers.get("content-type", ""), + ) + return status, ctype, data, (time.perf_counter() - started) * 1000 + + +def pct(values: list[float], q: float) -> float: + ordered = sorted(values) + return ordered[min(len(ordered) - 1, round(q * (len(ordered) - 1)))] + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("--direct", required=True, help="worker base URL") + parser.add_argument( + "--frontend", help="frontend base URL; omit to measure the worker only" + ) + parser.add_argument( + "--inputs", default=str(ROOT / "tests/cua_s1/data/text_inputs.json") + ) + parser.add_argument("--warmup", type=int, default=3) + parser.add_argument("--repeat", type=int, default=20) + parser.add_argument("--token", help="bearer token, if the worker requires one") + parser.add_argument("--out") + args = parser.parse_args() + + cases = json.loads(Path(args.inputs).read_text(encoding="utf-8")) + paths = {"direct": args.direct} + if args.frontend: + paths["frontend"] = args.frontend + results, mismatches = {}, 0 + for name, body in cases.items(): + raw = json.dumps(body, ensure_ascii=False).encode() + status, ctype, direct_body, _ = post(args.direct, raw, args.token) + row = {"status": status, "content_type": ctype} + if status == 200: + reply = json.loads(direct_body) + row["answers"], row["input_tokens"] = ( + reply["answers"], + reply["usage"]["input_tokens"], + ) + else: + row["body"] = direct_body.decode("utf-8", "replace")[:500] + if args.frontend: + f_status, f_ctype, f_body, _ = post(args.frontend, raw, args.token) + row["frontend_identical"] = (f_status, f_ctype, f_body) == ( + status, + ctype, + direct_body, + ) + mismatches += not row["frontend_identical"] + for label, url in paths.items(): + warm = [post(url, raw, args.token)[3] for _ in range(args.warmup)] + times = [post(url, raw, args.token)[3] for _ in range(args.repeat)] + row[label] = { + "warmup_ms": [round(t, 2) for t in warm], + "p50_ms": round(statistics.median(times), 2), + "p95_ms": round(pct(times, 0.95), 2), + "min_ms": round(min(times), 2), + "raw_ms": [round(t, 2) for t in times], + } + results[name] = row + line = f"{name}: status {status}, tokens {row.get('input_tokens')}" + for label in paths: + line += ( + f", {label} p50 {row[label]['p50_ms']} ms p95 {row[label]['p95_ms']} ms" + ) + if args.frontend: + line += f", identical {row['frontend_identical']}" + print(line, flush=True) + if args.out: + Path(args.out).write_text( + json.dumps(results, ensure_ascii=False, indent=1) + "\n" + ) + if args.frontend: + print( + f"{len(cases) - mismatches}/{len(cases)} cases byte-identical through the frontend" + ) + return 1 if mismatches else 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/recipe/cua_s1/compare_text_with_upstream.py b/recipe/cua_s1/compare_text_with_upstream.py new file mode 100644 index 00000000..8a50a546 --- /dev/null +++ b/recipe/cua_s1/compare_text_with_upstream.py @@ -0,0 +1,204 @@ +"""Compare the Cua-S1 worker with upstream `FourBModel` on the fixed input set. + +Needs a checkout of trycua/cua at the pinned commit (for `cua_s1.four_b` and +the jev-use chooser) and the pinned weights. The worker's model scores every +question first and is freed; then `FourBModel` is loaded with the same device +and dtype and scores the same questions. For every question it checks: + +- prompt token ids: worker vs upstream `build_prompt` plus the chat template; +- probabilities: worker vs `FourBModel.forward`, exact fp32 equality; +- for the two upstream fixtures, also worker vs the chooser's own path + (`S1DecisionModel.score`), matched by option key. + +Upstream `build_prompt` is given the worker's mapped labels, state and goal, +so the id check covers the prompt layout, chat template and tokenizer. The +request mapping itself (escaping, structured values, `null` labels) is +covered by the unit tests and, independently, by the two fixtures. + +Only one model is resident at a time. Both compute logits for every prompt +position, so the longest input (15,446 tokens) needs about 8 GB for logits in +bfloat16 and 15 GB in float32, on top of the weights. + +Run from the repository root: + + python recipe/cua_s1/compare_text_with_upstream.py --upstream ../cua \\ + --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text --device cuda +""" + +from __future__ import annotations + +import argparse +import gc +import json +import sys +import time +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(ROOT / "src")) + +FIXTURES = { + "fixture_positive": "jev-choice-request-v1.json", + "fixture_negative": "jev-choice-negative-v1.json", +} + + +def free(device: str) -> None: + import torch + + gc.collect() + if device.startswith("cuda"): + torch.cuda.empty_cache() + + +def peak_gib(device: str) -> float | None: + import torch + + if not device.startswith("cuda"): + return None + peak = torch.cuda.max_memory_allocated() / 2**30 + torch.cuda.reset_peak_memory_stats() + return round(peak, 2) + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument( + "--upstream", required=True, help="trycua/cua checkout at the pinned commit" + ) + parser.add_argument("--base", required=True) + parser.add_argument("--adapter", required=True) + parser.add_argument("--device", default="cuda") + parser.add_argument("--dtype", default="bfloat16") + parser.add_argument( + "--inputs", default=str(ROOT / "tests/cua_s1/data/text_inputs.json") + ) + parser.add_argument("--out", help="write one JSON line per question here") + parser.add_argument( + "--no-tf32", + action="store_true", + help="disable TF32 in cuBLAS and cuDNN (use for fp32 reference runs)", + ) + args = parser.parse_args() + + upstream = Path(args.upstream) + sys.path.insert(0, str(upstream / "libs/cua-s1/python/src")) + sys.path.insert(0, str(upstream / "libs/cua-driver/examples/jev-use/python")) + import torch + + if args.no_tf32: + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + from cua_s1.four_b import FourBModel, Option, assign_letters, build_prompt + from decision_models import DecisionRequest, S1DecisionModel + + from models.cua_s1.text.contract import ( + ACTION, + APP, + ROLE, + TASK_FAMILY, + map_request, + parse_body, + ) + from models.cua_s1.text.engine import TextEngine + + cases = json.loads(Path(args.inputs).read_text(encoding="utf-8")) + questions = [] + for name, body in cases.items(): + request = map_request(parse_body(json.dumps(body).encode())) + questions += [(name, request, question) for question in request.questions] + + # Pass 1: the worker. + engine = TextEngine(args.base, args.adapter, args.device, args.dtype) + print(f"worker loaded in {engine.load_seconds:.1f} s", flush=True) + worker = {} + for name, request, question in questions: + worker[name, question.name] = ( + engine.prompt_ids(request.state, question), + engine.score(request.state, question).probabilities, + ) + worker_peak = peak_gib(args.device) + del engine + free(args.device) + + # Pass 2: upstream FourBModel, and the chooser for the two fixtures. + started = time.perf_counter() + reference = FourBModel( + base_model=args.base, + lora_adapter_path=args.adapter, + device=args.device, + dtype=args.dtype, + modality="text", + ) + reference.load() + print(f"upstream loaded in {time.perf_counter() - started:.1f} s", flush=True) + fixture_dir = upstream / "libs/cua-driver/examples/jev-use/fixtures" + rows, failures = [], 0 + for name, request, question in questions: + options = [ + Option(element_id=k, role=ROLE, label=label, action=ACTION) + for k, label in zip(question.keys, question.labels, strict=True) + ] + kwargs = dict( + app=APP, + task_family=TASK_FAMILY, + ax_tree=request.state, + modality="text", + goal=question.goal or None, + ) + messages = build_prompt(assign_letters(options), **kwargs) + chat = reference._tokenizer.apply_chat_template( + messages, tokenize=False, add_generation_prompt=True + ) + upstream_ids = reference._tokenizer(chat)["input_ids"] + upstream_p = [r.probability for r in reference.forward(options, **kwargs)] + worker_ids, worker_p = worker[name, question.name] + row = { + "case": name, + "question": question.name, + "options": len(options), + "prompt_tokens": len(worker_ids), + "ids_equal": worker_ids == upstream_ids, + "probs_equal": worker_p == upstream_p, + "max_abs_diff": max( + abs(a - b) for a, b in zip(worker_p, upstream_p, strict=True) + ), + "worker": dict(zip(question.keys, worker_p, strict=True)), + "upstream": dict(zip(question.keys, upstream_p, strict=True)), + } + if name in FIXTURES: + raw = json.loads((fixture_dir / FIXTURES[name]).read_text(encoding="utf-8")) + chooser = ( + S1DecisionModel(reference, modality="text") + .score(DecisionRequest.from_validated(raw)) + .probabilities + ) + row["chooser_equal"] = all( + chooser.get(k) == p for k, p in row["worker"].items() + ) + ok = row["ids_equal"] and row["probs_equal"] and row.get("chooser_equal", True) + failures += not ok + rows.append(row) + print( + f"{'ok ' if ok else 'FAIL'} {name}/{question.name}: {len(options)} options, " + f"{len(worker_ids)} tokens, max |diff| {row['max_abs_diff']:.3g}", + flush=True, + ) + upstream_peak = peak_gib(args.device) + + if args.out: + with open(args.out, "w", encoding="utf-8") as f: + for row in rows: + f.write(json.dumps(row, ensure_ascii=False) + "\n") + if worker_peak is not None: + print(f"peak allocated: worker {worker_peak} GiB, upstream {upstream_peak} GiB") + print( + f"{len(rows) - failures}/{len(rows)} questions identical " + f"(device {args.device}, dtype {args.dtype}, torch {torch.__version__}, " + f"tf32 {'off' if args.no_tf32 else 'default'})" + ) + return 1 if failures else 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/recipe/cua_s1/requirements-text.txt b/recipe/cua_s1/requirements-text.txt new file mode 100644 index 00000000..cde47c55 --- /dev/null +++ b/recipe/cua_s1/requirements-text.txt @@ -0,0 +1,13 @@ +# Versions match upstream's `four-b` lock (trycua/cua libs/cua-s1/python/uv.lock +# at 0e75660ce4c2edda519e0c795fa3ad98abf4e76f), which the parity checks rely on. +torch==2.14.0 +transformers==5.17.0 +tokenizers==0.23.2 +peft==0.21.0 +accelerate==1.15.0 +safetensors==0.8.0 +huggingface-hub==1.32.0 +jinja2==3.1.6 +# HTTP serving. +fastapi==0.141.1 +uvicorn==0.54.0 diff --git a/recipe/cua_s1/text.md b/recipe/cua_s1/text.md new file mode 100644 index 00000000..90dd738b --- /dev/null +++ b/recipe/cua_s1/text.md @@ -0,0 +1,99 @@ +# Cua-S1 4B 0.2 text worker + +This recipe runs the Cua-S1 4B 0.2 `text` adapter behind the Rust frontend. The worker lives in [`src/models/cua_s1/text/`](../../src/models/cua_s1/text/), and [`src/models/cua_s1/README.md`](../../src/models/cua_s1/README.md) documents the inference contract and the request mapping. Only `choice` questions are supported. + +Run all commands from the repository root, on Linux with an NVIDIA GPU. + +## Install + +Use Python 3.12. The pinned versions match the upstream reference environment: + +```sh +python3.12 -m venv .venv +.venv/bin/python -m pip install -r recipe/cua_s1/requirements-text.txt +``` + +## Download the weights + +Download the pinned revisions (about 9.5 GB) into `weights/`: + +```sh +.venv/bin/hf download Qwen/Qwen3.5-4B \ + --revision 851bf6e806efd8d0a36b00ddf55e13ccb7b8cd0a --local-dir weights/Qwen3.5-4B +.venv/bin/hf download cua-ai/cua-s1-4b-0.2 \ + --revision 16818868b0cc7813808aae4e87b417657046ab79 --local-dir weights/cua-s1-4b-0.2 +``` + +To verify every file against upstream's lock, clone [trycua/cua](https://github.com/trycua/cua) next to this repository, check out `0e75660ce4c2edda519e0c795fa3ad98abf4e76f`, and run: + +```sh +.venv/bin/python ../cua/libs/cua-s1/ci/fetch_pinned_weights.py --dest weights --verify-only +``` + +## Start the worker + +```sh +PYTHONPATH=src .venv/bin/python -m models.cua_s1.text.server \ + --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text \ + --device cuda --dtype bfloat16 --host 127.0.0.1 --port 8000 +``` + +The worker loads the model and runs one warmup decision before it starts listening, so `GET /health` answers only once requests can be served; it then returns `{"status": "ready", "modality": "text", ...}`. The log reports load and warmup times separately. The worker refuses the `multimodal/` adapter and reports the adapter revision that `hf download` recorded. Every flag can also be set through an environment variable: `CUA_S1_BASE`, `CUA_S1_ADAPTER`, `CUA_S1_ADAPTER_REVISION`, `CUA_S1_DEVICE`, `CUA_S1_DTYPE`, `CUA_S1_HOST`, `CUA_S1_PORT`, `CUA_S1_MAX_BODY_BYTES`, `CUA_S1_MAX_QUESTIONS` and `CUA_S1_MAX_PROMPT_TOKENS`. `--adapter-revision` only sets the revision reported in `model` when the download metadata is missing; a value that contradicts the metadata stops the worker. Set `CUA_S1_API_KEY` to require `Authorization: Bearer ` on `/v1/systemone`. + +Oversized requests get `413`: bodies over 4 MiB, more than 64 questions, or a question whose prompt is over 16,384 tokens (`--max-body-bytes`, `--max-questions`, `--max-prompt-tokens`). The worker computes logits for every prompt position, as upstream does, so memory grows with prompt length: serving the 15,446-token test input in bfloat16 peaked at about 21.3 GiB in use on the card. Requests run one at a time, and the frontend gives up after 60 seconds. + +## Start the frontend + +The frontend is in [#2](https://github.com/ThinkFlowLab/system1-omni/pull/2), which is not merged yet. Build it from that pull request's branch: + +```sh +git fetch origin pull/2/head:frontend-pr2 +git worktree add ../system1-omni-frontend frontend-pr2 +(cd ../system1-omni-frontend && cargo build --release --locked) +OMNI_JEV_BIND=127.0.0.1:8080 \ +OMNI_JEV_BACKEND_URL=http://127.0.0.1:8000 \ + ../system1-omni-frontend/target/release/omni-jev +``` + +## Send a request + +```sh +curl http://127.0.0.1:8080/health +curl http://127.0.0.1:8080/v1/systemone \ + -H 'Content-Type: application/json' \ + -d '{"model":"cua-s1-4b-0.2","state":"Dialog: Delete 3 files permanently? Buttons: Delete, Cancel","questions":{"pick":{"type":"choice","instructions":"Keep the files.","criteria":{"delete":"Click Delete","cancel":"Click Cancel"}}}}' +``` + +The answer has the Jev choice shape. On an RTX 6000 Ada in bfloat16, the response is: + +```json +{"model":"cua-ai/cua-s1-4b-0.2@16818868b0cc7813808aae4e87b417657046ab79:text","answers":{"pick":{"type":"choice","choice":"cancel","probabilities":{"delete":0.0024726232513785362,"cancel":0.9975274205207825},"confidence":0.9750249565060322}},"usage":{"input_tokens":153,"output_tokens":0}} +``` + +## Check against upstream + +`compare_text_with_upstream.py` scores the fixed input set (`tests/cua_s1/data/text_inputs.json`) with the worker's model and then with upstream `FourBModel`, one model at a time, and compares the results. It needs the trycua/cua checkout from above and two extra packages for upstream's processor: + +```sh +.venv/bin/python -m pip install torchvision==0.29.0 pillow==11.3.0 +.venv/bin/python recipe/cua_s1/compare_text_with_upstream.py --upstream ../cua \ + --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text --device cuda +``` + +Every question must have identical prompt token ids and identical fp32 probabilities. + +With the worker and the frontend running, `bench_text.py` checks that the frontend returns the same bytes as the worker for every input, then measures warm latency on both paths: + +```sh +.venv/bin/python recipe/cua_s1/bench_text.py --direct http://127.0.0.1:8000 \ + --frontend http://127.0.0.1:8080 --warmup 3 --repeat 20 +``` + +## Tests + +The contract and HTTP tests need neither weights nor a GPU. The tokenizer tests also run when `CUA_S1_BASE` points to the downloaded base model; they read only its tokenizer files: + +```sh +.venv/bin/python -m pip install pytest +CUA_S1_BASE=weights/Qwen3.5-4B PYTHONPATH=src .venv/bin/python -m pytest tests/cua_s1/test_text_*.py +``` diff --git a/src/models/cua_s1/README.md b/src/models/cua_s1/README.md index 6a072600..f0712d26 100644 --- a/src/models/cua_s1/README.md +++ b/src/models/cua_s1/README.md @@ -2,7 +2,15 @@ This directory owns Cua-S1 4B 0.2 ([#10](https://github.com/ThinkFlowLab/system1-omni/issues/10)): request mapping, prompt construction, adapter selection, execution, and the answer-letter readout. This page records the pinned upstream revisions, the inference contract an implementation must match, and how its outputs will be compared with the upstream reference. -Status: planned; nothing is implemented or validated yet. The first target is the `text` adapter on CUDA, starting with a worker that loads the model directly through Hugging Face Transformers and PEFT. The `multimodal` adapter is deferred; see [Not covered yet](#not-covered-yet). +Status: a reference worker for the `text` adapter is in [`text/`](text/). It loads the model directly through Hugging Face Transformers and PEFT and serves `/v1/systemone`; setup and checks are in [`recipe/cua_s1/text.md`](../../../recipe/cua_s1/text.md). A worker for the `multimodal` adapter is proposed in [#12](https://github.com/ThinkFlowLab/system1-omni/pull/12). + +| Path | Contents | +| --- | --- | +| `text/contract.py` | Request validation, the `/v1/systemone` mapping, prompt construction and answers. No torch imports. | +| `text/engine.py` | Model and adapter loading and the answer-letter readout. | +| `text/server.py` | The HTTP worker (`GET /health`, `POST /v1/systemone`). | +| `text/adapter.py` | Finds and checks the local `text` adapter and its downloaded revision. | +| `tests/cua_s1/test_text_*.py` (repository root) | Tests that need neither weights nor a GPU, and tokenizer checks. The fixed input set is `tests/cua_s1/data/text_inputs.json`. | ## Pinned revisions @@ -76,19 +84,19 @@ The response `model` is `cua-ai/cua-s1-4b-0.2@:`, in An error rejects the whole request. Its body is `{"detail": ""}`, as the LAYA worker returns, and the message names the problem. -The status is `400` when the body is not a usable JSON object: invalid JSON or UTF-8, `NaN` or `Infinity`, a lone surrogate such as `\ud800`, nesting too deep to parse, or a key repeated in any object. +The status is `400` when the body is not a usable JSON object: invalid JSON or UTF-8, `NaN`, `Infinity` or a number out of range, a lone surrogate such as `\ud800`, nesting too deep to parse, or a key repeated in any object. The status is `422` when a well-formed request cannot be answered: - a `score` or `noul` question, since the adapters were trained only on closed-option choices; - a question without an `instructions` field (`null` is allowed), or a `choice` with no options or more than 26 options; - a `criteria` value that is a number or a boolean; -- an empty `state`; +- an empty `state` (`""`, `{}` or `[]`); - a `model` other than `cua-s1-4b-0.2`. ## Validation -**Inputs.** The fixed input set is upstream's two checked-in fixtures, converted to `/v1/systemone` requests with the chooser's rendered regions as `state`, plus `/v1/systemone` choice requests that will be checked in with the worker. These cover 1 to 26 options, short and long states, string and structured `state`, `instructions` and `criteria`, `null` criteria, non-ASCII text, and text that spells a special token. Each input is scored once per configuration. +**Inputs.** The fixed input set is upstream's two checked-in fixtures, converted to `/v1/systemone` requests with the chooser's rendered regions as `state`, plus `/v1/systemone` choice requests, in `tests/cua_s1/data/text_inputs.json`. These cover 1 to 26 options, short and long states, string and structured `state`, `instructions` and `criteria`, `null` criteria, non-ASCII text, and text that spells a special token. Each input is scored once per configuration. **Tolerances.** These are declared before any comparison is run: diff --git a/src/models/cua_s1/text/THIRD_PARTY_NOTICES.md b/src/models/cua_s1/text/THIRD_PARTY_NOTICES.md new file mode 100644 index 00000000..f054d433 --- /dev/null +++ b/src/models/cua_s1/text/THIRD_PARTY_NOTICES.md @@ -0,0 +1,23 @@ +The system message, prompt layout and fixed values in `contract.py`, and the two upstream fixtures converted in `tests/cua_s1/data/text_inputs.json`, come from [trycua/cua](https://github.com/trycua/cua) at `0e75660ce4c2edda519e0c795fa3ad98abf4e76f` under the following license. + +MIT License + +Copyright (c) 2025 Cua AI, Inc. + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/src/models/cua_s1/text/adapter.py b/src/models/cua_s1/text/adapter.py new file mode 100644 index 00000000..472e533a --- /dev/null +++ b/src/models/cua_s1/text/adapter.py @@ -0,0 +1,51 @@ +"""Locate and check the local Cua-S1 `text` adapter. No torch imports.""" + +from __future__ import annotations + +import json +import re +from pathlib import Path + + +def text_adapter_dir(adapter_root: str | Path) -> Path: + """Return the `text` adapter directory under the adapter root. + + Accepts the repository root (`/text`) or the `text/` directory + itself, and refuses the `multimodal/` adapter: PEFT only warns about keys + it cannot place, so loading the wrong adapter would otherwise go unnoticed. + """ + root = Path(adapter_root) + path = root / "text" if (root / "text" / "adapter_config.json").exists() else root + config_file = path / "adapter_config.json" + if not config_file.exists(): + raise RuntimeError(f"no adapter_config.json under {root}") + config = json.loads(config_file.read_text()) + if config.get("base_model_name_or_path") != "Qwen/Qwen3.5-4B": + raise RuntimeError(f"{config_file}: base model is not Qwen/Qwen3.5-4B") + if {"linear_fc1", "linear_fc2"} & set(config.get("target_modules") or []): + raise RuntimeError( + f"{config_file}: this is the multimodal adapter, not the text adapter" + ) + return path + + +def downloaded_revision(adapter_root: str | Path) -> str | None: + """The commit that `hf download --local-dir` recorded for the text adapter, if any. + + `hf download` keeps its metadata under the repository root, so this also + looks one level up when `adapter_root` is the `text/` directory itself. + """ + root = Path(adapter_root) + places = [(root, "text/"), (root, "")] + if root.name == "text": + places.insert(0, (root.parent, "text/")) + for base, prefix in places: + cache = base / ".cache" / "huggingface" / "download" + try: + meta = (cache / f"{prefix}adapter_model.safetensors.metadata").read_text() + first = meta.splitlines()[0].strip() + except (OSError, IndexError): + continue + if re.fullmatch(r"[0-9a-f]{40}", first): + return first + return None diff --git a/src/models/cua_s1/text/contract.py b/src/models/cua_s1/text/contract.py new file mode 100644 index 00000000..a0a2d1d3 --- /dev/null +++ b/src/models/cua_s1/text/contract.py @@ -0,0 +1,273 @@ +"""Request mapping, prompt construction and answers for Cua-S1 4B 0.2. + +This module follows the contract in `src/models/cua_s1/README.md`. It has no +torch or Transformers imports, so it can be tested without weights. +""" + +from __future__ import annotations + +import json +import math +import string +from dataclasses import dataclass +from typing import Any + +MODEL_NAME = "cua-s1-4b-0.2" +ADAPTER_REPO = "cua-ai/cua-s1-4b-0.2" +ADAPTER_REVISION = "16818868b0cc7813808aae4e87b417657046ab79" +BASE_REPO = "Qwen/Qwen3.5-4B" +BASE_REVISION = "851bf6e806efd8d0a36b00ddf55e13ccb7b8cd0a" + +LETTERS = string.ascii_uppercase +MAX_OPTIONS = len(LETTERS) + +# The system message, the user message layout and the fixed values below are +# copied from trycua/cua at 0e75660ce4c2edda519e0c795fa3ad98abf4e76f: +# `libs/cua-s1/python/src/cua_s1/four_b.py` (SYSTEM_PROMPT, build_prompt, +# _describe_option) and `libs/cua-driver/examples/jev-use/python/ +# decision_models.py` (S1DecisionModel.score). MIT License, Copyright (c) 2025 +# Cua AI, Inc.; see THIRD_PARTY_NOTICES.md. +SYSTEM_PROMPT = ( + "You are a one-pass computer-use decision model. You are shown the " + "current state of a screen and a fixed, closed list of candidate " + "(element, action) options, each given a single letter. Choose exactly " + "one option: the single best next action to take. Answer with ONLY that " + "option's letter -- no words, no punctuation, no explanation." +) +APP = "Cua Driver" +TASK_FAMILY = "closed-candidate decision" +ROLE = "Decision" +ACTION = "select" + + +class RequestError(ValueError): + """A request the worker rejects. `status` is the HTTP status to return.""" + + def __init__(self, message: str, status: int = 422) -> None: + super().__init__(message) + self.status = status + + +def _object_pairs(pairs: list[tuple[str, Any]]) -> dict[str, Any]: + obj: dict[str, Any] = {} + for key, value in pairs: + if key in obj: + raise RequestError(f"duplicate key {key!r} in a JSON object", status=400) + obj[key] = value + return obj + + +def _reject_constant(name: str) -> Any: + raise RequestError(f"{name} is not valid JSON", status=400) + + +def _finite_float(text: str) -> float: + value = float(text) + if not math.isfinite(value): + raise RequestError(f"number {text} is out of range", status=400) + return value + + +def _check_unicode(value: Any) -> None: + """Reject lone surrogates (for example a `\\ud800` escape): they cannot be + encoded as UTF-8, so they cannot be tokenized or echoed back.""" + if isinstance(value, str): + value.encode("utf-8") + elif isinstance(value, dict): + for key, item in value.items(): + key.encode("utf-8") + _check_unicode(item) + elif isinstance(value, list): + for item in value: + _check_unicode(item) + + +def parse_body(raw: bytes) -> dict[str, Any]: + """Decode a request body, keeping key order and rejecting duplicate keys.""" + try: + text = raw.decode("utf-8") + body = json.loads( + text, + object_pairs_hook=_object_pairs, + parse_constant=_reject_constant, + parse_float=_finite_float, + ) + _check_unicode(body) + except RequestError: + raise + except RecursionError as error: + raise RequestError("request body is nested too deeply", status=400) from error + except UnicodeError as error: + raise RequestError( + "request body must be valid UTF-8 text", status=400 + ) from error + except ValueError as error: + # json.JSONDecodeError, a UTF-8 byte order mark, or an integer too long + # for Python to convert. + raise RequestError("request body must be valid JSON", status=400) from error + if not isinstance(body, dict): + raise RequestError("request body must be a JSON object", status=400) + return body + + +def as_text(value: Any) -> str: + """Render `state` or `instructions` as prompt text. + + A string is used as is; an object or array is serialized the way Python's + `json.dumps(value, ensure_ascii=False)` does. + """ + if isinstance(value, str): + return value + return json.dumps(value, ensure_ascii=False) + + +def escape_label(value: str) -> str: + """Escape an option label the way upstream's chooser does.""" + return json.dumps(value, ensure_ascii=False)[1:-1] + + +@dataclass(frozen=True) +class Question: + """One `choice` question mapped onto the prompt fields.""" + + name: str + goal: str + keys: tuple[str, ...] + labels: tuple[str, ...] + + +@dataclass(frozen=True) +class Request: + state: str + questions: tuple[Question, ...] + + +def _check_json_value(value: Any, where: str, allow_null: bool) -> None: + if value is None: + if not allow_null: + raise RequestError(f"{where} must not be null") + return + if isinstance(value, bool) or isinstance(value, (int, float)): + raise RequestError(f"{where} must be a string, an object or an array") + if not isinstance(value, (str, dict, list)): + raise RequestError(f"{where} must be a string, an object or an array") + + +def map_request(body: dict[str, Any], *, max_questions: int = 64) -> Request: + """Validate a `/v1/systemone` body and map it onto prompt fields.""" + model = body.get("model") + if model != MODEL_NAME: + raise RequestError(f"'model' must be {MODEL_NAME!r}") + + if "state" not in body: + raise RequestError("'state' is required") + state_value = body["state"] + _check_json_value(state_value, "'state'", allow_null=False) + if state_value in ("", {}, []): + raise RequestError("'state' must not be empty") + state = as_text(state_value) + + questions = body.get("questions") + if not isinstance(questions, dict) or not questions: + raise RequestError("'questions' must be a non-empty object") + if len(questions) > max_questions: + raise RequestError( + f"too many questions ({len(questions)} > {max_questions})", status=413 + ) + + # Check every question type before the per-question checks, so a `score` + # or `noul` question anywhere rejects the whole request with that reason. + for name, question in questions.items(): + if not isinstance(question, dict): + raise RequestError(f"question {name!r} must be an object") + kind = question.get("type") + if kind in ("score", "noul"): + raise RequestError( + f"question {name!r}: type {kind!r} is not supported; " + "Cua-S1 4B 0.2 answers 'choice' questions only" + ) + if kind != "choice": + raise RequestError(f"question {name!r}: unknown type {kind!r}") + + mapped = [] + for name, question in questions.items(): + where = f"question {name!r}" + if "instructions" not in question: + raise RequestError(f"{where}: 'instructions' is required") + instructions = question["instructions"] + _check_json_value(instructions, f"{where}: 'instructions'", allow_null=True) + goal = "" if instructions is None else as_text(instructions) + + criteria = question.get("criteria") + if not isinstance(criteria, dict): + raise RequestError(f"{where}: 'criteria' must be an object") + if not criteria: + raise RequestError(f"{where}: 'criteria' must have at least one option") + if len(criteria) > MAX_OPTIONS: + raise RequestError( + f"{where}: {len(criteria)} options; at most {MAX_OPTIONS} are supported" + ) + keys, labels = [], [] + for key, value in criteria.items(): + _check_json_value(value, f"{where}: option {key!r}", allow_null=True) + if value is None: + text = key + else: + text = as_text(value) + keys.append(key) + labels.append(escape_label(text)) + mapped.append( + Question(name=name, goal=goal, keys=tuple(keys), labels=tuple(labels)) + ) + return Request(state=state, questions=tuple(mapped)) + + +def build_messages(state: str, question: Question) -> list[dict[str, str]]: + """Chat messages for one question, matching upstream `build_prompt` (text).""" + option_lines = "\n".join( + f'{letter}. {ROLE} "{label}" -> {ACTION}' + for letter, label in zip(LETTERS, question.labels, strict=False) + ) + user = ( + (f"Goal: {question.goal}\n\n" if question.goal else "") + + f"App: {APP}\nTask family: {TASK_FAMILY}\n\n" + + f"Accessibility tree:\n{state}\n\n" + + f"Options:\n{option_lines}\n\nAnswer with a single letter." + ) + return [ + {"role": "system", "content": SYSTEM_PROMPT}, + {"role": "user", "content": user}, + ] + + +def confidence(probabilities: list[float]) -> float: + """Normalized entropy, `1 - H(p) / ln(n)`, as the LAYA worker reports it.""" + n = len(probabilities) + if n < 2: + return 1.0 + entropy = -sum(p * math.log(min(max(p, 1e-12), 1.0)) for p in probabilities) + return min(max(1.0 - entropy / math.log(n), 0.0), 1.0) + + +def answer(question: Question, probabilities: list[float]) -> dict[str, Any]: + """The Jev choice answer. Ties go to the earliest option.""" + if len(probabilities) != len(question.keys) or not all( + math.isfinite(p) and 0.0 <= p <= 1.0 for p in probabilities + ): + raise ValueError(f"model returned invalid probabilities: {probabilities}") + if not math.isclose(sum(probabilities), 1.0, abs_tol=1e-5): + raise ValueError(f"model probabilities do not sum to one: {probabilities}") + best = 0 + for index, p in enumerate(probabilities): + if p > probabilities[best]: + best = index + return { + "type": "choice", + "choice": question.keys[best], + "probabilities": dict(zip(question.keys, probabilities, strict=True)), + "confidence": confidence(probabilities), + } + + +def model_identity(revision: str = ADAPTER_REVISION, modality: str = "text") -> str: + return f"{ADAPTER_REPO}@{revision}:{modality}" diff --git a/src/models/cua_s1/text/engine.py b/src/models/cua_s1/text/engine.py new file mode 100644 index 00000000..cdcc351d --- /dev/null +++ b/src/models/cua_s1/text/engine.py @@ -0,0 +1,84 @@ +"""Load Qwen3.5-4B with the Cua-S1 `text` adapter and score one prompt. + +The calls mirror upstream `cua_s1.four_b.FourBModel` (text modality): the same +model class, an unmerged PEFT adapter, the chat template with its default +generation prompt, full logits, and a fp32 softmax over the letter logits at +the last position. Keeping them the same is what makes the worker's +probabilities bitwise identical to the reference in the same environment. +""" + +from __future__ import annotations + +import time +from dataclasses import dataclass + +import torch +from peft import PeftModel +from transformers import AutoModelForCausalLM, AutoTokenizer + +from .adapter import text_adapter_dir +from .contract import LETTERS, Question, build_messages + + +@dataclass +class Scored: + probabilities: list[float] + prompt_tokens: int + + +class TextEngine: + def __init__( + self, base_model: str, adapter_root: str, device: str, dtype: str + ) -> None: + self.device = device + self.dtype = dtype + started = time.perf_counter() + self.tokenizer = AutoTokenizer.from_pretrained(base_model) + model = AutoModelForCausalLM.from_pretrained( + base_model, dtype=getattr(torch, dtype), device_map=device + ) + model = PeftModel.from_pretrained(model, str(text_adapter_dir(adapter_root))) + model.eval() + self.model = model + self.load_seconds = time.perf_counter() - started + self.letter_ids = self._letter_ids() + + def _letter_ids(self) -> list[int]: + ids = [] + for letter in LETTERS: + tokens = self.tokenizer.encode(letter, add_special_tokens=False) + if len(tokens) != 1: + raise RuntimeError(f"letter {letter!r} is not a single token: {tokens}") + ids.append(tokens[0]) + return ids + + def encode(self, state: str, question: Question): + """Tokenized prompt for one question, on CPU. + + The Qwen3.5 tokenizer adds no special tokens here (contract point 4); + the chat template already contains them. + """ + chat_text = self.tokenizer.apply_chat_template( + build_messages(state, question), tokenize=False, add_generation_prompt=True + ) + return self.tokenizer(chat_text, return_tensors="pt") + + def prompt_ids(self, state: str, question: Question) -> list[int]: + return self.encode(state, question)["input_ids"][0].tolist() + + @torch.no_grad() + def score_encoded(self, inputs, n_options: int) -> Scored: + inputs = inputs.to(self.model.device) + out = self.model(**inputs) + final_logits = out.logits[0, -1, :] + letter_ids = self.letter_ids[:n_options] + option_logits = final_logits[ + torch.tensor(letter_ids, device=final_logits.device) + ] + probabilities = torch.softmax(option_logits.float(), dim=-1).tolist() + return Scored( + probabilities=probabilities, prompt_tokens=int(inputs["input_ids"].shape[1]) + ) + + def score(self, state: str, question: Question) -> Scored: + return self.score_encoded(self.encode(state, question), len(question.keys)) diff --git a/src/models/cua_s1/text/server.py b/src/models/cua_s1/text/server.py new file mode 100644 index 00000000..a272eec5 --- /dev/null +++ b/src/models/cua_s1/text/server.py @@ -0,0 +1,234 @@ +"""HTTP worker for Cua-S1 4B 0.2 (`text` adapter) behind the Rust frontend. + +Routes: `GET /health` and `POST /v1/systemone`. The model is loaded before the +server starts listening, and one forward pass runs at a time. + + PYTHONPATH=src python -m models.cua_s1.text.server --base --adapter +""" + +from __future__ import annotations + +import argparse +import asyncio +import hmac +import json +import os +import sys +import time +import traceback +from concurrent.futures import ThreadPoolExecutor +from typing import Any + +# FastAPI reads the handler annotations at runtime, so `Request` must be a +# module-level name while `from __future__ import annotations` is in effect. +from fastapi import FastAPI, Request +from fastapi.responses import JSONResponse + +from .contract import ( + ADAPTER_REVISION, + MODEL_NAME, + RequestError, + answer, + map_request, + model_identity, + parse_body, +) + +WARMUP_REQUEST = { + "model": MODEL_NAME, + "state": "Dialog: 'Update installed.' Button: OK", + "questions": { + "warmup": { + "type": "choice", + "instructions": "Close the dialog.", + "criteria": {"ok": "Click OK", "wait": "Wait"}, + } + }, +} + + +def build_app( + engine: Any, + *, + api_key: str | None, + max_body_bytes: int, + max_questions: int, + max_prompt_tokens: int, + revision: str, +): + app = FastAPI() + pool = ThreadPoolExecutor(max_workers=1) + identity = model_identity(revision) + expected_auth = ( + f"Bearer {api_key}".encode("utf-8", "surrogateescape") if api_key else b"" + ) + + def error(status: int, message: str) -> JSONResponse: + return JSONResponse({"detail": message}, status_code=status) + + def authorized(request: Request) -> bool: + if not api_key: + return True + supplied = request.headers.get("authorization", "").encode( + "utf-8", "surrogateescape" + ) + return hmac.compare_digest(supplied, expected_auth) + + @app.get("/health") + def health(): + return { + "status": "ready", + "modality": "text", + "model": identity, + "device": engine.device, + "dtype": engine.dtype, + } + + def decide(mapped): + # Tokenize every question first, so an over-long prompt is rejected + # before any forward pass runs. + encoded = [] + for question in mapped.questions: + inputs = engine.encode(mapped.state, question) + n = int(inputs["input_ids"].shape[1]) + if max_prompt_tokens and n > max_prompt_tokens: + raise RequestError( + f"question {question.name!r}: prompt is {n} tokens, " + f"over the {max_prompt_tokens}-token limit", + status=413, + ) + encoded.append((question, inputs)) + answers, prompt_tokens = {}, 0 + for question, inputs in encoded: + scored = engine.score_encoded(inputs, len(question.keys)) + answers[question.name] = answer(question, scored.probabilities) + prompt_tokens += scored.prompt_tokens + return { + "model": identity, + "answers": answers, + "usage": {"input_tokens": prompt_tokens, "output_tokens": 0}, + } + + @app.post("/v1/systemone") + async def systemone(request: Request): + if not authorized(request): + return error(401, "invalid or missing bearer token") + length = request.headers.get("content-length") + if length and length.isdigit() and int(length) > max_body_bytes: + return error(413, "request body too large") + raw = bytearray() + async for chunk in request.stream(): + raw.extend(chunk) + if len(raw) > max_body_bytes: + return error(413, "request body too large") + try: + mapped = map_request(parse_body(bytes(raw)), max_questions=max_questions) + loop = asyncio.get_running_loop() + return await loop.run_in_executor(pool, decide, mapped) + except RequestError as exc: + return error(exc.status, str(exc)) + except Exception: + traceback.print_exc(file=sys.stderr) + return error(500, "inference failed") + + def warmup() -> None: + """Run one decision on the worker thread through the full request path.""" + mapped = map_request(WARMUP_REQUEST) + json.dumps(pool.submit(decide, mapped).result(), allow_nan=False) + + app.state.warmup = warmup + return app + + +def main(argv: list[str] | None = None) -> None: + env = os.environ.get + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument( + "--base", + default=env("CUA_S1_BASE"), + help="local Qwen/Qwen3.5-4B directory (env CUA_S1_BASE)", + ) + parser.add_argument( + "--adapter", + default=env("CUA_S1_ADAPTER"), + help="local cua-ai/cua-s1-4b-0.2 directory (env CUA_S1_ADAPTER)", + ) + parser.add_argument( + "--adapter-revision", + default=env("CUA_S1_ADAPTER_REVISION"), + help="adapter revision to report when the download metadata is missing", + ) + parser.add_argument("--device", default=env("CUA_S1_DEVICE", "cuda")) + parser.add_argument( + "--dtype", + default=env("CUA_S1_DTYPE", "bfloat16"), + choices=["bfloat16", "float16", "float32"], + ) + parser.add_argument("--host", default=env("CUA_S1_HOST", "127.0.0.1")) + parser.add_argument("--port", type=int, default=int(env("CUA_S1_PORT", "8000"))) + parser.add_argument( + "--max-body-bytes", + type=int, + default=int(env("CUA_S1_MAX_BODY_BYTES", str(4 << 20))), + ) + parser.add_argument( + "--max-questions", type=int, default=int(env("CUA_S1_MAX_QUESTIONS", "64")) + ) + parser.add_argument( + "--max-prompt-tokens", + type=int, + default=int(env("CUA_S1_MAX_PROMPT_TOKENS", "16384")), + help="per question; 0 disables the check", + ) + args = parser.parse_args(argv) + if not args.base or not args.adapter: + parser.error("--base and --adapter are required") + + import uvicorn + + from .adapter import downloaded_revision, text_adapter_dir + from .engine import TextEngine + + # Fail before loading weights if this is not the text adapter. + text_adapter_dir(args.adapter) + detected = downloaded_revision(args.adapter) + if detected and args.adapter_revision and detected != args.adapter_revision: + parser.error( + f"--adapter-revision {args.adapter_revision} does not match the " + f"downloaded revision {detected}" + ) + revision = detected or args.adapter_revision or ADAPTER_REVISION + if revision != ADAPTER_REVISION: + print( + f"warning: adapter revision {revision} is not the pinned {ADAPTER_REVISION}", + flush=True, + ) + if not detected: + print( + "note: no download metadata under --adapter; the adapter revision is not verified", + flush=True, + ) + + engine = TextEngine(args.base, args.adapter, args.device, args.dtype) + print( + f"loaded in {engine.load_seconds:.1f} s on {args.device} ({args.dtype})", + flush=True, + ) + app = build_app( + engine, + api_key=env("CUA_S1_API_KEY") or None, + max_body_bytes=args.max_body_bytes, + max_questions=args.max_questions, + max_prompt_tokens=args.max_prompt_tokens, + revision=revision, + ) + # One decision before listening, so the first real request does not pay + # for lazy weight loading or first-call kernel setup on the worker thread. + started = time.perf_counter() + app.state.warmup() + print(f"warmed up in {time.perf_counter() - started:.1f} s", flush=True) + uvicorn.run(app, host=args.host, port=args.port, log_level="warning") + + +if __name__ == "__main__": + main() diff --git a/tests/cua_s1/data/text_inputs.json b/tests/cua_s1/data/text_inputs.json new file mode 100644 index 00000000..1a56460b --- /dev/null +++ b/tests/cua_s1/data/text_inputs.json @@ -0,0 +1,289 @@ +{ + "fixture_positive": { + "model": "cua-s1-4b-0.2", + "state": "Visual-region-derived observation for capture \"capture-fixture-1\":\n\"submit-text\": text 'Submit' at (300,240,100,40) confidence=0.96 interactive=true", + "questions": { + "pick": { + "type": "choice", + "instructions": "Submit the verified form.", + "criteria": { + "submit-form": "Submit using the unique validated visual region.", + "reobserve": "Discard this decision set and obtain a fresh observation.", + "abstain": "Stop without acting if no supplied action is safe." + } + } + } + }, + "fixture_negative": { + "model": "cua-s1-4b-0.2", + "state": "Visual-region-derived observation for capture \"synthetic-negative-1\":\n\"save\": text 'Save' at (10,10,80,30) confidence=0.98 interactive=true", + "questions": { + "pick": { + "type": "choice", + "instructions": "Choose exactly one candidate by applying its condition to the supplied state. The host alone authorizes any selected action.", + "criteria": { + "region:save": "Select only if exactly one supplied region has id=save, kind=text, exact_text=\"Send\", confidence=0.98, and it is the sole exact Send match at or above 0.80.", + "reobserve": "Select only when no action candidate condition matches and no regions are supplied. Do not act; request one fresh bounded observation.", + "abstain": "Select only when no action candidate condition matches and one or more regions are supplied. Do not act; stop." + } + } + } + }, + "one_option": { + "model": "cua-s1-4b-0.2", + "state": "Dialog: 'Update installed.' Button: OK", + "questions": { + "pick": { + "type": "choice", + "instructions": "Close the dialog.", + "criteria": { + "ok": "Click OK" + } + } + } + }, + "two_options": { + "model": "cua-s1-4b-0.2", + "state": "Dialog: 'Delete 3 files permanently?' Buttons: Delete, Cancel", + "questions": { + "pick": { + "type": "choice", + "instructions": "Keep the files.", + "criteria": { + "delete": "Click Delete", + "cancel": "Click Cancel" + } + } + } + }, + "max_26_options": { + "model": "cua-s1-4b-0.2", + "state": "Toolbar of a document editor. Selected text: 'quarterly results'.\nbutton 'Undo' enabled=true\nbutton 'Redo' enabled=true\nbutton 'Cut' enabled=true\nbutton 'Copy' enabled=true\nbutton 'Paste' enabled=true\nbutton 'Bold' enabled=true\nbutton 'Italic' enabled=true\nbutton 'Underline' enabled=true\nbutton 'Strikethrough' enabled=true\nbutton 'Font color' enabled=true\nbutton 'Highlight' enabled=true\nbutton 'Align left' enabled=true\nbutton 'Center' enabled=true\nbutton 'Align right' enabled=true\nbutton 'Justify' enabled=true\nbutton 'Bullets' enabled=true\nbutton 'Numbering' enabled=true\nbutton 'Indent' enabled=true\nbutton 'Outdent' enabled=true\nbutton 'Insert link' enabled=true\nbutton 'Insert image' enabled=true\nbutton 'Insert table' enabled=true\nbutton 'Comment' enabled=true\nbutton 'Find' enabled=true\nbutton 'Replace' enabled=true\nbutton 'Print' enabled=true", + "questions": { + "pick": { + "type": "choice", + "instructions": "Make the selected text bold.", + "criteria": { + "undo": "Click the 'Undo' button", + "redo": "Click the 'Redo' button", + "cut": "Click the 'Cut' button", + "copy": "Click the 'Copy' button", + "paste": "Click the 'Paste' button", + "bold": "Click the 'Bold' button", + "italic": "Click the 'Italic' button", + "underline": "Click the 'Underline' button", + "strikethrough": "Click the 'Strikethrough' button", + "font-color": "Click the 'Font color' button", + "highlight": "Click the 'Highlight' button", + "align-left": "Click the 'Align left' button", + "center": "Click the 'Center' button", + "align-right": "Click the 'Align right' button", + "justify": "Click the 'Justify' button", + "bullets": "Click the 'Bullets' button", + "numbering": "Click the 'Numbering' button", + "indent": "Click the 'Indent' button", + "outdent": "Click the 'Outdent' button", + "insert-link": "Click the 'Insert link' button", + "insert-image": "Click the 'Insert image' button", + "insert-table": "Click the 'Insert table' button", + "comment": "Click the 'Comment' button", + "find": "Click the 'Find' button", + "replace": "Click the 'Replace' button", + "print": "Click the 'Print' button" + } + } + } + }, + "long_state": { + "model": "cua-s1-4b-0.2", + "state": "Orders table (web admin), 300 rows, sorted by order number.\nrow 0: cell 'Order #10000' | cell 'Customer 0' | cell '10.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 1: cell 'Order #10001' | cell 'Customer 1' | cell '17.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 2: cell 'Order #10002' | cell 'Customer 2' | cell '24.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 3: cell 'Order #10003' | cell 'Customer 3' | cell '31.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 4: cell 'Order #10004' | cell 'Customer 4' | cell '38.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 5: cell 'Order #10005' | cell 'Customer 5' | cell '45.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 6: cell 'Order #10006' | cell 'Customer 6' | cell '52.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 7: cell 'Order #10007' | cell 'Customer 7' | cell '59.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 8: cell 'Order #10008' | cell 'Customer 8' | cell '66.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 9: cell 'Order #10009' | cell 'Customer 9' | cell '73.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 10: cell 'Order #10010' | cell 'Customer 10' | cell '80.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 11: cell 'Order #10011' | cell 'Customer 11' | cell '87.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 12: cell 'Order #10012' | cell 'Customer 12' | cell '94.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 13: cell 'Order #10013' | cell 'Customer 13' | cell '11.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 14: cell 'Order #10014' | cell 'Customer 14' | cell '18.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 15: cell 'Order #10015' | cell 'Customer 15' | cell '25.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 16: cell 'Order #10016' | cell 'Customer 16' | cell '32.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 17: cell 'Order #10017' | cell 'Customer 17' | cell '39.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 18: cell 'Order #10018' | cell 'Customer 18' | cell '46.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 19: cell 'Order #10019' | cell 'Customer 19' | cell '53.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 20: cell 'Order #10020' | cell 'Customer 20' | cell '60.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 21: cell 'Order #10021' | cell 'Customer 21' | cell '67.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 22: cell 'Order #10022' | cell 'Customer 22' | cell '74.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 23: cell 'Order #10023' | cell 'Customer 23' | cell '81.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 24: cell 'Order #10024' | cell 'Customer 24' | cell '88.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 25: cell 'Order #10025' | cell 'Customer 25' | cell '95.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 26: cell 'Order #10026' | cell 'Customer 26' | cell '12.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 27: cell 'Order #10027' | cell 'Customer 27' | cell '19.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 28: cell 'Order #10028' | cell 'Customer 28' | cell '26.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 29: cell 'Order #10029' | cell 'Customer 29' | cell '33.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 30: cell 'Order #10030' | cell 'Customer 30' | cell '40.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 31: cell 'Order #10031' | cell 'Customer 31' | cell '47.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 32: cell 'Order #10032' | cell 'Customer 32' | cell '54.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 33: cell 'Order #10033' | cell 'Customer 33' | cell '61.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 34: cell 'Order #10034' | cell 'Customer 34' | cell '68.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 35: cell 'Order #10035' | cell 'Customer 35' | cell '75.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 36: cell 'Order #10036' | cell 'Customer 36' | cell '82.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 37: cell 'Order #10037' | cell 'Customer 0' | cell '89.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 38: cell 'Order #10038' | cell 'Customer 1' | cell '96.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 39: cell 'Order #10039' | cell 'Customer 2' | cell '13.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 40: cell 'Order #10040' | cell 'Customer 3' | cell '20.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 41: cell 'Order #10041' | cell 'Customer 4' | cell '27.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 42: cell 'Order #10042' | cell 'Customer 5' | cell '34.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 43: cell 'Order #10043' | cell 'Customer 6' | cell '41.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 44: cell 'Order #10044' | cell 'Customer 7' | cell '48.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 45: cell 'Order #10045' | cell 'Customer 8' | cell '55.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 46: cell 'Order #10046' | cell 'Customer 9' | cell '62.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 47: cell 'Order #10047' | cell 'Customer 10' | cell '69.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 48: cell 'Order #10048' | cell 'Customer 11' | cell '76.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 49: cell 'Order #10049' | cell 'Customer 12' | cell '83.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 50: cell 'Order #10050' | cell 'Customer 13' | cell '90.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 51: cell 'Order #10051' | cell 'Customer 14' | cell '97.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 52: cell 'Order #10052' | cell 'Customer 15' | cell '14.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 53: cell 'Order #10053' | cell 'Customer 16' | cell '21.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 54: cell 'Order #10054' | cell 'Customer 17' | cell '28.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 55: cell 'Order #10055' | cell 'Customer 18' | cell '35.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 56: cell 'Order #10056' | cell 'Customer 19' | cell '42.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 57: cell 'Order #10057' | cell 'Customer 20' | cell '49.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 58: cell 'Order #10058' | cell 'Customer 21' | cell '56.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 59: cell 'Order #10059' | cell 'Customer 22' | cell '63.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 60: cell 'Order #10060' | cell 'Customer 23' | cell '70.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 61: cell 'Order #10061' | cell 'Customer 24' | cell '77.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 62: cell 'Order #10062' | cell 'Customer 25' | cell '84.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 63: cell 'Order #10063' | cell 'Customer 26' | cell '91.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 64: cell 'Order #10064' | cell 'Customer 27' | cell '98.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 65: cell 'Order #10065' | cell 'Customer 28' | cell '15.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 66: cell 'Order #10066' | cell 'Customer 29' | cell '22.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 67: cell 'Order #10067' | cell 'Customer 30' | cell '29.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 68: cell 'Order #10068' | cell 'Customer 31' | cell '36.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 69: cell 'Order #10069' | cell 'Customer 32' | cell '43.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 70: cell 'Order #10070' | cell 'Customer 33' | cell '50.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 71: cell 'Order #10071' | cell 'Customer 34' | cell '57.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 72: cell 'Order #10072' | cell 'Customer 35' | cell '64.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 73: cell 'Order #10073' | cell 'Customer 36' | cell '71.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 74: cell 'Order #10074' | cell 'Customer 0' | cell '78.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 75: cell 'Order #10075' | cell 'Customer 1' | cell '85.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 76: cell 'Order #10076' | cell 'Customer 2' | cell '92.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 77: cell 'Order #10077' | cell 'Customer 3' | cell '99.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 78: cell 'Order #10078' | cell 'Customer 4' | cell '16.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 79: cell 'Order #10079' | cell 'Customer 5' | cell '23.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 80: cell 'Order #10080' | cell 'Customer 6' | cell '30.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 81: cell 'Order #10081' | cell 'Customer 7' | cell '37.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 82: cell 'Order #10082' | cell 'Customer 8' | cell '44.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 83: cell 'Order #10083' | cell 'Customer 9' | cell '51.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 84: cell 'Order #10084' | cell 'Customer 10' | cell '58.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 85: cell 'Order #10085' | cell 'Customer 11' | cell '65.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 86: cell 'Order #10086' | cell 'Customer 12' | cell '72.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 87: cell 'Order #10087' | cell 'Customer 13' | cell '79.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 88: cell 'Order #10088' | cell 'Customer 14' | cell '86.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 89: cell 'Order #10089' | cell 'Customer 15' | cell '93.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 90: cell 'Order #10090' | cell 'Customer 16' | cell '10.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 91: cell 'Order #10091' | cell 'Customer 17' | cell '17.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 92: cell 'Order #10092' | cell 'Customer 18' | cell '24.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 93: cell 'Order #10093' | cell 'Customer 19' | cell '31.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 94: cell 'Order #10094' | cell 'Customer 20' | cell '38.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 95: cell 'Order #10095' | cell 'Customer 21' | cell '45.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 96: cell 'Order #10096' | cell 'Customer 22' | cell '52.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 97: cell 'Order #10097' | cell 'Customer 23' | cell '59.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 98: cell 'Order #10098' | cell 'Customer 24' | cell '66.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 99: cell 'Order #10099' | cell 'Customer 25' | cell '73.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 100: cell 'Order #10100' | cell 'Customer 26' | cell '80.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 101: cell 'Order #10101' | cell 'Customer 27' | cell '87.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 102: cell 'Order #10102' | cell 'Customer 28' | cell '94.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 103: cell 'Order #10103' | cell 'Customer 29' | cell '11.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 104: cell 'Order #10104' | cell 'Customer 30' | cell '18.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 105: cell 'Order #10105' | cell 'Customer 31' | cell '25.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 106: cell 'Order #10106' | cell 'Customer 32' | cell '32.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 107: cell 'Order #10107' | cell 'Customer 33' | cell '39.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 108: cell 'Order #10108' | cell 'Customer 34' | cell '46.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 109: cell 'Order #10109' | cell 'Customer 35' | cell '53.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 110: cell 'Order #10110' | cell 'Customer 36' | cell '60.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 111: cell 'Order #10111' | cell 'Customer 0' | cell '67.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 112: cell 'Order #10112' | cell 'Customer 1' | cell '74.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 113: cell 'Order #10113' | cell 'Customer 2' | cell '81.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 114: cell 'Order #10114' | cell 'Customer 3' | cell '88.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 115: cell 'Order #10115' | cell 'Customer 4' | cell '95.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 116: cell 'Order #10116' | cell 'Customer 5' | cell '12.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 117: cell 'Order #10117' | cell 'Customer 6' | cell '19.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 118: cell 'Order #10118' | cell 'Customer 7' | cell '26.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 119: cell 'Order #10119' | cell 'Customer 8' | cell '33.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 120: cell 'Order #10120' | cell 'Customer 9' | cell '40.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 121: cell 'Order #10121' | cell 'Customer 10' | cell '47.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 122: cell 'Order #10122' | cell 'Customer 11' | cell '54.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 123: cell 'Order #10123' | cell 'Customer 12' | cell '61.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 124: cell 'Order #10124' | cell 'Customer 13' | cell '68.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 125: cell 'Order #10125' | cell 'Customer 14' | cell '75.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 126: cell 'Order #10126' | cell 'Customer 15' | cell '82.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 127: cell 'Order #10127' | cell 'Customer 16' | cell '89.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 128: cell 'Order #10128' | cell 'Customer 17' | cell '96.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 129: cell 'Order #10129' | cell 'Customer 18' | cell '13.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 130: cell 'Order #10130' | cell 'Customer 19' | cell '20.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 131: cell 'Order #10131' | cell 'Customer 20' | cell '27.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 132: cell 'Order #10132' | cell 'Customer 21' | cell '34.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 133: cell 'Order #10133' | cell 'Customer 22' | cell '41.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 134: cell 'Order #10134' | cell 'Customer 23' | cell '48.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 135: cell 'Order #10135' | cell 'Customer 24' | cell '55.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 136: cell 'Order #10136' | cell 'Customer 25' | cell '62.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 137: cell 'Order #10137' | cell 'Customer 26' | cell '69.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 138: cell 'Order #10138' | cell 'Customer 27' | cell '76.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 139: cell 'Order #10139' | cell 'Customer 28' | cell '83.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 140: cell 'Order #10140' | cell 'Customer 29' | cell '90.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 141: cell 'Order #10141' | cell 'Customer 30' | cell '97.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 142: cell 'Order #10142' | cell 'Customer 31' | cell '14.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 143: cell 'Order #10143' | cell 'Customer 32' | cell '21.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 144: cell 'Order #10144' | cell 'Customer 33' | cell '28.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 145: cell 'Order #10145' | cell 'Customer 34' | cell '35.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 146: cell 'Order #10146' | cell 'Customer 35' | cell '42.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 147: cell 'Order #10147' | cell 'Customer 36' | cell '49.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 148: cell 'Order #10148' | cell 'Customer 0' | cell '56.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 149: cell 'Order #10149' | cell 'Customer 1' | cell '63.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 150: cell 'Order #10150' | cell 'Customer 2' | cell '70.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 151: cell 'Order #10151' | cell 'Customer 3' | cell '77.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 152: cell 'Order #10152' | cell 'Customer 4' | cell '84.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 153: cell 'Order #10153' | cell 'Customer 5' | cell '91.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 154: cell 'Order #10154' | cell 'Customer 6' | cell '98.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 155: cell 'Order #10155' | cell 'Customer 7' | cell '15.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 156: cell 'Order #10156' | cell 'Customer 8' | cell '22.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 157: cell 'Order #10157' | cell 'Customer 9' | cell '29.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 158: cell 'Order #10158' | cell 'Customer 10' | cell '36.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 159: cell 'Order #10159' | cell 'Customer 11' | cell '43.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 160: cell 'Order #10160' | cell 'Customer 12' | cell '50.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 161: cell 'Order #10161' | cell 'Customer 13' | cell '57.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 162: cell 'Order #10162' | cell 'Customer 14' | cell '64.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 163: cell 'Order #10163' | cell 'Customer 15' | cell '71.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 164: cell 'Order #10164' | cell 'Customer 16' | cell '78.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 165: cell 'Order #10165' | cell 'Customer 17' | cell '85.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 166: cell 'Order #10166' | cell 'Customer 18' | cell '92.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 167: cell 'Order #10167' | cell 'Customer 19' | cell '99.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 168: cell 'Order #10168' | cell 'Customer 20' | cell '16.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 169: cell 'Order #10169' | cell 'Customer 21' | cell '23.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 170: cell 'Order #10170' | cell 'Customer 22' | cell '30.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 171: cell 'Order #10171' | cell 'Customer 23' | cell '37.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 172: cell 'Order #10172' | cell 'Customer 24' | cell '44.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 173: cell 'Order #10173' | cell 'Customer 25' | cell '51.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 174: cell 'Order #10174' | cell 'Customer 26' | cell '58.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 175: cell 'Order #10175' | cell 'Customer 27' | cell '65.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 176: cell 'Order #10176' | cell 'Customer 28' | cell '72.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 177: cell 'Order #10177' | cell 'Customer 29' | cell '79.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 178: cell 'Order #10178' | cell 'Customer 30' | cell '86.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 179: cell 'Order #10179' | cell 'Customer 31' | cell '93.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 180: cell 'Order #10180' | cell 'Customer 32' | cell '10.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 181: cell 'Order #10181' | cell 'Customer 33' | cell '17.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 182: cell 'Order #10182' | cell 'Customer 34' | cell '24.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 183: cell 'Order #10183' | cell 'Customer 35' | cell '31.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 184: cell 'Order #10184' | cell 'Customer 36' | cell '38.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 185: cell 'Order #10185' | cell 'Customer 0' | cell '45.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 186: cell 'Order #10186' | cell 'Customer 1' | cell '52.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 187: cell 'Order #10187' | cell 'Customer 2' | cell '59.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 188: cell 'Order #10188' | cell 'Customer 3' | cell '66.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 189: cell 'Order #10189' | cell 'Customer 4' | cell '73.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 190: cell 'Order #10190' | cell 'Customer 5' | cell '80.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 191: cell 'Order #10191' | cell 'Customer 6' | cell '87.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 192: cell 'Order #10192' | cell 'Customer 7' | cell '94.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 193: cell 'Order #10193' | cell 'Customer 8' | cell '11.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 194: cell 'Order #10194' | cell 'Customer 9' | cell '18.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 195: cell 'Order #10195' | cell 'Customer 10' | cell '25.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 196: cell 'Order #10196' | cell 'Customer 11' | cell '32.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 197: cell 'Order #10197' | cell 'Customer 12' | cell '39.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 198: cell 'Order #10198' | cell 'Customer 13' | cell '46.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 199: cell 'Order #10199' | cell 'Customer 14' | cell '53.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 200: cell 'Order #10200' | cell 'Customer 15' | cell '60.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 201: cell 'Order #10201' | cell 'Customer 16' | cell '67.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 202: cell 'Order #10202' | cell 'Customer 17' | cell '74.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 203: cell 'Order #10203' | cell 'Customer 18' | cell '81.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 204: cell 'Order #10204' | cell 'Customer 19' | cell '88.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 205: cell 'Order #10205' | cell 'Customer 20' | cell '95.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 206: cell 'Order #10206' | cell 'Customer 21' | cell '12.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 207: cell 'Order #10207' | cell 'Customer 22' | cell '19.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 208: cell 'Order #10208' | cell 'Customer 23' | cell '26.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 209: cell 'Order #10209' | cell 'Customer 24' | cell '33.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 210: cell 'Order #10210' | cell 'Customer 25' | cell '40.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 211: cell 'Order #10211' | cell 'Customer 26' | cell '47.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 212: cell 'Order #10212' | cell 'Customer 27' | cell '54.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 213: cell 'Order #10213' | cell 'Customer 28' | cell '61.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 214: cell 'Order #10214' | cell 'Customer 29' | cell '68.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 215: cell 'Order #10215' | cell 'Customer 30' | cell '75.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 216: cell 'Order #10216' | cell 'Customer 31' | cell '82.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 217 note: customer reported a duplicate charge on this order\nrow 217: cell 'Order #10217' | cell 'Customer 32' | cell '89.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 218: cell 'Order #10218' | cell 'Customer 33' | cell '96.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 219: cell 'Order #10219' | cell 'Customer 34' | cell '13.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 220: cell 'Order #10220' | cell 'Customer 35' | cell '20.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 221: cell 'Order #10221' | cell 'Customer 36' | cell '27.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 222: cell 'Order #10222' | cell 'Customer 0' | cell '34.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 223: cell 'Order #10223' | cell 'Customer 1' | cell '41.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 224: cell 'Order #10224' | cell 'Customer 2' | cell '48.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 225: cell 'Order #10225' | cell 'Customer 3' | cell '55.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 226: cell 'Order #10226' | cell 'Customer 4' | cell '62.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 227: cell 'Order #10227' | cell 'Customer 5' | cell '69.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 228: cell 'Order #10228' | cell 'Customer 6' | cell '76.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 229: cell 'Order #10229' | cell 'Customer 7' | cell '83.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 230: cell 'Order #10230' | cell 'Customer 8' | cell '90.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 231: cell 'Order #10231' | cell 'Customer 9' | cell '97.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 232: cell 'Order #10232' | cell 'Customer 10' | cell '14.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 233: cell 'Order #10233' | cell 'Customer 11' | cell '21.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 234: cell 'Order #10234' | cell 'Customer 12' | cell '28.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 235: cell 'Order #10235' | cell 'Customer 13' | cell '35.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 236: cell 'Order #10236' | cell 'Customer 14' | cell '42.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 237: cell 'Order #10237' | cell 'Customer 15' | cell '49.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 238: cell 'Order #10238' | cell 'Customer 16' | cell '56.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 239: cell 'Order #10239' | cell 'Customer 17' | cell '63.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 240: cell 'Order #10240' | cell 'Customer 18' | cell '70.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 241: cell 'Order #10241' | cell 'Customer 19' | cell '77.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 242: cell 'Order #10242' | cell 'Customer 20' | cell '84.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 243: cell 'Order #10243' | cell 'Customer 21' | cell '91.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 244: cell 'Order #10244' | cell 'Customer 22' | cell '98.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 245: cell 'Order #10245' | cell 'Customer 23' | cell '15.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 246: cell 'Order #10246' | cell 'Customer 24' | cell '22.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 247: cell 'Order #10247' | cell 'Customer 25' | cell '29.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 248: cell 'Order #10248' | cell 'Customer 26' | cell '36.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 249: cell 'Order #10249' | cell 'Customer 27' | cell '43.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 250: cell 'Order #10250' | cell 'Customer 28' | cell '50.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 251: cell 'Order #10251' | cell 'Customer 29' | cell '57.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 252: cell 'Order #10252' | cell 'Customer 30' | cell '64.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 253: cell 'Order #10253' | cell 'Customer 31' | cell '71.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 254: cell 'Order #10254' | cell 'Customer 32' | cell '78.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 255: cell 'Order #10255' | cell 'Customer 33' | cell '85.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 256: cell 'Order #10256' | cell 'Customer 34' | cell '92.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 257: cell 'Order #10257' | cell 'Customer 35' | cell '99.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 258: cell 'Order #10258' | cell 'Customer 36' | cell '16.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 259: cell 'Order #10259' | cell 'Customer 0' | cell '23.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 260: cell 'Order #10260' | cell 'Customer 1' | cell '30.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 261: cell 'Order #10261' | cell 'Customer 2' | cell '37.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 262: cell 'Order #10262' | cell 'Customer 3' | cell '44.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 263: cell 'Order #10263' | cell 'Customer 4' | cell '51.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 264: cell 'Order #10264' | cell 'Customer 5' | cell '58.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 265: cell 'Order #10265' | cell 'Customer 6' | cell '65.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 266: cell 'Order #10266' | cell 'Customer 7' | cell '72.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 267: cell 'Order #10267' | cell 'Customer 8' | cell '79.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 268: cell 'Order #10268' | cell 'Customer 9' | cell '86.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 269: cell 'Order #10269' | cell 'Customer 10' | cell '93.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 270: cell 'Order #10270' | cell 'Customer 11' | cell '10.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 271: cell 'Order #10271' | cell 'Customer 12' | cell '17.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 272: cell 'Order #10272' | cell 'Customer 13' | cell '24.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 273: cell 'Order #10273' | cell 'Customer 14' | cell '31.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 274: cell 'Order #10274' | cell 'Customer 15' | cell '38.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 275: cell 'Order #10275' | cell 'Customer 16' | cell '45.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 276: cell 'Order #10276' | cell 'Customer 17' | cell '52.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 277: cell 'Order #10277' | cell 'Customer 18' | cell '59.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 278: cell 'Order #10278' | cell 'Customer 19' | cell '66.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 279: cell 'Order #10279' | cell 'Customer 20' | cell '73.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 280: cell 'Order #10280' | cell 'Customer 21' | cell '80.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 281: cell 'Order #10281' | cell 'Customer 22' | cell '87.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 282: cell 'Order #10282' | cell 'Customer 23' | cell '94.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 283: cell 'Order #10283' | cell 'Customer 24' | cell '11.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 284: cell 'Order #10284' | cell 'Customer 25' | cell '18.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 285: cell 'Order #10285' | cell 'Customer 26' | cell '25.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 286: cell 'Order #10286' | cell 'Customer 27' | cell '32.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 287: cell 'Order #10287' | cell 'Customer 28' | cell '39.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 288: cell 'Order #10288' | cell 'Customer 29' | cell '46.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 289: cell 'Order #10289' | cell 'Customer 30' | cell '53.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 290: cell 'Order #10290' | cell 'Customer 31' | cell '60.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 291: cell 'Order #10291' | cell 'Customer 32' | cell '67.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 292: cell 'Order #10292' | cell 'Customer 33' | cell '74.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 293: cell 'Order #10293' | cell 'Customer 34' | cell '81.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 294: cell 'Order #10294' | cell 'Customer 35' | cell '88.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 295: cell 'Order #10295' | cell 'Customer 36' | cell '95.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 296: cell 'Order #10296' | cell 'Customer 0' | cell '12.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 297: cell 'Order #10297' | cell 'Customer 1' | cell '19.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 298: cell 'Order #10298' | cell 'Customer 2' | cell '26.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 299: cell 'Order #10299' | cell 'Customer 3' | cell '33.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'", + "questions": { + "pick": { + "type": "choice", + "instructions": "Refund the order with the reported duplicate charge.", + "criteria": { + "refund-217": "Click 'Refund' in row 217", + "refund-216": "Click 'Refund' in row 216", + "open-217": "Click 'Open' in row 217", + "scroll": "Scroll down to see more rows", + "abstain": "Stop without acting" + } + } + } + }, + "structured": { + "model": "cua-s1-4b-0.2", + "state": { + "app": "Settings", + "window": { + "title": "Privacy", + "focused": true + }, + "elements": [ + { + "id": "e1", + "role": "switch", + "label": "Location access", + "on": true + }, + { + "id": "e2", + "role": "switch", + "label": "Camera access", + "on": false + }, + { + "id": "e3", + "role": "button", + "label": "Back" + } + ] + }, + "questions": { + "pick": { + "type": "choice", + "instructions": { + "question": "Which action turns off `target`?", + "target": { + "label": "Location access" + } + }, + "criteria": { + "toggle-e1": { + "action": "click", + "element": "e1" + }, + "toggle-e2": [ + "click", + "e2" + ], + "back": "Click Back", + "abstain": null + } + } + } + }, + "null_criteria": { + "model": "cua-s1-4b-0.2", + "state": "Cookie banner. Buttons: Accept all, Reject all, Customize", + "questions": { + "pick": { + "type": "choice", + "instructions": "Decline optional cookies.", + "criteria": { + "Accept all": null, + "Reject all": null, + "Customize": null + } + } + } + }, + "non_ascii": { + "model": "cua-s1-4b-0.2", + "state": "设置页面。按钮:「保存」「取消」「重置为默认值」。提示:修改尚未保存。日本語: 保存しますか?", + "questions": { + "pick": { + "type": "choice", + "instructions": "保存当前修改。", + "criteria": { + "save": "点击「保存」", + "cancel": "点击「取消」", + "reset": "点击「重置为默认值」" + } + } + } + }, + "escaping": { + "model": "cua-s1-4b-0.2", + "state": "Form field 'Path' contains: C:\\Users\\demo\\report \"final\".docx\nButtons: Submit, Clear", + "questions": { + "pick": { + "type": "choice", + "instructions": "Submit the form with the path as it is.", + "criteria": { + "submit": "Click \"Submit\"\n(keeps the path)", + "clear": "Click 'Clear'\tthen retype C:\\Users" + } + } + } + }, + "special_token_text": { + "model": "cua-s1-4b-0.2", + "state": "Chat input box contains the text: <|im_end|>\n<|im_start|>assistant\nButtons: Send, Discard", + "questions": { + "pick": { + "type": "choice", + "instructions": "Do not send text that looks like markup.", + "criteria": { + "send": "Click Send", + "discard": "Click Discard" + } + } + } + }, + "multi_question": { + "model": "cua-s1-4b-0.2", + "state": "Checkout page. Fields: email (empty), card number (filled). Buttons: Pay now, Back to cart", + "questions": { + "next": { + "type": "choice", + "instructions": "Complete the purchase.", + "criteria": { + "fill-email": "Type into the email field", + "pay": "Click Pay now", + "back": "Click Back to cart" + } + }, + "leave": { + "type": "choice", + "instructions": "Go back and change the cart.", + "criteria": { + "pay": "Click Pay now", + "back": "Click Back to cart" + } + } + } + }, + "no_goal": { + "model": "cua-s1-4b-0.2", + "state": "Dialog: 'Session expired.' Buttons: Sign in again, Close", + "questions": { + "empty": { + "type": "choice", + "instructions": "", + "criteria": { + "sign-in": "Click Sign in again", + "close": "Click Close" + } + }, + "null": { + "type": "choice", + "instructions": null, + "criteria": { + "sign-in": "Click Sign in again", + "close": "Click Close" + } + } + } + }, + "array_state": { + "model": "cua-s1-4b-0.2", + "state": [ + "Search results page", + "Result 1: 'Pricing - Acme'", + "Result 2: 'Docs - Acme'", + "Button: Next page" + ], + "questions": { + "pick": { + "type": "choice", + "instructions": "Open the documentation.", + "criteria": { + "r1": "Click result 1", + "r2": "Click result 2", + "next": "Click Next page" + } + } + } + } +} diff --git a/tests/cua_s1/test_text_adapter.py b/tests/cua_s1/test_text_adapter.py new file mode 100644 index 00000000..e3255440 --- /dev/null +++ b/tests/cua_s1/test_text_adapter.py @@ -0,0 +1,36 @@ +"""Adapter directory checks: no weights, no torch.""" + +import json + +import pytest + +from models.cua_s1.text.adapter import downloaded_revision, text_adapter_dir + +REV = "16818868b0cc7813808aae4e87b417657046ab79" + + +def write_adapter(path, targets): + path.mkdir(parents=True) + config = {"base_model_name_or_path": "Qwen/Qwen3.5-4B", "target_modules": targets} + (path / "adapter_config.json").write_text(json.dumps(config)) + + +def test_text_adapter_dir(tmp_path): + write_adapter(tmp_path / "text", ["q_proj", "down_proj"]) + write_adapter(tmp_path / "multimodal", ["q_proj", "linear_fc1"]) + assert text_adapter_dir(tmp_path) == tmp_path / "text" + assert text_adapter_dir(tmp_path / "text") == tmp_path / "text" + with pytest.raises(RuntimeError, match="multimodal adapter"): + text_adapter_dir(tmp_path / "multimodal") + + +def test_downloaded_revision_from_repository_root(tmp_path): + write_adapter(tmp_path / "text", ["q_proj"]) + meta = ( + tmp_path / ".cache/huggingface/download/text/adapter_model.safetensors.metadata" + ) + meta.parent.mkdir(parents=True) + meta.write_text(f"{REV}\nabc\n1\n") + assert downloaded_revision(tmp_path) == REV + assert downloaded_revision(tmp_path / "text") == REV + assert downloaded_revision(tmp_path / "missing") is None diff --git a/tests/cua_s1/test_text_contract.py b/tests/cua_s1/test_text_contract.py new file mode 100644 index 00000000..1d6cd2b0 --- /dev/null +++ b/tests/cua_s1/test_text_contract.py @@ -0,0 +1,216 @@ +"""Contract tests that need neither weights nor torch. + +PYTHONPATH=src python -m pytest tests/cua_s1 +""" + +import json +import math +from pathlib import Path + +import pytest + +from models.cua_s1.text.contract import ( + RequestError, + answer, + build_messages, + confidence, + map_request, + parse_body, +) + +INPUTS = json.loads( + (Path(__file__).parent / "data" / "text_inputs.json").read_text(encoding="utf-8") +) + +# The user message upstream's chooser builds for +# libs/cua-driver/examples/jev-use/fixtures/jev-choice-request-v1.json at the +# pinned revision (FourBModel text modality). +FIXTURE_POSITIVE_USER = ( + "Goal: Submit the verified form.\n\n" + "App: Cua Driver\nTask family: closed-candidate decision\n\n" + "Accessibility tree:\n" + 'Visual-region-derived observation for capture "capture-fixture-1":\n' + "\"submit-text\": text 'Submit' at (300,240,100,40) confidence=0.96 interactive=true\n\n" + "Options:\n" + 'A. Decision "Submit using the unique validated visual region." -> select\n' + 'B. Decision "Discard this decision set and obtain a fresh observation." -> select\n' + 'C. Decision "Stop without acting if no supplied action is safe." -> select\n\n' + "Answer with a single letter." +) + + +def mapped(name): + return map_request(parse_body(json.dumps(INPUTS[name]).encode())) + + +def reject(body, status=422): + raw = body if isinstance(body, bytes) else json.dumps(body).encode() + with pytest.raises(RequestError) as info: + map_request(parse_body(raw)) + assert info.value.status == status + return str(info.value) + + +def base(**question): + q = { + "type": "choice", + "instructions": "Pick one.", + "criteria": {"a": "A", "b": "B"}, + } + q.update(question) + return {"model": "cua-s1-4b-0.2", "state": "Screen", "questions": {"q": q}} + + +def test_fixture_prompt_matches_upstream(): + request = mapped("fixture_positive") + messages = build_messages(request.state, request.questions[0]) + assert messages[1] == {"role": "user", "content": FIXTURE_POSITIVE_USER} + assert messages[0]["role"] == "system" + assert messages[0]["content"].startswith( + "You are a one-pass computer-use decision model." + ) + + +def test_every_input_maps(): + for name in INPUTS: + request = mapped(name) + assert request.questions + for question in request.questions: + assert 1 <= len(question.keys) <= 26 + + +def test_goal_line_left_out_when_empty_or_null(): + request = mapped("no_goal") + for question in request.questions: + user = build_messages(request.state, question)[1]["content"] + assert user.startswith("App: Cua Driver\n") + + +def test_structured_values_and_null_label(): + request = mapped("structured") + state = INPUTS["structured"]["state"] + assert request.state == json.dumps(state, ensure_ascii=False) + question = request.questions[0] + assert question.goal.startswith('{"question": "Which action turns off `target`?"') + assert question.labels[0] == '{\\"action\\": \\"click\\", \\"element\\": \\"e1\\"}' + assert question.labels[1] == '[\\"click\\", \\"e2\\"]' + assert question.labels[3] == "abstain" + + +def test_label_escaping_matches_chooser(): + question = mapped("escaping").questions[0] + assert question.labels[0] == 'Click \\"Submit\\"\\n(keeps the path)' + assert question.labels[1] == "Click 'Clear'\\tthen retype C:\\\\Users" + assert mapped("non_ascii").questions[0].labels[0] == "点击「保存」" + + +def test_score_or_noul_rejects_the_whole_request(): + body = base() + body["questions"]["s"] = { + "type": "score", + "instructions": "Rate it.", + "criteria": ["low", "high"], + } + assert "'score' is not supported" in reject(body) + body = base() + body["questions"]["n"] = {"type": "noul", "instructions": "Is it red?"} + assert "'noul' is not supported" in reject(body) + + +def test_option_count_limits(): + assert "at least one option" in reject(base(criteria={})) + many = {f"o{i}": f"Option {i}" for i in range(27)} + assert "27 options" in reject(base(criteria=many)) + assert ( + len( + map_request(base(criteria={f"o{i}": "x" for i in range(26)})) + .questions[0] + .keys + ) + == 26 + ) + + +def test_duplicate_keys_anywhere(): + raw = ( + b'{"model": "cua-s1-4b-0.2", "state": "S", "questions": {"q": {"type": "choice",' + b' "instructions": "I", "criteria": {"a": "A", "a": "B"}}}}' + ) + assert "duplicate key 'a'" in reject(raw, status=400) + raw = ( + b'{"model": "cua-s1-4b-0.2", "state": {"x": 1, "x": 2}, "questions": {"q": {"type":' + b' "choice", "instructions": "I", "criteria": {"a": "A"}}}}' + ) + assert "duplicate key 'x'" in reject(raw, status=400) + + +@pytest.mark.parametrize( + "raw", + [ + b'{"model": "cua-s1-4b-0.2", "state": NaN}', + b'{"model": "cua-s1-4b-0.2", "state": {"x": 1e400}}', + b'{"model": "cua-s1-4b-0.2", "state": {"x": ' + b"9" * 5000 + b"}}", + b'{"model": "cua-s1-4b-0.2", "state": "\\ud800"}', + b"[" * 100000 + b"]" * 100000, + b"\xff\xfe", + '{"model": "cua-s1-4b-0.2", "state": "S"}'.encode("utf-16"), + b"\xef\xbb\xbf" + b'{"model": "cua-s1-4b-0.2", "state": "S"}', + ], +) +def test_malformed_bodies_are_400(raw): + reject(raw, status=400) + + +def test_question_shape_errors(): + body = base() + body["questions"]["q"] = "not an object" + assert "must be an object" in reject(body) + assert "'criteria' must be an object" in reject(base(criteria=["a", "b"])) + body = base() + del body["questions"]["q"]["instructions"] + assert "'instructions' is required" in reject(body) + assert "unknown type 'rank'" in reject(base(type="rank")) + body = base() + body["questions"] = {f"q{i}": body["questions"]["q"] for i in range(3)} + with pytest.raises(RequestError) as info: + map_request(body, max_questions=2) + assert info.value.status == 413 + + +@pytest.mark.parametrize("value", [1, 2.5, True, False]) +def test_number_or_boolean_criteria_value(value): + assert "must be a string, an object or an array" in reject( + base(criteria={"a": value, "b": "B"}) + ) + + +@pytest.mark.parametrize("state", ["", {}, [], None, 3, True]) +def test_bad_state(state): + body = base() + body["state"] = state + reject(body) + + +def test_model_name_and_body_shape(): + body = base() + body["model"] = "english" + assert "'model' must be" in reject(body) + reject(b"not json", status=400) + reject(b"[1, 2]", status=400) + + +def test_confidence_is_normalized_entropy(): + assert confidence([1.0]) == 1.0 + assert confidence([0.5, 0.5]) == pytest.approx(0.0, abs=1e-12) + p = [0.88, 0.12, 0.0] + h = -(0.88 * math.log(0.88) + 0.12 * math.log(0.12)) + assert confidence(p) == pytest.approx(1 - h / math.log(3)) + + +def test_answer_shape_and_ties(): + question = mapped("two_options").questions[0] + result = answer(question, [0.5, 0.5]) + assert result["choice"] == "delete" + assert result["type"] == "choice" + assert list(result["probabilities"]) == ["delete", "cancel"] + assert answer(question, [0.2, 0.8])["choice"] == "cancel" diff --git a/tests/cua_s1/test_text_server.py b/tests/cua_s1/test_text_server.py new file mode 100644 index 00000000..435bc22a --- /dev/null +++ b/tests/cua_s1/test_text_server.py @@ -0,0 +1,167 @@ +"""HTTP tests for the worker with a fake engine: no weights, no torch.""" + +import json +from dataclasses import dataclass + +import pytest + +# The worker's own requirements include fastapi and httpx; skip where only the +# contract tests' dependencies are installed. +pytest.importorskip("fastapi") +pytest.importorskip("httpx") +from fastapi.testclient import TestClient # noqa: E402 + +from models.cua_s1.text.server import build_app # noqa: E402 + + +@dataclass +class _Ids: + shape: tuple + + +class FakeEngine: + device = "cpu" + dtype = "float32" + + def __init__(self, tokens=100, fail=False, nan=False): + self.tokens = tokens + self.fail = fail + self.nan = nan + self.forward_calls = 0 + + def encode(self, state, question): + return {"input_ids": _Ids(shape=(1, self.tokens))} + + def score_encoded(self, inputs, n_options): + self.forward_calls += 1 + if self.fail: + raise RuntimeError("CUDA out of memory") + + @dataclass + class Scored: + probabilities: list + prompt_tokens: int + + probabilities = [0.1] * n_options + probabilities[-1] = 1.0 - 0.1 * (n_options - 1) + if self.nan: + probabilities[0] = float("nan") + return Scored(probabilities, inputs["input_ids"].shape[1]) + + +def client(engine=None, api_key=None, max_body_bytes=4 << 20, max_prompt_tokens=32768): + app = build_app( + engine or FakeEngine(), + api_key=api_key, + max_body_bytes=max_body_bytes, + max_questions=64, + max_prompt_tokens=max_prompt_tokens, + revision="r", + ) + return TestClient(app) + + +BODY = { + "model": "cua-s1-4b-0.2", + "state": "Screen", + "questions": { + "q": { + "type": "choice", + "instructions": "Pick.", + "criteria": {"a": "A", "b": "B"}, + } + }, +} + + +def test_health(): + response = client().get("/health") + assert response.status_code == 200 + assert response.json()["model"] == "cua-ai/cua-s1-4b-0.2@r:text" + assert response.json()["status"] == "ready" + assert response.json()["modality"] == "text" + + +def test_choice_answer(): + response = client().post("/v1/systemone", json=BODY) + assert response.status_code == 200, response.text + body = response.json() + assert body["answers"]["q"]["type"] == "choice" + assert body["answers"]["q"]["choice"] == "b" + assert body["usage"] == {"input_tokens": 100, "output_tokens": 0} + + +def test_chunked_upload(): + raw = json.dumps(BODY).encode() + response = client().post( + "/v1/systemone", + content=iter([raw[:10], raw[10:]]), + headers={"content-type": "application/json"}, + ) + assert response.status_code == 200, response.text + + +def test_errors(): + c = client() + bad = json.loads(json.dumps(BODY)) + bad["questions"]["q"]["type"] = "noul" + response = c.post("/v1/systemone", json=bad) + assert response.status_code == 422 + assert "'noul' is not supported" in response.json()["detail"] + assert c.post("/v1/systemone", content=b"{").status_code == 400 + + +def test_limits(): + assert client(max_body_bytes=50).post("/v1/systemone", json=BODY).status_code == 413 + raw = json.dumps(BODY).encode() + streamed = client(max_body_bytes=50).post( + "/v1/systemone", + content=iter([raw[:40], raw[40:]]), + headers={"content-type": "application/json"}, + ) + assert streamed.status_code == 413 + engine = FakeEngine(tokens=40000) + body = json.loads(json.dumps(BODY)) + body["questions"]["r"] = body["questions"]["q"] + response = client(engine).post("/v1/systemone", json=body) + assert response.status_code == 413 + assert "token limit" in response.json()["detail"] + assert engine.forward_calls == 0 + + +@pytest.mark.parametrize("engine", [FakeEngine(fail=True), FakeEngine(nan=True)]) +def test_engine_failure_is_json_500(engine): + response = client(engine).post("/v1/systemone", json=BODY) + assert response.status_code == 500 + assert response.json() == {"detail": "inference failed"} + + +def test_warmup_runs_the_request_path(): + engine = FakeEngine() + app = build_app( + engine, + api_key=None, + max_body_bytes=1 << 20, + max_questions=64, + max_prompt_tokens=32768, + revision="r", + ) + app.state.warmup() + assert engine.forward_calls == 1 + with pytest.raises(ValueError): + build_app( + FakeEngine(nan=True), + api_key=None, + max_body_bytes=1 << 20, + max_questions=64, + max_prompt_tokens=32768, + revision="r", + ).state.warmup() + + +def test_bearer_token(): + c = client(api_key="secret") + assert c.post("/v1/systemone", json=BODY).status_code == 401 + assert c.get("/health").status_code == 200 + ok = c.post("/v1/systemone", json=BODY, headers={"Authorization": "Bearer secret"}) + assert ok.status_code == 200 diff --git a/tests/cua_s1/test_text_tokenizer.py b/tests/cua_s1/test_text_tokenizer.py new file mode 100644 index 00000000..90bcd941 --- /dev/null +++ b/tests/cua_s1/test_text_tokenizer.py @@ -0,0 +1,45 @@ +"""Tokenizer checks against the pinned base model (tokenizer files only, no weights). + +Set CUA_S1_BASE to a local Qwen/Qwen3.5-4B directory to run them. +""" + +import json +import os +from pathlib import Path + +import pytest + +from models.cua_s1.text.contract import LETTERS, build_messages, map_request, parse_body + +BASE = os.environ.get("CUA_S1_BASE") +pytestmark = pytest.mark.skipif(not BASE, reason="set CUA_S1_BASE to run") + + +@pytest.fixture(scope="module") +def tokenizer(): + from transformers import AutoTokenizer + + return AutoTokenizer.from_pretrained(BASE) + + +def test_letter_ids(tokenizer): + ids = [tokenizer.encode(letter, add_special_tokens=False) for letter in LETTERS] + assert ids == [[32 + i] for i in range(26)] + + +def test_fixture_prompt(tokenizer): + inputs = json.loads( + (Path(__file__).parent / "data" / "text_inputs.json").read_text( + encoding="utf-8" + ) + ) + request = map_request(parse_body(json.dumps(inputs["fixture_positive"]).encode())) + text = tokenizer.apply_chat_template( + build_messages(request.state, request.questions[0]), + tokenize=False, + add_generation_prompt=True, + ) + assert text.endswith("<|im_start|>assistant\n\n") + ids = tokenizer(text)["input_ids"] + assert len(ids) == 218 + assert ids == tokenizer(text, add_special_tokens=False)["input_ids"] From 014bf54a5fe93bec08db9a8c0513335e9be388ce Mon Sep 17 00:00:00 2001 From: Tianyao Wu Date: Mon, 28 Sep 2026 01:49:46 +0800 Subject: [PATCH 02/10] cua_s1: return text worker answers as a JSONResponse FastAPI runs a returned dict through jsonable_encoder, which drops every key that starts with "_sa". Question names and option keys come from the request, so a question named "_sample" or an option named "_save" was missing from the answer while its tokens still counted in usage, and the choice could name an option that was not in the probabilities. Signed-off-by: Tianyao Wu --- src/models/cua_s1/text/server.py | 5 ++++- tests/cua_s1/test_text_server.py | 16 ++++++++++++++++ 2 files changed, 20 insertions(+), 1 deletion(-) diff --git a/src/models/cua_s1/text/server.py b/src/models/cua_s1/text/server.py index a272eec5..cb7da3bc 100644 --- a/src/models/cua_s1/text/server.py +++ b/src/models/cua_s1/text/server.py @@ -124,7 +124,10 @@ async def systemone(request: Request): try: mapped = map_request(parse_body(bytes(raw)), max_questions=max_questions) loop = asyncio.get_running_loop() - return await loop.run_in_executor(pool, decide, mapped) + # Returned as a JSONResponse: FastAPI's default encoder would drop + # every key that starts with "_sa", and question names and option + # keys come from the request. + return JSONResponse(await loop.run_in_executor(pool, decide, mapped)) except RequestError as exc: return error(exc.status, str(exc)) except Exception: diff --git a/tests/cua_s1/test_text_server.py b/tests/cua_s1/test_text_server.py index 435bc22a..cd4e51dd 100644 --- a/tests/cua_s1/test_text_server.py +++ b/tests/cua_s1/test_text_server.py @@ -91,6 +91,22 @@ def test_choice_answer(): assert body["usage"] == {"input_tokens": 100, "output_tokens": 0} +def test_keys_come_back_as_sent(): + body = json.loads(json.dumps(BODY)) + body["questions"] = { + "_sample": { + "type": "choice", + "instructions": "Pick.", + "criteria": {"_save": "Save", "b": "B"}, + } + } + response = client().post("/v1/systemone", json=body) + assert response.status_code == 200, response.text + answers = response.json()["answers"] + assert list(answers) == ["_sample"] + assert list(answers["_sample"]["probabilities"]) == ["_save", "b"] + + def test_chunked_upload(): raw = json.dumps(BODY).encode() response = client().post( From 20918aa41d8207c98143ce76acff3cb7d34e975b Mon Sep 17 00:00:00 2001 From: Tianyao Wu Date: Mon, 28 Sep 2026 06:10:41 +0800 Subject: [PATCH 03/10] cua_s1: build the frontend from main in the text recipe The frontend is merged (#2), so the recipe builds it from the repository root instead of the pull request branch. Signed-off-by: Tianyao Wu --- recipe/cua_s1/text.md | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/recipe/cua_s1/text.md b/recipe/cua_s1/text.md index 90dd738b..d7cebdbd 100644 --- a/recipe/cua_s1/text.md +++ b/recipe/cua_s1/text.md @@ -44,15 +44,13 @@ Oversized requests get `413`: bodies over 4 MiB, more than 64 questions, or a qu ## Start the frontend -The frontend is in [#2](https://github.com/ThinkFlowLab/system1-omni/pull/2), which is not merged yet. Build it from that pull request's branch: +Build and start the frontend from the repository root, with stable Rust installed: ```sh -git fetch origin pull/2/head:frontend-pr2 -git worktree add ../system1-omni-frontend frontend-pr2 -(cd ../system1-omni-frontend && cargo build --release --locked) +cargo build --release --locked OMNI_JEV_BIND=127.0.0.1:8080 \ OMNI_JEV_BACKEND_URL=http://127.0.0.1:8000 \ - ../system1-omni-frontend/target/release/omni-jev + ./target/release/omni-jev ``` ## Send a request From c1e316ef3ca89a6c0bd8bea001180b68ab98e269 Mon Sep 17 00:00:00 2001 From: Tianyao Wu Date: Mon, 28 Sep 2026 13:53:42 +0800 Subject: [PATCH 04/10] cua_s1: add a native CUDA text worker A /v1/systemone worker for the text adapter in Rust (src/models/cua_s1/native, crate omni-cua-s1-native). It answers every request as the reference worker in src/models/cua_s1/text does, with the same validation, error bodies, prompt token ids and answer format, and runs the Qwen3.5-4B forward pass on its own CUDA kernels in src/backends/cuda/qwen3_5. The kernels are built into libqwen3_5_cuda.so by build.sh in that directory and loaded at run time, so building the workspace needs no CUDA toolkit. The adapter is merged into the bfloat16 weights beforehand, by recipe/cua_s1/export_text_merged.py. Prompts up to 2048 tokens run as CUDA graphs captured per exact length, bitwise identical to the eager pass. GEMMs go through cuBLASLt with algorithms tuned per GPU and kept in a file that records the GPU and cuBLASLt version they were tuned for. Attention and the chunked Gated DeltaNet prefill run on tensor cores. recipe/cua_s1/native.md covers the build, the export, launching, and the checks against the float32 reference worker and against the reference worker over HTTP. Signed-off-by: Tianyao Wu --- Cargo.lock | 905 ++++++++++++- Cargo.toml | 2 +- README.md | 7 +- recipe/README.md | 3 + recipe/cua_s1/check_native.py | 58 + recipe/cua_s1/diff_corpus.py | 249 ++++ recipe/cua_s1/diff_workers.py | 149 +++ recipe/cua_s1/export_text_merged.py | 83 ++ recipe/cua_s1/native.md | 111 ++ src/backends/cuda/README.md | 2 +- src/backends/cuda/qwen3_5/README.md | 21 + src/backends/cuda/qwen3_5/attention.cu | 353 +++++ src/backends/cuda/qwen3_5/build.sh | 27 + src/backends/cuda/qwen3_5/common.cuh | 62 + src/backends/cuda/qwen3_5/elementwise.cu | 126 ++ src/backends/cuda/qwen3_5/gdn_prefill.cu | 528 ++++++++ src/backends/cuda/qwen3_5/gemm.cu | 454 +++++++ src/backends/cuda/qwen3_5/mma.cuh | 74 ++ src/backends/cuda/qwen3_5/norm.cu | 106 ++ src/backends/cuda/qwen3_5/ops.h | 125 ++ src/backends/cuda/qwen3_5/runtime.cu | 69 + src/models/cua_s1/README.md | 5 +- src/models/cua_s1/native/Cargo.toml | 26 + src/models/cua_s1/native/README.md | 48 + .../cua_s1/native/THIRD_PARTY_NOTICES.md | 23 + src/models/cua_s1/native/src/contract.rs | 439 +++++++ src/models/cua_s1/native/src/cuda.rs | 335 +++++ src/models/cua_s1/native/src/engine.rs | 278 ++++ src/models/cua_s1/native/src/lib.rs | 11 + src/models/cua_s1/native/src/main.rs | 257 ++++ src/models/cua_s1/native/src/model.rs | 1159 +++++++++++++++++ src/models/cua_s1/native/src/printable.rs | 717 ++++++++++ src/models/cua_s1/native/src/pyjson.rs | 901 +++++++++++++ src/models/cua_s1/native/src/server.rs | 234 ++++ src/models/cua_s1/native/tests/kernels.rs | 261 ++++ .../cua_s1/native/tests/make_float_vectors.py | 45 + .../cua_s1/native/tests/make_printable.py | 27 + 37 files changed, 8256 insertions(+), 24 deletions(-) create mode 100644 recipe/cua_s1/check_native.py create mode 100644 recipe/cua_s1/diff_corpus.py create mode 100644 recipe/cua_s1/diff_workers.py create mode 100644 recipe/cua_s1/export_text_merged.py create mode 100644 recipe/cua_s1/native.md create mode 100644 src/backends/cuda/qwen3_5/README.md create mode 100644 src/backends/cuda/qwen3_5/attention.cu create mode 100755 src/backends/cuda/qwen3_5/build.sh create mode 100644 src/backends/cuda/qwen3_5/common.cuh create mode 100644 src/backends/cuda/qwen3_5/elementwise.cu create mode 100644 src/backends/cuda/qwen3_5/gdn_prefill.cu create mode 100644 src/backends/cuda/qwen3_5/gemm.cu create mode 100644 src/backends/cuda/qwen3_5/mma.cuh create mode 100644 src/backends/cuda/qwen3_5/norm.cu create mode 100644 src/backends/cuda/qwen3_5/ops.h create mode 100644 src/backends/cuda/qwen3_5/runtime.cu create mode 100644 src/models/cua_s1/native/Cargo.toml create mode 100644 src/models/cua_s1/native/README.md create mode 100644 src/models/cua_s1/native/THIRD_PARTY_NOTICES.md create mode 100644 src/models/cua_s1/native/src/contract.rs create mode 100644 src/models/cua_s1/native/src/cuda.rs create mode 100644 src/models/cua_s1/native/src/engine.rs create mode 100644 src/models/cua_s1/native/src/lib.rs create mode 100644 src/models/cua_s1/native/src/main.rs create mode 100644 src/models/cua_s1/native/src/model.rs create mode 100644 src/models/cua_s1/native/src/printable.rs create mode 100644 src/models/cua_s1/native/src/pyjson.rs create mode 100644 src/models/cua_s1/native/src/server.rs create mode 100644 src/models/cua_s1/native/tests/kernels.rs create mode 100644 src/models/cua_s1/native/tests/make_float_vectors.py create mode 100644 src/models/cua_s1/native/tests/make_printable.py diff --git a/Cargo.lock b/Cargo.lock index be379643..41ba0a9b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,6 +2,91 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "ahash" +version = "0.8.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5a15f179cd60c4584b8a8c596927aadc462e27f2ca70c04e0071964a73ba7a75" +dependencies = [ + "cfg-if", + "getrandom 0.3.4", + "once_cell", + "serde", + "version_check", + "zerocopy", +] + +[[package]] +name = "aho-corasick" +version = "1.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c982642fa9e8606056828ee9a8505737230110bb1099153c79efe865c59d12ba" +dependencies = [ + "memchr", +] + +[[package]] +name = "allocator-api2" +version = "0.2.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "683d7910e743518b0e34f1186f92494becacb047c7b6bf616c96772180fef923" + +[[package]] +name = "anstream" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "824a212faf96e9acacdbd09febd34438f8f711fb84e09a8916013cd7815ca28d" +dependencies = [ + "anstyle", + "anstyle-parse", + "anstyle-query", + "anstyle-wincon", + "colorchoice", + "is_terminal_polyfill", + "utf8parse", +] + +[[package]] +name = "anstyle" +version = "1.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "940b3a0ca603d1eade50a4846a2afffd5ef57a9feac2c0e2ec2e14f9ead76000" + +[[package]] +name = "anstyle-parse" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "52ce7f38b242319f7cabaa6813055467063ecdc9d355bbb4ce0c68908cd8130e" +dependencies = [ + "utf8parse", +] + +[[package]] +name = "anstyle-query" +version = "1.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" +dependencies = [ + "windows-sys 0.61.2", +] + +[[package]] +name = "anstyle-wincon" +version = "3.0.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" +dependencies = [ + "anstyle", + "once_cell_polyfill", + "windows-sys 0.61.2", +] + +[[package]] +name = "anyhow" +version = "1.0.104" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "330a5ed07fa54e4702c9d6c4174f74427fc0ef6e214bbd677ae50a5099946470" + [[package]] name = "atomic-waker" version = "1.1.2" @@ -60,6 +145,12 @@ dependencies = [ "tracing", ] +[[package]] +name = "base64" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9e1b586273c5702936fe7b7d6896644d8be71e6314cfe09d3167c95f712589e8" + [[package]] name = "base64" version = "0.22.1" @@ -78,6 +169,15 @@ version = "2.13.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3ded4057c258ba199e2d26386d3af3780957ecaee6c4ef4041c6b4b8b97c0b06" +[[package]] +name = "block-buffer" +version = "0.10.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71" +dependencies = [ + "generic-array", +] + [[package]] name = "bumpalo" version = "3.20.3" @@ -90,6 +190,15 @@ version = "1.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fc652a48c352aef3ea3aed32080501cf3ef6ed5da78602a020c991775b0aff04" +[[package]] +name = "castaway" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dec551ab6e7578819132c713a93c022a05d60159dc86e7a7050223577484c55a" +dependencies = [ + "rustversion", +] + [[package]] name = "cc" version = "1.4.7" @@ -119,8 +228,78 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06" dependencies = [ "cfg-if", - "cpufeatures", - "rand_core", + "cpufeatures 0.3.1", + "rand_core 0.10.1", +] + +[[package]] +name = "clap" +version = "4.6.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "aa8876b300ab35ba921adea3dfd70157a46249b33f95c9084ae5709785478946" +dependencies = [ + "clap_builder", + "clap_derive", +] + +[[package]] +name = "clap_builder" +version = "4.6.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ec0797fb7aeb1406c84efac526901f7ec3ead2124f946b494e72879d4b54704d" +dependencies = [ + "anstream", + "anstyle", + "clap_lex", + "strsim", +] + +[[package]] +name = "clap_derive" +version = "4.6.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f9c751b79415d4e559e3d1fcf128e09e720eb673a06d26cf6f392d37d75b66e0" +dependencies = [ + "heck", + "proc-macro2", + "quote", + "syn 3.0.6", +] + +[[package]] +name = "clap_lex" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1c133bc6a41be0d194c306b5506d15e6feeea7b1d6604bd3f8310dfb2ca96486" + +[[package]] +name = "colorchoice" +version = "1.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" + +[[package]] +name = "compact_str" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9dfdd1c2274d9aa354115b09dc9a901d6c5576818cdf70d14cae2bdb47df00ab" +dependencies = [ + "castaway", + "cfg-if", + "itoa", + "rustversion", + "ryu", + "serde", + "static_assertions", +] + +[[package]] +name = "cpufeatures" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280" +dependencies = [ + "libc", ] [[package]] @@ -132,6 +311,132 @@ dependencies = [ "libc", ] +[[package]] +name = "crossbeam-deque" +version = "0.8.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "622f3fc73690be383c7214310406f28a90e6edeadc3cea882f9d71e495b9711a" +dependencies = [ + "crossbeam-epoch", + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-epoch" +version = "0.9.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc74980687109a3b14c72fd458107bf0baa1da1a1a805e178d15501ba9b86d9d" +dependencies = [ + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-utils" +version = "0.8.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a31eee39dddec8330830986fcd7625edb5a24ec90ea038215273bbc3adb08ac6" + +[[package]] +name = "crunchy" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5" + +[[package]] +name = "crypto-common" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" +dependencies = [ + "generic-array", + "typenum", +] + +[[package]] +name = "darling" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc7f46116c46ff9ab3eb1597a45688b6715c6e628b5c133e288e709a29bcb4ee" +dependencies = [ + "darling_core", + "darling_macro", +] + +[[package]] +name = "darling_core" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0d00b9596d185e565c2207a0b01f8bd1a135483d02d9b7b0a54b11da8d53412e" +dependencies = [ + "fnv", + "ident_case", + "proc-macro2", + "quote", + "strsim", + "syn 2.0.119", +] + +[[package]] +name = "darling_macro" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc34b93ccb385b40dc71c6fceac4b2ad23662c7eeb248cf10d529b7e055b6ead" +dependencies = [ + "darling_core", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "dary_heap" +version = "0.3.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b1e3a325bc115f096c8b77bbf027a7c2592230e70be2d985be950d3d5e60ebe" +dependencies = [ + "serde", +] + +[[package]] +name = "derive_builder" +version = "0.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "507dfb09ea8b7fa618fcf76e953f4f5e192547945816d5358edffe39f6f94947" +dependencies = [ + "derive_builder_macro", +] + +[[package]] +name = "derive_builder_core" +version = "0.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2d5bcf7b024d6835cfb3d473887cd966994907effbe9227e8c8219824d06c4e8" +dependencies = [ + "darling", + "proc-macro2", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "derive_builder_macro" +version = "0.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ab63b0e2bf4d5928aff72e83a7dace85d7bba5fe12dcc3c5a572d78caffd3f3c" +dependencies = [ + "derive_builder_core", + "syn 2.0.119", +] + +[[package]] +name = "digest" +version = "0.10.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" +dependencies = [ + "block-buffer", + "crypto-common", +] + [[package]] name = "displaydoc" version = "0.2.7" @@ -140,9 +445,21 @@ checksum = "c6232dd377dcc64799954cbd3a9bb882e9cdc1308ccd87b1c098f1fb2eaf82a8" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] +[[package]] +name = "either" +version = "1.18.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "252afb9ae5eaa683babdc6a068b3f5726eb19e05070c731f9b2a23a7c3e8ed34" + +[[package]] +name = "equivalent" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" + [[package]] name = "errno" version = "0.3.14" @@ -153,12 +470,36 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "esaxx-rs" +version = "0.1.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d817e038c30374a4bcb22f94d0a8a0e216958d4c3dcde369b1439fec4bdda6e6" + +[[package]] +name = "fastrand" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "da7c62ceae207dd37ea5b845da6a0696c799f85e97da1ab5b7910be3c1c80223" + [[package]] name = "find-msvc-tools" version = "0.1.13" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ef25905e51abafe4dcea6c15fec58c57b601cdbd0ee53d22ea1d3016c587d39b" +[[package]] +name = "fnv" +version = "1.0.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f9eec918d3f24069decb9af1554cad7c880e2da24a9afd88aca000531ab82c1" + +[[package]] +name = "foldhash" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77ce24cb58228fbb8aa041425bb1050850ac19177686ea6e0f41a70416f56fdb" + [[package]] name = "form_urlencoded" version = "1.2.2" @@ -197,7 +538,7 @@ checksum = "9fb9654ba8355388abeb8dcb4fc62f511300867002afc858860463bdd9fe0c44" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -228,6 +569,16 @@ dependencies = [ "slab", ] +[[package]] +name = "generic-array" +version = "0.14.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" +dependencies = [ + "typenum", + "version_check", +] + [[package]] name = "getrandom" version = "0.2.17" @@ -241,6 +592,18 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "getrandom" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" +dependencies = [ + "cfg-if", + "libc", + "r-efi 5.3.0", + "wasip2", +] + [[package]] name = "getrandom" version = "0.4.3" @@ -250,11 +613,47 @@ dependencies = [ "cfg-if", "js-sys", "libc", - "r-efi", - "rand_core", + "r-efi 6.0.0", + "rand_core 0.10.1", "wasm-bindgen", ] +[[package]] +name = "half" +version = "2.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ea2d84b969582b4b1864a92dc5d27cd2b77b622a8d79306834f1be5ba20d84b" +dependencies = [ + "cfg-if", + "crunchy", + "zerocopy", +] + +[[package]] +name = "hashbrown" +version = "0.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100" +dependencies = [ + "allocator-api2", + "equivalent", + "foldhash", + "serde", + "serde_core", +] + +[[package]] +name = "hashbrown" +version = "0.17.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" + +[[package]] +name = "heck" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" + [[package]] name = "http" version = "1.5.0" @@ -444,6 +843,12 @@ dependencies = [ "zerovec", ] +[[package]] +name = "ident_case" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9e0384b61958566e926dc50660321d12159025e767c18e043daf26b70104c39" + [[package]] name = "idna" version = "1.1.0" @@ -465,12 +870,37 @@ dependencies = [ "icu_properties", ] +[[package]] +name = "indexmap" +version = "2.14.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cc4e190f5d26ca7051642629da2c52fc03bde85a03197c99408dcd291734c855" +dependencies = [ + "equivalent", + "hashbrown 0.17.1", +] + [[package]] name = "ipnet" version = "2.12.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "791930b43c0d5973160d90a8f3894509f2b273430f5c5c73b668636d0287c5c0" +[[package]] +name = "is_terminal_polyfill" +version = "1.70.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695" + +[[package]] +name = "itertools" +version = "0.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2b192c782037fadd9cfa75548310488aabdbf3d2da73885b31bd0abd03351285" +dependencies = [ + "either", +] + [[package]] name = "itoa" version = "1.0.18" @@ -494,6 +924,22 @@ version = "0.2.189" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3eaf3ede3fee6db1a4c2ee091bf8a8b4dccdc6d17f656fb07896ee72867612f2" +[[package]] +name = "libloading" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d7c4b02199fee7c5d21a5ae7d8cfa79a6ef5bb2fc834d6e9058e89c825efdc55" +dependencies = [ + "cfg-if", + "windows-link", +] + +[[package]] +name = "linux-raw-sys" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" + [[package]] name = "litemap" version = "0.8.3" @@ -512,6 +958,22 @@ version = "0.1.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4050469837a6ff301cd14c1f8f24f88549e6d548f24f64e2148eb0f72cebc51f" +[[package]] +name = "macro_rules_attribute" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b3ae8f6d608c795738406608304d30a2dfbdc8e58e44f7ba43236da5208ded3c" +dependencies = [ + "macro_rules_attribute-proc_macro", + "pastey", +] + +[[package]] +name = "macro_rules_attribute-proc_macro" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc04a4c58212d57930a24bf47d3fa87485264a3a054e9c10e042eb373573ad3c" + [[package]] name = "matchit" version = "0.8.4" @@ -524,12 +986,27 @@ version = "2.8.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98" +[[package]] +name = "memmap2" +version = "0.9.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d1219ed1b7f229ee7104d281dd01d6802fe28bb6e95d292942c4daacdeb798c0" +dependencies = [ + "libc", +] + [[package]] name = "mime" version = "0.3.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6877bb514081ee2a7ff5ef9de3281f14a4dd4bceac4c09388074a6b5df8a139a" +[[package]] +name = "minimal-lexical" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" + [[package]] name = "mio" version = "1.2.3" @@ -541,6 +1018,56 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "monostate" +version = "0.1.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3341a273f6c9d5bef1908f17b7267bbab0e95c9bf69a0d4dcf8e9e1b2c76ef67" +dependencies = [ + "monostate-impl", + "serde", + "serde_core", +] + +[[package]] +name = "monostate-impl" +version = "0.1.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e4db6d5580af57bf992f59068d4ea26fd518574ff48d7639b255a36f9de6e7e9" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "nom" +version = "7.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d273983c5a657a70a3e8f2a01329822f3b8c8172b73826411a55751e404a0a4a" +dependencies = [ + "memchr", + "minimal-lexical", +] + +[[package]] +name = "omni-cua-s1-native" +version = "0.1.0" +dependencies = [ + "anyhow", + "axum", + "clap", + "half", + "http-body-util", + "libloading", + "memmap2", + "safetensors", + "serde_json", + "sha2", + "tokenizers", + "tokio", +] + [[package]] name = "omni-jev" version = "0.1.0" @@ -557,6 +1084,46 @@ version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" +[[package]] +name = "once_cell_polyfill" +version = "1.70.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" + +[[package]] +name = "onig" +version = "6.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0cc3cbf698f9438986c11a880c90a6d04b9de27575afd28bbf45b154b6c709e2" +dependencies = [ + "bitflags", + "libc", + "once_cell", + "onig_sys", +] + +[[package]] +name = "onig_sys" +version = "69.9.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e68317604e77e53b85896388e1a803c1d21b74c899ec9e5e1112db90735edd7" +dependencies = [ + "cc", + "pkg-config", +] + +[[package]] +name = "paste" +version = "1.0.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a" + +[[package]] +name = "pastey" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2ee67f1008b1ba2321834326597b8e186293b049a023cdef258527550b9935b4" + [[package]] name = "percent-encoding" version = "2.3.2" @@ -569,6 +1136,12 @@ version = "0.2.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" +[[package]] +name = "pkg-config" +version = "0.3.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f6b464fbc74e149a392436b17d523f769e057cb6877f6a5c4618bc6f11800548" + [[package]] name = "potential_utf" version = "0.1.6" @@ -578,6 +1151,15 @@ dependencies = [ "zerovec", ] +[[package]] +name = "ppv-lite86" +version = "0.2.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85eae3c4ed2f50dcfe72643da4befc30deadb458a9b590d720cde2f2b1e97da9" +dependencies = [ + "zerocopy", +] + [[package]] name = "proc-macro2" version = "1.0.107" @@ -616,7 +1198,7 @@ dependencies = [ "bytes", "getrandom 0.4.3", "lru-slab", - "rand", + "rand 0.10.3", "rand_pcg", "ring", "rustc-hash", @@ -652,12 +1234,28 @@ dependencies = [ "proc-macro2", ] +[[package]] +name = "r-efi" +version = "5.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" + [[package]] name = "r-efi" version = "6.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" +[[package]] +name = "rand" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9ef1d0d795eb7d84685bca4f72f3649f064e6641543d3a8c415898726a57b41" +dependencies = [ + "rand_chacha", + "rand_core 0.9.5", +] + [[package]] name = "rand" version = "0.10.3" @@ -666,7 +1264,26 @@ checksum = "65c9fb96cbc91e3478eaae79a69fcd3f1ae4ad052e471fe6732fff548984b4af" dependencies = [ "chacha20", "getrandom 0.4.3", - "rand_core", + "rand_core 0.10.1", +] + +[[package]] +name = "rand_chacha" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb" +dependencies = [ + "ppv-lite86", + "rand_core 0.9.5", +] + +[[package]] +name = "rand_core" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "76afc826de14238e6e8c374ddcc1fa19e374fd8dd986b0d2af0d02377261d83c" +dependencies = [ + "getrandom 0.3.4", ] [[package]] @@ -681,9 +1298,69 @@ version = "0.10.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "caa0f4137e1c0a72f4c651489402276c8e8e1cf081f3b0ba156d2cbeef09e86a" dependencies = [ - "rand_core", + "rand_core 0.10.1", +] + +[[package]] +name = "rayon" +version = "1.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fb39b166781f92d482534ef4b4b1b2568f42613b53e5b6c160e24cfbfa30926d" +dependencies = [ + "either", + "rayon-core", +] + +[[package]] +name = "rayon-cond" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2964d0cf57a3e7a06e8183d14a8b527195c706b7983549cd5462d5aa3747438f" +dependencies = [ + "either", + "itertools", + "rayon", +] + +[[package]] +name = "rayon-core" +version = "1.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "22e18b0f0062d30d4230b2e85ff77fdfe4326feb054b9783a3460d8435c8ab91" +dependencies = [ + "crossbeam-deque", + "crossbeam-utils", ] +[[package]] +name = "regex" +version = "1.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f020237b6c8eed93db2e2cb53c00c60a8e1bc73da7d073199a1180401450218d" +dependencies = [ + "aho-corasick", + "memchr", + "regex-automata", + "regex-syntax", +] + +[[package]] +name = "regex-automata" +version = "0.4.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ad8553b9b26413251cbf30e620595c7a41b3887f03da04579c0e6b0d6a06b4b2" +dependencies = [ + "aho-corasick", + "memchr", + "regex-syntax", +] + +[[package]] +name = "regex-syntax" +version = "0.8.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" + [[package]] name = "reqwest" version = "0.12.28" @@ -745,6 +1422,19 @@ version = "2.1.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6b1e7f9a428571be2dc5bc0505c13fb6bf936822b894ec87abf8a08a4e51742d" +[[package]] +name = "rustix" +version = "1.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "891efababe418670775f199f0d233d84843c227a0949a883ce15b37c78d6629d" +dependencies = [ + "bitflags", + "errno", + "libc", + "linux-raw-sys", + "windows-sys 0.61.2", +] + [[package]] name = "rustls" version = "0.23.45" @@ -792,6 +1482,19 @@ version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" +[[package]] +name = "safetensors" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "79b079b829cb27a1c3c374341345ed2e8b2c0c839034522cee576c140bd7f846" +dependencies = [ + "hashbrown 0.16.1", + "libc", + "serde", + "serde_json", + "tempfile", +] + [[package]] name = "serde" version = "1.0.229" @@ -799,6 +1502,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4148590afebada386688f18773da617792bf2ef03ffc1e4cbd2b1d45b023e0ba" dependencies = [ "serde_core", + "serde_derive", ] [[package]] @@ -818,7 +1522,7 @@ checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -827,6 +1531,7 @@ version = "1.0.151" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c841b55ecdae098c80dcae9cf767f6f8a0c2cdb3416bbef72181df4d0fe73f14" dependencies = [ + "indexmap", "itoa", "memchr", "serde", @@ -857,6 +1562,17 @@ dependencies = [ "serde", ] +[[package]] +name = "sha2" +version = "0.10.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" +dependencies = [ + "cfg-if", + "cpufeatures 0.2.17", + "digest", +] + [[package]] name = "shlex" version = "2.0.1" @@ -895,18 +1611,53 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "spm_precompiled" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5851699c4033c63636f7ea4cf7b7c1f1bf06d0cc03cfb42e711de5a5c46cf326" +dependencies = [ + "base64 0.13.1", + "nom", + "serde", + "unicode-segmentation", +] + [[package]] name = "stable_deref_trait" version = "1.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" +[[package]] +name = "static_assertions" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a2eb9349b6444b326872e140eb1cf5e7c522154d69e7a0ffb0fb81c06b37543f" + +[[package]] +name = "strsim" +version = "0.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" + [[package]] name = "subtle" version = "2.6.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" +[[package]] +name = "syn" +version = "2.0.119" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "872831b642d1a07999a962a351ed35b955ea2cfc8f3862091e2a240a84f17297" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + [[package]] name = "syn" version = "3.0.6" @@ -935,7 +1686,20 @@ checksum = "901704edd0dfe137f1987838ee4f259e4e063c31371bdb423f7ae38ec6f77f02" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", +] + +[[package]] +name = "tempfile" +version = "3.27.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd" +dependencies = [ + "fastrand", + "getrandom 0.4.3", + "once_cell", + "rustix", + "windows-sys 0.61.2", ] [[package]] @@ -955,7 +1719,7 @@ checksum = "fe5197923287db20a58125f0bc85c062f7f2c892de97b18c356f9efb14b28524" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -974,6 +1738,39 @@ version = "1.13.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fd3ca314f692efd6c868f8408f53fe444634a845f96c028b97d35f6a1f79f0ee" +[[package]] +name = "tokenizers" +version = "0.22.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b238e22d44a15349529690fb07bd645cf58149a1b1e44d6cb5bd1641ff1a6223" +dependencies = [ + "ahash", + "aho-corasick", + "compact_str", + "dary_heap", + "derive_builder", + "esaxx-rs", + "getrandom 0.3.4", + "itertools", + "log", + "macro_rules_attribute", + "monostate", + "onig", + "paste", + "rand 0.9.5", + "rayon", + "rayon-cond", + "regex", + "regex-syntax", + "serde", + "serde_json", + "spm_precompiled", + "thiserror", + "unicode-normalization-alignments", + "unicode-segmentation", + "unicode_categories", +] + [[package]] name = "tokio" version = "1.53.1" @@ -998,7 +1795,7 @@ checksum = "78773a2a397f451582ce068015985c33193cf6dea8b74d2a639fe457b2f07b0e" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -1096,12 +1893,39 @@ version = "0.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" +[[package]] +name = "typenum" +version = "1.20.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6f5e870be6c3b371b77fe0ee0bafb859fa4964b4404c27de1d380043c4dda20" + [[package]] name = "unicode-ident" version = "1.0.26" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d245f478577f809a851594d02313b640fb437e0bb33866753cff937863096954" +[[package]] +name = "unicode-normalization-alignments" +version = "0.1.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "43f613e4fa046e69818dd287fdc4bc78175ff20331479dab6e1b0f98d57062de" +dependencies = [ + "smallvec", +] + +[[package]] +name = "unicode-segmentation" +version = "1.13.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6f5d3c3b1bf09027a88a6bc961fc00497d651009560b5463668dc81b0fa87a8" + +[[package]] +name = "unicode_categories" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39ec24b3121d976906ece63c9daad25b85969647682eee313cb5779fdd69e14e" + [[package]] name = "untrusted" version = "0.9.0" @@ -1126,6 +1950,18 @@ version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" +[[package]] +name = "utf8parse" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" + +[[package]] +name = "version_check" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" + [[package]] name = "want" version = "0.3.1" @@ -1141,6 +1977,15 @@ version = "0.11.1+wasi-snapshot-preview1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" +[[package]] +name = "wasip2" +version = "1.0.4+wasi-0.2.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b67efb37e106e55ce722a510d6b5f9c17f083e5fc79afc2badeb12cc313d9487" +dependencies = [ + "wit-bindgen", +] + [[package]] name = "wasm-bindgen" version = "0.2.129" @@ -1184,7 +2029,7 @@ dependencies = [ "bumpalo", "proc-macro2", "quote", - "syn", + "syn 3.0.6", "wasm-bindgen-shared", ] @@ -1327,6 +2172,12 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" +[[package]] +name = "wit-bindgen" +version = "0.57.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e" + [[package]] name = "writeable" version = "0.6.4" @@ -1352,10 +2203,30 @@ checksum = "33811428bee40dbceb6d545e95754741d17a6aef9a4849f0fd62e2ba4f412a78" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", "synstructure", ] +[[package]] +name = "zerocopy" +version = "0.8.59" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6df92bf3d9227be3d53173901ddbffac2babc27ae50f397776ffd6dc33f800cb" +dependencies = [ + "zerocopy-derive", +] + +[[package]] +name = "zerocopy-derive" +version = "0.8.59" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ac4f328cf2f05d084e496c3e9c3f33ed0a183656a16e1fcec4d464d8373aec82" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "zerofrom" version = "0.1.8" @@ -1373,7 +2244,7 @@ checksum = "f75b4683f6c7f45248d4d64056a24298c6281e0993356d7d1b4a1a962ef10d4a" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", "synstructure", ] @@ -1413,7 +2284,7 @@ checksum = "34df6fc39dbd26ddc9c10e6a2984476e13acce22e64e4487636ef494369225da" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 8036e4bc..8c100b4b 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,3 +1,3 @@ [workspace] -members = ["src/frontend"] +members = ["src/frontend", "src/models/cua_s1/native"] resolver = "3" diff --git a/README.md b/README.md index f029b38e..402249fb 100644 --- a/README.md +++ b/README.md @@ -2,7 +2,7 @@ A community-maintained inference engine for prefill-only System1-Omni models, designed around a Rust frontend, model-owned execution, and high-performance CUDA and Metal backends. -The Rust frontend forwards requests to a separately running model worker. In-repository model engines and GPU backends are not implemented yet. +The Rust frontend forwards requests to a separately running model worker. The Cua-S1 4B 0.2 `text` adapter has a native worker with CUDA kernels in this repository; other in-repository model engines and GPU backends are not implemented yet. ## Run the frontend @@ -48,15 +48,16 @@ Implementation code lives under `src/`; recipes and documentation stay at the re | [`recipe/`](recipe/) | Model setup instructions, launch commands, configuration examples, and example requests. | | [`docs/`](docs/) | Project documentation and architecture assets. | -The frontend is a Cargo workspace member. Model and backend directories currently document planned work; they do not prescribe process boundaries. +The frontend and the Cua-S1 native worker are Cargo workspace members. The other model and backend directories currently document planned work; they do not prescribe process boundaries. ## Supported models -LAYA can run as an external Python worker for text requests. Its in-repository model engine is still planned: +LAYA can run as an external Python worker for text requests; its in-repository model engine is still planned. The Cua-S1 4B 0.2 `text` adapter runs as a Python worker or as a native worker on CUDA: | Model | Status | | --- | --- | | LAYA | [External worker](recipe/laya/README.md); model engine 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 875486cb..2c7d1ac7 100644 --- a/recipe/README.md +++ b/recipe/README.md @@ -4,6 +4,9 @@ Rust frontend and compare direct and proxied responses. - [Cua-S1 4B 0.2 text worker](cua_s1/text.md): download the pinned weights, start the worker, connect the Rust frontend and check the worker against upstream. +- [Cua-S1 4B 0.2 native text worker](cua_s1/native.md): build the CUDA library and + the Rust worker, export the merged weights, and check the worker against the + reference worker. Recipes contain setup, launch commands and examples. Reusable implementation code belongs under `src/`. diff --git a/recipe/cua_s1/check_native.py b/recipe/cua_s1/check_native.py new file mode 100644 index 00000000..e6732054 --- /dev/null +++ b/recipe/cua_s1/check_native.py @@ -0,0 +1,58 @@ +"""Check `omni-cua-s1-native --score-all` output against the float32 reference. + + python recipe/cua_s1/check_native.py scores.jsonl [scores from another start.jsonl] + +The parity directory holds parity_float32.jsonl and parity_bfloat16.jsonl, written by +compare_text_with_upstream.py --dtype float32 / bfloat16 --out (see native.md). + +The rule is the one src/models/cua_s1/README.md declares for a native engine: the +largest per-option difference from the float32 worker is at most 2 x (bfloat16 +worker vs float32) + 0.01, and the top option matches float32 wherever float32's +top-two margin is at least 0.05. The served result (from a graph when the +prompt fits) must be bitwise identical to the eager one. With a second file (another +process start), both runs must be bitwise identical too. +""" +import json +import sys +from pathlib import Path + + +def load_parity(path): + rows = [json.loads(l) for l in path.read_text().splitlines() if l] + return {(r["case"], r["question"]): r["worker"] for r in rows} + + +def load(path): + return [json.loads(l) for l in Path(path).read_text().splitlines() if l] + + +runs = load(sys.argv[1]) +parity = Path(sys.argv[2]) +fp32, bf16 = load_parity(parity / "parity_float32.jsonl"), load_parity(parity / "parity_bfloat16.jsonl") +diff = lambda a, b: max(abs(a[k] - b[k]) for k in b) +allowance = 2 * max(diff(bf16[k], fp32[k]) for k in fp32) + 0.01 +eager = {(r["case"], r["question"]): r for r in runs if r["mode"] == "eager"} +served = {(r["case"], r["question"]): r for r in runs if r["mode"] == "served"} +missing = sorted(set(fp32) - set(eager)) +worst, worst_at, flips = 0.0, None, [] +for key, r in eager.items(): + ref, got = fp32[key], r["probabilities"] + d = diff(got, ref) + if d > worst: + worst, worst_at = d, f"{key[0]}/{key[1]}" + top = sorted(ref.values(), reverse=True) + margin = top[0] - (top[1] if len(top) > 1 else 0.0) + if margin >= 0.05 and max(ref, key=ref.get) != max(got, key=got.get): + flips.append(f"{key[0]}/{key[1]}") +graph_keys = [k for k, r in served.items() if r["graph"]] +mismatch = [f"{k[0]}/{k[1]}" for k in served if served[k]["probabilities"] != eager[k]["probabilities"]] +print(f"{len(eager)} questions; allowance {allowance:.4f}") +print(f"largest |eager - fp32| {worst:.4f} ({worst_at}); top-option changes: {flips or 'none'}; missing: {missing or 'none'}") +print(f"served from a graph: {len(graph_keys)}; served differs from eager: {mismatch or 'none'}") +ok = worst <= allowance and not flips and not missing and not mismatch +if len(sys.argv) > 3: + other = load(sys.argv[3]) + same = len(other) == len(runs) and all(a == b for a, b in zip(runs, other)) + print(f"identical to the other start: {same}") + ok = ok and same +print("PASS" if ok else "FAIL") diff --git a/recipe/cua_s1/diff_corpus.py b/recipe/cua_s1/diff_corpus.py new file mode 100644 index 00000000..e9343a94 --- /dev/null +++ b/recipe/cua_s1/diff_corpus.py @@ -0,0 +1,249 @@ +"""Request bodies for comparing the Python and native Cua-S1 text workers. + + python recipe/cua_s1/diff_corpus.py tests/cua_s1/data/text_inputs.json corpus.jsonl [n_fuzz] + +Each line is {"name": ..., "body": }. The set has +the fixed input set, hand-written edge cases for every error path of +`contract.parse_body` / `map_request` / the server, and seeded random bodies. +""" + +from __future__ import annotations + +import base64 +import json +import random +import sys + +M = "cua-s1-4b-0.2" + + +def req(state="Button: OK", questions=None, model=M, **extra): + body = {"model": model, "state": state} + body["questions"] = questions if questions is not None else {"q": choice()} + body.update(extra) + return body + + +def choice(instructions="Press OK.", criteria=None, type_="choice"): + q = {"type": type_, "instructions": instructions} + q["criteria"] = criteria if criteria is not None else {"ok": "OK", "cancel": "Cancel"} + return q + + +def build(inputs_path: str, n_fuzz: int) -> list[tuple[str, bytes]]: + cases: list[tuple[str, bytes]] = [] + + def add(name, body): + if isinstance(body, (dict, list)): + body = json.dumps(body, ensure_ascii=False) + if isinstance(body, str): + body = body.encode("utf-8", "surrogatepass") + cases.append((name, body)) + + fixtures = json.load(open(inputs_path, encoding="utf-8")) + for name, body in fixtures.items(): + add(f"fixture/{name}", body) + add(f"fixture_ascii_compact/{name}", json.dumps(body, ensure_ascii=True, separators=(",", ":"))) + + # --- JSON syntax and decoding --- + ok = json.dumps(req()) + raw = { + "empty": "", "space": " ", "open": "{", "close": "}", "array": "[]", "null": "null", + "number": "1", "string": '"x"', "true": "true", "extra": ok + " x", "bom": "" + ok, + "ws": " \t\n" + ok + "\r\n", "trailing_comma_obj": ok[:-1] + ",}", + "single_quotes": ok.replace('"', "'"), "comment": "// c\n" + ok, + "nan": ok.replace('"Button: OK"', "NaN"), "inf": ok.replace('"Button: OK"', "Infinity"), + "neg_inf": ok.replace('"Button: OK"', "-Infinity"), "nan_nested": ok.replace('"OK"', "[1, NaN]"), + "float_range": ok.replace('"OK"', "[1e400]"), "float_range_neg": ok.replace('"OK"', "[-1E309]"), + "float_range_then_garbage": ok.replace('"OK"', "[1e400x]"), + "float_underflow": ok.replace('"OK"', "[1e-400, 1.0e+308, -0.0, 5e-324]"), + "int_4300": ok.replace('"OK"', "[" + "9" * 4300 + "]"), + "int_4301": ok.replace('"OK"', "[" + "9" * 4301 + "]"), + "int_neg_4301": ok.replace('"OK"', "[-" + "9" * 4301 + "]"), + "leading_zero": ok.replace('"OK"', "[01]"), "dot_no_digit": ok.replace('"OK"', "[1.]"), + "exp_no_digit": ok.replace('"OK"', "[1e]"), "exp_sign_no_digit": ok.replace('"OK"', "[1e+]"), + "minus_alone": ok.replace('"OK"', "[-]"), "plus": ok.replace('"OK"', "[+1]"), + "hex": ok.replace('"OK"', "[0x10]"), + "raw_control": ok.replace("Press OK.", "Press\u0001OK."), "raw_tab": ok.replace("Press OK.", "Press\tOK."), + "raw_del": ok.replace("Press OK.", "Press\u007fOK."), + "bad_escape": ok.replace("Press OK.", "Press \\x OK."), "short_u": ok.replace("Press OK.", "\\u12"), + "bad_u": ok.replace("Press OK.", "\\uZZZZ"), + "escapes": ok.replace("Press OK.", "\\/\\b\\f\\n\\r\\t\\\"\\\\ \\u00e9\\u0000"), + "missing_colon": ok.replace('"model":', '"model"'), "unterminated": ok[:-3], + "dup_top": '{"model": "' + M + '", "model": "' + M + '", "state": "s", "questions": {}}', + "dup_criteria": ok.replace('"cancel": "Cancel"', '"ok": "Again"'), + "dup_then_syntax": ok.replace('"cancel": "Cancel"', '"ok": "Again"') + "x", + "syntax_inside_dup_object": ok.replace('"cancel": "Cancel"', '"ok": "Again", x'), + "dup_state_nested": ok.replace('"Button: OK"', '{"a": {"b": 1, "b": 2}}'), + "dup_after_nan": ok.replace('"cancel": "Cancel"', '"ok": NaN'), + "lone_high": ok.replace("Press OK.", "\\ud800"), "lone_low": ok.replace("Press OK.", "\\udc00x"), + "pair": ok.replace("Press OK.", "\\ud83d\\ude00"), "reversed_pair": ok.replace("Press OK.", "\\ude00\\ud83d"), + "high_then_bmp": ok.replace("Press OK.", "\\ud800\\u0041"), + "lone_in_key": ok.replace('"ok": "OK"', '"\\udfff": "OK"'), + "lone_dup_key": ok.replace('"ok": "OK", "cancel": "Cancel"', '"\\ud800": 1, "\\ud800": 2'), + "lone_then_syntax": ok.replace("Press OK.", "\\ud800") + ",", + } + for name, text in raw.items(): + add(f"json/{name}", text) + add("json/invalid_utf8", ok.encode().replace(b"Press", b"Pr\xffss")) + add("json/overlong_utf8", ok.encode().replace(b"Press", b"Pr\xc0\xafss")) + add("json/utf8_surrogate", ok.encode().replace(b"Press", b"Pr\xed\xa0\x80ss")) + add("json/utf8_bom_bytes", b"\xef\xbb\xbf" + ok.encode()) + add("json/latin1", ok.encode().replace(b"Press", b"Pr\xe9ss")) + + def nested(d, leaf='"x"', open_="[", close="]"): + return open_ * d + leaf + close * d + + for d in (1, 50, 966, 967, 968, 969, 1500): + add(f"depth/state_list_{d}", ok.replace('"Button: OK"', nested(d))) + for d in (964, 965, 966): + add(f"depth/criteria_{d}", ok.replace('"OK"', nested(d))) + add(f"depth/type_{d}", ok.replace('"choice"', nested(d))) + add(f"depth/instructions_obj_{d}", ok.replace('"Press OK."', nested(d, '"x"', '{"a":', "}"))) + + head = '{"model": "cua-s1-4b-0.2", "state": ' + tail = ', "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "A"}}}}' + for d in (967, 968, 969, 970): + add(f"depth2/empty_list_leaf_{d}", head + nested(d, "") + tail) + add(f"depth2/empty_obj_leaf_{d}", head + nested(d, "{}") + tail) + add(f"depth2/number_leaf_{d}", head + nested(d, "1") + tail) + for d in (9988, 9989, 9990, 9991): + add(f"depth2/parse_limit_garbage_{d}", head + nested(d, "1") + tail + "x") + add("depth2/deep_then_dup", head + nested(5000, "1") + ', "state": 1' + tail) + add("depth2/deep_then_nan", head + nested(3000, "1") + ', "z": NaN' + tail) + add("depth2/surrogate_before_deep", '{"model": "\\ud800", "state": ' + nested(1500, "1") + tail) + add("depth2/deep_before_surrogate", head + nested(1500, "1") + ', "z": "\\ud800"' + tail) + add("depth2/deep_before_surrogate_key", '{"a": ' + nested(1500, "1") + ', "\\ud800": 1}') + add("depth2/hundred_thousand", "[" * 100000) + add("depth2/deep_top_level_array", nested(5000, "1")) + sa = {"_sample": choice(criteria={"_save": "Save", "_sa": None, "b": "B"}), "q": choice()} + add("keys/_sa_names_and_options", req(questions=sa)) + add("keys/_sa_state_keys", req(state={"_sa": 1, "_sample": [2]})) + + # --- mapping errors, in Python's check order --- + sem = { + "no_model": {"state": "s", "questions": {"q": choice()}}, + "model_null": req(model=None), "model_case": req(model="Cua-S1-4B-0.2"), + "model_space": req(model=M + " "), "model_number": req(model=1), + "model_wrong_and_no_state": {"model": "x"}, + "no_state": {"model": M, "questions": {"q": choice()}}, + "state_null": req(state=None), "state_true": req(state=True), "state_zero": req(state=0), + "state_float": req(state=1.5), "state_empty": req(state=""), "state_empty_obj": req(state={}), + "state_empty_list": req(state=[]), "state_spaces": req(state=" "), + "state_obj": req(state={"a": None, "b": [1, 2.5, -0.0, 1e22, 12345678901234567890, True], "é": "
"}), + "state_list": req(state=["Button: OK", {"k": "v"}, 3.14e-07]), + "no_questions": {"model": M, "state": "s"}, "questions_null": req(questions=None) | {"questions": None}, + "questions_list": req(questions=[]), "questions_empty": req(questions={}), + "questions_str": req(questions="q"), + "questions_65": req(questions={f"q{i}": choice() for i in range(65)}), + "questions_65_bad_types": req(questions={f"q{i}": {"type": "score"} for i in range(65)}), + "questions_12": req(questions={f"q{i}": choice(f"Pick {i}.") for i in range(12)}), + "question_list": req(questions={"q": []}), "question_null": req(questions={"q": None}), + "question_str": req(questions={"q": "choice"}), + "score_second": req(questions={"a": {"type": "choice"}, "b": {"type": "score"}}), + "noul": req(questions={"a": {"type": "noul"}}), + "type_missing": req(questions={"a": {"instructions": "x", "criteria": {"a": "A"}}}), + "missing_instructions_then_score": req(questions={"a": {"type": "choice"}, "b": {"type": "noul"}}), + "no_instructions": req(questions={"a": {"type": "choice", "criteria": {"x": "X"}}}), + "instructions_null": req(questions={"a": choice(None)}), + "instructions_empty": req(questions={"a": choice("")}), + "instructions_zero": req(questions={"a": choice(0)}), + "instructions_false": req(questions={"a": choice(False)}), + "instructions_obj": req(questions={"a": choice({"x": 1.5e-7, "y": [None]})}), + "instructions_empty_obj": req(questions={"a": choice({})}), + "instructions_list": req(questions={"a": choice([])}), + "no_criteria": req(questions={"a": {"type": "choice", "instructions": "x"}}), + "criteria_null": req(questions={"a": choice(criteria=None) | {"criteria": None}}), + "criteria_list": req(questions={"a": choice(criteria=[])}), + "criteria_empty": req(questions={"a": choice(criteria={})}), + "criteria_26": req(questions={"a": choice(criteria={f"o{i}": f"Option {i}" for i in range(26)})}), + "criteria_27": req(questions={"a": choice(criteria={f"o{i}": f"Option {i}" for i in range(27)})}), + "criteria_values": req(questions={"a": choice(criteria={ + "n": None, "o": {"k": [1, 2]}, "l": [], "e": {}, "q": 'quote"s', "b": "back\\slash", + "nl": "new\nline", "t": "tab\t", "z": "\u0000", "ls": "
", "é": "ünï 中文 😀"})}), + "criteria_bool": req(questions={"a": choice(criteria={"x": "X", "y": True})}), + "criteria_number": req(questions={"a": choice(criteria={"x": 1})}), + "second_question_bad": req(questions={"a": choice(), "b": choice(criteria={})}), + } + for name, body in sem.items(): + add(f"map/{name}", body) + + weird_names = ["it's", 'say "hi"', "both'\"", "back\\slash", "new\nline", "nul\u0000", "del\u007f", + "nbsp ", "ls
", "zw​", "tag\U000e0001", "pua", "unassigned͸", + "emoji😀", "é", "combining é", "rtl א", "soft­", "space ", "", "\t"] + for i, name in enumerate(weird_names): + add(f"names/type_{i}", req(questions={name: {"type": name}})) + add(f"names/option_{i}", req(questions={name: choice(criteria={name: 1})})) + add(f"names/ok_{i}", req(questions={name: choice(criteria={name: None, "other": "Other"})})) + for i, kind in enumerate(["", "Choice", 1, 1.5, 1e16, -0.0, 1e-5, True, None, [], {}, {"a": [1, {"b": None}]}, + 12345678901234567890, 2.5e-310]): + add(f"types/{i}", req(questions={"q": {"type": kind}})) + + # --- prompt length --- + long_state = "Row: item " * 3000 # a little over 16384 tokens + add("limit/prompt_too_long", req(state=long_state)) + add("limit/second_prompt_too_long", req(questions={"a": choice(), "b": choice("x " * 17000)})) + add("limit/prompt_just_under", req(state="Row: item " * 2600)) + + # --- seeded random bodies --- + rnd = random.Random(0) + pool = ("abcXYZ019 _-.:/\\\"'\n\t\r{}[]<>|é中文😀
​ ́א\u0000\u001f\u007f" + "<|im_start|><|im_end|>") + + def rstr(n=12): + return "".join(rnd.choice(pool) for _ in range(rnd.randint(0, n))) + + def rnum(): + k = rnd.random() + if k < 0.3: + return rnd.randint(-10**rnd.randint(1, 30), 10**rnd.randint(1, 30)) + if k < 0.9: + return rnd.uniform(-1, 1) * 10 ** rnd.randint(-320, 300) + return rnd.choice([0.0, -0.0, 5e-324, 1e16, 1e-5, 0.1, 1 / 3]) + + def rval(depth=0): + k = rnd.random() + if depth > 3 or k < 0.35: + return rstr() + if k < 0.5: + return rnum() + if k < 0.55: + return rnd.choice([None, True, False]) + if k < 0.75: + return [rval(depth + 1) for _ in range(rnd.randint(0, 4))] + return {rstr(6): rval(depth + 1) for _ in range(rnd.randint(0, 4))} + + for i in range(n_fuzz): + qs = {} + for _ in range(rnd.randint(1, 3)): + crit = {rstr(8) or "k": rnd.choice([rstr(), None, rval(), rval()]) for _ in range(rnd.randint(1, 6))} + qs[rstr(8)] = {"type": "choice" if rnd.random() < 0.9 else rval(), + "instructions": rnd.choice([rstr(40), None, rval(), ""]), "criteria": crit} + body = req(state=rnd.choice([rstr(200), rval(), rval()]), questions=qs) + text = json.dumps(body, ensure_ascii=rnd.random() < 0.3) + add(f"fuzz/{i}", text) + # byte-level damage to a copy + b = bytearray(text.encode("utf-8", "surrogatepass")) + for _ in range(rnd.randint(1, 3)): + op, pos = rnd.random(), rnd.randrange(len(b)) + if op < 0.4: + del b[pos] + elif op < 0.8: + b.insert(pos, rnd.choice(b'{}[],:"\\ 0e-.\x00\xff')) + else: + b[pos] = rnd.randrange(256) + add(f"fuzz_damaged/{i}", bytes(b)) + return cases + + +def main() -> None: + n_fuzz = int(sys.argv[3]) if len(sys.argv) > 3 else 150 + cases = build(sys.argv[1], n_fuzz) + with open(sys.argv[2], "w") as f: + for name, body in cases: + f.write(json.dumps({"name": name, "body": base64.b64encode(body).decode()}) + "\n") + print(f"{len(cases)} bodies") + + +if __name__ == "__main__": + main() diff --git a/recipe/cua_s1/diff_workers.py b/recipe/cua_s1/diff_workers.py new file mode 100644 index 00000000..19ce7706 --- /dev/null +++ b/recipe/cua_s1/diff_workers.py @@ -0,0 +1,149 @@ +"""Send every corpus body to the Python and the native worker and compare the answers. + + python recipe/cua_s1/diff_workers.py corpus.jsonl out.jsonl + +Errors must match exactly: status, content type and body bytes. For answers, the +model identity, usage, question and option order and the answer type must match; +probabilities are compared numerically, and the choice may differ only where the +Python worker's top-two margin is under 0.05. +""" + +from __future__ import annotations + +import base64 +import http.client +import json +import sys +import time + + +def request(port, method, path, body=None, headers=None, chunked=False): + conn = http.client.HTTPConnection("127.0.0.1", port, timeout=600) + headers = dict(headers or {}) + if body is not None and not chunked: + headers.setdefault("content-type", "application/json") + started = time.perf_counter() + if chunked: + conn.putrequest(method, path) + conn.putheader("transfer-encoding", "chunked") + conn.putheader("content-type", "application/json") + conn.endheaders() + step = 1 << 16 + try: + for i in range(0, len(body), step): + part = body[i : i + step] + conn.send(f"{len(part):x}\r\n".encode() + part + b"\r\n") + conn.send(b"0\r\n\r\n") + except (BrokenPipeError, ConnectionResetError): + pass + else: + try: + conn.request(method, path, body=body, headers=headers) + except (BrokenPipeError, ConnectionResetError): + pass # the server answered before reading the whole body + try: + r = conn.getresponse() + except (ConnectionResetError, http.client.RemoteDisconnected) as e: + conn.close() + return -1, None, repr(e).encode(), (time.perf_counter() - started) * 1000 + data = r.read() + ms = (time.perf_counter() - started) * 1000 + ctype = r.getheader("content-type") + conn.close() + return r.status, ctype, data, ms + + +def ordered(data: bytes): + return json.loads(data, object_pairs_hook=lambda pairs: pairs) + + +def compare(name, py, rs): + """Return (ok, note, max_prob_diff).""" + ps, pc, pb, _ = py + rs_, rc, rb, _ = rs + if ps != 200 or rs_ != 200: + same = (ps, pc, pb) == (rs_, rc, rb) + return same, "" if same else f"python {ps} {pb[:200]!r} | native {rs_} {rb[:200]!r}", 0.0 + p, r = ordered(pb), ordered(rb) + pd, rd = dict(p), dict(r) + notes, worst = [], 0.0 + if [k for k, _ in p] != [k for k, _ in r]: + notes.append("top-level keys differ") + if pd["model"] != rd["model"] or pd["usage"] != rd["usage"]: + notes.append(f"model/usage differ: {pd['model']} {pd['usage']} vs {rd['model']} {rd['usage']}") + pa, ra = pd["answers"], rd["answers"] + if [k for k, _ in pa] != [k for k, _ in ra]: + notes.append("question order differs") + for (qn, pans), (_, rans) in zip(pa, ra): + pans, rans = dict(pans), dict(rans) + pp, rp = pans["probabilities"], rans["probabilities"] + if [k for k, _ in pp] != [k for k, _ in rp] or pans["type"] != rans["type"]: + notes.append(f"{qn}: option keys or type differ") + continue + pv, rv = [v for _, v in pp], [v for _, v in rp] + worst = max(worst, max(abs(a - b) for a, b in zip(pv, rv))) + top = sorted(pv, reverse=True) + margin = top[0] - (top[1] if len(top) > 1 else 0.0) + if pans["choice"] != rans["choice"] and margin >= 0.05: + notes.append(f"{qn}: choice {pans['choice']!r} vs {rans['choice']!r} (margin {margin:.3f})") + return not notes, "; ".join(notes), worst + + +def main() -> None: + corpus, pport, rport, out_path = sys.argv[1], int(sys.argv[2]), int(sys.argv[3]), sys.argv[4] + cases = [json.loads(l) for l in open(corpus)] + extra = [] + max_body = 4 << 20 + big = b'{"model": "cua-s1-4b-0.2", "state": "' + b"x" * (max_body + 1) + b'"}' + extra.append(("http/too_large_content_length", "POST", "/v1/systemone", big, False)) + extra.append(("http/too_large_chunked", "POST", "/v1/systemone", big, True)) + fill = max_body - len(b'{"model": "cua-s1-4b-0.2", "state": "", "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "A"}}}}') + exact = b'{"model": "cua-s1-4b-0.2", "state": "' + b"ab " * (fill // 3) + b"a" * (fill % 3) + b'", "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "A"}}}}' + assert len(exact) == max_body, len(exact) + extra.append(("http/exactly_max_body", "POST", "/v1/systemone", exact, False)) + extra.append(("http/chunked_small", "POST", "/v1/systemone", base64.b64decode(cases[0]["body"]), True)) + extra.append(("http/get_systemone", "GET", "/v1/systemone", None, False)) + extra.append(("http/post_health", "POST", "/health", b"{}", False)) + extra.append(("http/not_found", "GET", "/nope", None, False)) + + results, failures, worst, times = [], [], 0.0, {"py": [], "rs": []} + items = [(c["name"], "POST", "/v1/systemone", base64.b64decode(c["body"]), False) for c in cases] + extra + for i, (name, method, path, body, chunked) in enumerate(items): + py = request(pport, method, path, body, chunked=chunked) + rs = request(rport, method, path, body, chunked=chunked) + ok, note, diff = compare(name, py, rs) + worst = max(worst, diff) + if py[0] == 200 and rs[0] == 200: + times["py"].append(py[3]) + times["rs"].append(rs[3]) + results.append({"name": name, "ok": ok, "note": note, "python_status": py[0], + "native_status": rs[0], "max_prob_diff": diff, "python_ms": py[3], "native_ms": rs[3]}) + if not ok: + failures.append((name, note)) + if (i + 1) % 100 == 0: + print(f"{i + 1}/{len(items)} done, {len(failures)} failures", flush=True) + + hp = request(pport, "GET", "/health") + hr = request(rport, "GET", "/health") + hpj, hrj = json.loads(hp[2]), json.loads(hr[2]) + health_note = {k: (hpj.get(k), hrj.get(k)) for k in {**hpj, **hrj} if hpj.get(k) != hrj.get(k)} + + with open(out_path, "w") as f: + for r in results: + f.write(json.dumps(r, ensure_ascii=False) + "\n") + statuses = {} + for r in results: + statuses[r["python_status"]] = statuses.get(r["python_status"], 0) + 1 + print(f"{len(results)} requests, python statuses {statuses}") + print(f"mismatches: {len(failures)}") + for name, note in failures[:40]: + print(f"- {name}: {note}") + print(f"largest probability difference on answered requests: {worst:.4f}") + if times["py"]: + s = lambda v: sorted(v)[len(v) // 2] + print(f"answered requests: python median {s(times['py']):.1f} ms, native median {s(times['rs']):.1f} ms") + print(f"/health differences (python, native): {health_note}") + + +if __name__ == "__main__": + main() diff --git a/recipe/cua_s1/export_text_merged.py b/recipe/cua_s1/export_text_merged.py new file mode 100644 index 00000000..da38237c --- /dev/null +++ b/recipe/cua_s1/export_text_merged.py @@ -0,0 +1,83 @@ +"""Export Qwen3.5-4B with the Cua-S1 `text` adapter merged, for the native worker. + + PYTHONPATH=src .venv/bin/python recipe/cua_s1/export_text_merged.py \ + --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text \ + --out weights/cua-s1-4b-0.2-text-merged + +Run it in the environment of the text worker (recipe/cua_s1/text.md). It loads the +model as that worker does, merges the adapter into the bfloat16 weights with PEFT's +`merge_and_unload` on the given device, and writes the checkpoint and tokenizer +files to --out, plus `cua_s1_export.json` with the revisions it was made from (read +from the Hugging Face download metadata) and the SHA-256 of the tokenizer.json it +wrote. The native worker reads that record and checks the tokenizer against it: +Transformers writes the pre-tokenizer rule it actually uses into this file, which +is not the one in the base repository's tokenizer.json. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import time +from pathlib import Path + +import peft +import torch +import transformers + +from models.cua_s1.text.adapter import downloaded_revision +from models.cua_s1.text.contract import ADAPTER_REPO, BASE_REPO, LETTERS +from models.cua_s1.text.engine import TextEngine + + +def base_revision(base: Path) -> str | None: + """The commit Hugging Face recorded when it downloaded config.json.""" + meta = base / ".cache/huggingface/download/config.json.metadata" + return meta.read_text().splitlines()[0].strip() if meta.exists() else None + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("--base", required=True, type=Path) + parser.add_argument("--adapter", required=True, type=Path) + parser.add_argument("--out", required=True, type=Path) + parser.add_argument("--device", default="cuda") + args = parser.parse_args() + + started = time.perf_counter() + engine = TextEngine(str(args.base), str(args.adapter), args.device, "bfloat16") + model = engine.model.merge_and_unload().eval() + model.save_pretrained(args.out, safe_serialization=True, max_shard_size="5GB") + engine.tokenizer.save_pretrained(args.out) + tokenizer = (args.out / "tokenizer.json").read_bytes() + record = { + "format": "cua-s1-text-merged/1", + "base": {"repo": BASE_REPO, "revision": base_revision(args.base)}, + "adapter": { + "repo": ADAPTER_REPO, + "revision": downloaded_revision(args.adapter), + "subfolder": "text", + }, + "merge": { + "method": "peft merge_and_unload", + "device": args.device, + "dtype": "bfloat16", + "torch": torch.__version__, + "transformers": transformers.__version__, + "peft": peft.__version__, + }, + "letters": LETTERS, + "tokenizer": { + "file": "tokenizer.json", + "sha256": hashlib.sha256(tokenizer).hexdigest(), + "saved_by": f"transformers {transformers.__version__}", + }, + } + (args.out / "cua_s1_export.json").write_text(json.dumps(record, indent=2) + "\n") + print(f"exported to {args.out} in {time.perf_counter() - started:.1f} s") + print(json.dumps(record)) + + +if __name__ == "__main__": + main() diff --git a/recipe/cua_s1/native.md b/recipe/cua_s1/native.md new file mode 100644 index 00000000..38073854 --- /dev/null +++ b/recipe/cua_s1/native.md @@ -0,0 +1,111 @@ +# Cua-S1 4B 0.2 native text worker + +The native worker ([`src/models/cua_s1/native/`](../../src/models/cua_s1/native/)) serves the `text` adapter like the reference worker in [`text.md`](text.md), with the forward pass on the CUDA kernels in [`src/backends/cuda/qwen3_5/`](../../src/backends/cuda/qwen3_5/). It needs an NVIDIA GPU with compute capability 8.0 or newer and was measured on an RTX 6000 Ada (sm_89). + +Run the commands from the repository root. The reference worker's setup from `text.md` is needed once, to export the merged weights and for the checks. + +## Build + +```sh +src/backends/cuda/qwen3_5/build.sh target/release 89 # needs nvcc and cuBLASLt +cargo build --release --locked -p omni-cua-s1-native +``` + +Pass your GPU's compute capability to `build.sh` (89 for Ada, 80 for A100, 90 for H100). Only CUDA 13.2 on sm_89 has been run. The worker finds `libqwen3_5_cuda.so` next to its executable; `--cuda-lib` (or `CUA_S1_CUDA_LIB`) points elsewhere. + +## Export the merged weights + +The worker loads Qwen3.5-4B with the `text` adapter already merged into the bfloat16 weights. With the weights downloaded as in `text.md`, in the reference worker's environment: + +```sh +PYTHONPATH=src .venv/bin/python recipe/cua_s1/export_text_merged.py \ + --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text \ + --out weights/cua-s1-4b-0.2-text-merged +``` + +This writes about 8.5 GB: the checkpoint, the tokenizer files, and `cua_s1_export.json`, which records the revisions and the SHA-256 of `tokenizer.json`. The worker refuses a `tokenizer.json` that does not match it, and warns if the revisions are not the pinned ones. + +## Start the worker + +```sh +target/release/omni-cua-s1-native --model weights/cua-s1-4b-0.2-text-merged \ + --gemm-plans weights/gemm-plans.json --gemm-search --port 8000 +``` + +The first start tunes the GEMMs for this GPU and writes the choices to `--gemm-plans`; with `--gemm-search` that takes about a minute. Later starts read the file and are ready in about 2 seconds, and give the same results each time. Without `--gemm-search`, tuning takes a few seconds and the worker is somewhat slower for prompts of 200 to 2,000 tokens. The file records the GPU, the cuBLASLt version and `--graph-max-tokens`; the worker refuses a file that does not match and leaves it alone, so keep one per GPU model and CUDA version, and remove it to tune again. + +The same options exist as environment variables (`CUA_S1_MODEL`, `CUA_S1_PORT`, `CUA_S1_GEMM_PLANS`, ...; see `--help`), as do the reference worker's request limits and `CUA_S1_API_KEY`. The Rust frontend and the requests are as in `text.md`. + +## Check it + +Tests without a GPU, then the kernel checks (attention against a float32 kernel, the gated delta rule against a float64 token-by-token reference): + +```sh +cargo test -p omni-cua-s1-native +CUA_S1_CUDA_LIB=$PWD/target/release/libqwen3_5_cuda.so \ + cargo test --release -p omni-cua-s1-native --test kernels -- --ignored +``` + +Accuracy against the float32 reference worker, under the tolerance in [`src/models/cua_s1/README.md`](../../src/models/cua_s1/README.md#validation). First write the reference results with the reference worker's check (`text.md`), once in each dtype: + +```sh +mkdir -p parity +.venv/bin/python recipe/cua_s1/compare_text_with_upstream.py --upstream ../cua \ + --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text --device cuda \ + --dtype float32 --no-tf32 --out parity/parity_float32.jsonl +# the same with --dtype bfloat16 --out parity/parity_bfloat16.jsonl +``` + +Then score the fixed input set eagerly and as served, in two separate processes, and compare: + +```sh +for run in 1 2; do + target/release/omni-cua-s1-native --model weights/cua-s1-4b-0.2-text-merged \ + --gemm-plans weights/gemm-plans.json --score-all tests/cua_s1/data/text_inputs.json > scores-$run.jsonl +done +python3 recipe/cua_s1/check_native.py scores-1.jsonl parity scores-2.jsonl +``` + +It passes when the largest difference and the top options are within the tolerance, the served results (CUDA graphs) equal the eager ones bit for bit, and the two processes agree bit for bit. + +The HTTP behaviour against the reference worker: start the reference worker on port 8001 and the native worker on port 8002, then + +```sh +python3 recipe/cua_s1/diff_corpus.py tests/cua_s1/data/text_inputs.json corpus.jsonl 150 +python3 recipe/cua_s1/diff_workers.py corpus.jsonl 8001 8002 diff.jsonl +``` + +sends the fixed input set, edge cases for every error the workers return, and random bodies to both. Errors must be identical (status, content type and body); answers must have the same model, usage, keys and types, and the same choice wherever the reference's top-two margin is at least 0.05. + +Latency, through the frontend and directly, as for the reference worker: + +```sh +.venv/bin/python recipe/cua_s1/bench_text.py --direct http://127.0.0.1:8000 \ + --frontend http://127.0.0.1:8080 --warmup 3 --repeat 20 +``` + +`--bench ... --bench-gap-ms 50` times the forward pass alone, with idle time between passes. + +## Results + +On one RTX 6000 Ada (48 GB, sm_89), CUDA 13.2, driver 595.91.07, with the pinned revisions: + +- Accuracy: over the 16 questions of the fixed input set, the largest difference from the float32 reference worker is 0.0126 (the allowance is 2 × 0.0145 + 0.01 = 0.039), and no top option changes. Served and eager results are bitwise identical, and so are two separate processes with the same plans file. +- HTTP: 567 requests to both workers (the fixed input set, error cases and random bodies) give identical status and body for every error and the same keys, types and choices for every answer; the largest probability difference is 0.029. +- Frontend: all 14 requests return identical status, content type and body bytes directly and through the frontend. +- Startup: 2.0 s to a ready `/health` with a plans file (about 70 s on the first start with `--gemm-search`); the first request after that took 16 ms. The card peaked at 12.6 GiB in use while serving. +- Latency, p50 in milliseconds, one request at a time. `bench_text.py` sends requests back to back, which keeps this card at its 300 W power limit; the forward-only numbers leave 50 ms idle before each pass, closer to decisions that arrive one by one. The reference worker's numbers come from the same two methods. + +| Case | Prompt tokens | Reference worker, `bench_text.py` | Native, `bench_text.py` | Reference worker, forward only | Native, forward only | +| --- | --- | --- | --- | --- | --- | +| `one_option` | 139 | 46.2 | 15.7 | 44.6 | 12.2 | +| `fixture_positive` | 218 | 49.2 | 18.8 | 47.4 | 13.8 | +| `fixture_negative` | 292 | 52.7 | 25.6 | 51.1 | 17.0 | +| `max_26_options` | 712 | 109.3 | 50.3 | 103.0 | 33.4 | +| `long_state` | 15446 | 3084.7 | 1364.5 | 3047.6 | 1308.9 | + +## Not covered + +- `score` and `noul` questions, and the `multimodal` adapter, as in the reference worker. +- More than one request at a time: the worker answers one decision at a time, like the reference worker. +- GPUs other than sm_89, and CUDA versions other than 13.2, have not been run. The GEMM plans are per GPU and cuBLASLt version. diff --git a/src/backends/cuda/README.md b/src/backends/cuda/README.md index 1ba4c67c..577209d6 100644 --- a/src/backends/cuda/README.md +++ b/src/backends/cuda/README.md @@ -4,4 +4,4 @@ Planned home for high-performance NVIDIA GPU operations and kernel integration. Model orchestration, batching policy, state management, and kernel selection remain with the model engine. CUDA and Metal implementations do not need identical internal structures or a universal tensor abstraction. -Status: planned; no CUDA implementation or validated hardware coverage yet. +Status: [`qwen3_5/`](qwen3_5/) has the operations of a prefill-only Qwen3.5 forward pass, used by the Cua-S1 native worker and measured on sm_89. Other models are planned. diff --git a/src/backends/cuda/qwen3_5/README.md b/src/backends/cuda/qwen3_5/README.md new file mode 100644 index 00000000..8863ab7e --- /dev/null +++ b/src/backends/cuda/qwen3_5/README.md @@ -0,0 +1,21 @@ +# Qwen3.5 prefill operations + +CUDA kernels for a prefill-only Qwen3.5 forward pass, built into `libqwen3_5_cuda.so`: + +```sh +src/backends/cuda/qwen3_5/build.sh [compute capability, default 89] +``` + +The library has a C interface ([`ops.h`](ops.h)): the operations, plus the few CUDA runtime calls a caller needs (allocation, copies, streams, graph capture), so that a Rust model engine can load it at run time and build without a CUDA toolkit. The Cua-S1 native worker ([`src/models/cua_s1/native/`](../../../models/cua_s1/native/)) uses it; the layer loop, buffers, CUDA graphs and GEMM algorithm choice stay in that model engine. + +| File | Operations | +| --- | --- | +| `norm.cu` | Zero-centred RMSNorm, the residual add fused with the next norm, and the gated RMSNorm of the Gated DeltaNet output. | +| `elementwise.cu` | Embedding lookup; the Gated DeltaNet causal convolution with SiLU and its gates; the attention output gate; SiLU(gate) * up. | +| `attention.cu` | q/k RMSNorm and partial rotary embedding; causal attention with grouped KV heads (head dim 256) on tensor cores, FlashAttention-2 style. | +| `gdn_prefill.cu` | The chunked gated delta rule (chunks of 64) in three kernels, split the way flash-linear-attention splits it: per-chunk preparation, the state carried from chunk to chunk, and the per-chunk output. | +| `gemm.cu` | bfloat16 GEMMs through cuBLASLt with float32 accumulation, algorithm tuning, and saving and loading the tuned choices. | +| `runtime.cu` | The CUDA runtime calls. | +| `mma.cuh`, `common.cuh` | `mma.sync`, `ldmatrix` and `cp.async` helpers, and shared device helpers. | + +The norm, elementwise and q/k preparation kernels round to bfloat16 at the same points as the Transformers implementation (`modeling_qwen3_5.py`). Attention and the Gated DeltaNet prefill keep some intermediate results in bfloat16, as FlashAttention and flash-linear-attention do, where Transformers' float32 fallback for the gated delta rule does not; the model is checked end to end against the float32 reference worker. Tensor-core kernels need sm_80 or newer; they are measured on sm_89 (RTX 6000 Ada). Split-K reductions that accumulate into the output in place are not used, and plans that use them are refused on import, so a given GEMM algorithm always gives the same result. `build.sh` also embeds PTX, but only sm_89 has been run. diff --git a/src/backends/cuda/qwen3_5/attention.cu b/src/backends/cuda/qwen3_5/attention.cu new file mode 100644 index 00000000..75489deb --- /dev/null +++ b/src/backends/cuda/qwen3_5/attention.cu @@ -0,0 +1,353 @@ +// Full-attention layers of Qwen3.5: q/k preparation and causal attention. +// +// cs1_attention is a FlashAttention-2 style kernel on tensor cores (mma.sync +// m16n8k16, bfloat16 in, float32 accumulation): a block takes 64 queries of one head, +// four warps of 16 rows each, and walks the keys up to its last query in tiles of +// 32, keeping the output and the online softmax in registers. The probabilities are +// rounded to bfloat16 for the P*V product, as in flash attention; the running sums +// stay float32. cs1_attention_simple is a plain float32 version kept for checking. +#include "common.cuh" +#include "mma.cuh" +#include "ops.h" + +namespace cs1 { +namespace { + +constexpr int DH = 256; // head dim +constexpr int PER = DH / 32; // values per lane +constexpr int QB = 16; // queries per block, two per warp +constexpr int KB = 32; // keys per shared-memory tile +constexpr int ATTN_THREADS = 256; + +// One warp per (token, head), q heads first, then k heads. Each lane holds 8 +// consecutive dims, so the rotary partner of dim i < 32 (dim i + 32) sits in lane ^ 4. +__global__ void attn_prep_kernel(const bf16* __restrict__ qg, const bf16* __restrict__ kr, int ld, + const bf16* __restrict__ qw, const bf16* __restrict__ kw, + const bf16* __restrict__ cos, const bf16* __restrict__ sin, + bf16* __restrict__ q, bf16* __restrict__ gate, bf16* __restrict__ k, int T, int Hq, + int Hk, int half, float eps) { + const int warp = blockIdx.x * (blockDim.x / 32) + threadIdx.x / 32, lane = threadIdx.x & 31; + const int heads = Hq + Hk; + if (warp >= T * heads) return; + const int t = warp / heads, hh = warp % heads; + const bool is_q = hh < Hq; + const int h = is_q ? hh : hh - Hq; + const bf16* src = is_q ? qg + (size_t)t * ld + (size_t)h * 2 * DH : kr + (size_t)t * ld + (size_t)h * DH; + const bf16* w = is_q ? qw : kw; + const int d0 = lane * PER; + + float x[PER]; + load8(src + d0, x); + float ss = 0.f; +#pragma unroll + for (int i = 0; i < PER; i++) ss += x[i] * x[i]; + const float inv = rsqrtf(warp_sum(ss) / DH + eps); + float wv[PER]; + load8(w + d0, wv); +#pragma unroll + for (int i = 0; i < PER; i++) x[i] = round_bf16(x[i] * inv * (1.f + wv[i])); + + // rotate_half on the first 2 * half dims: out = x * cos + rotate_half(x) * sin, + // each product and the sum rounded to bfloat16 as in the reference. + const int rot = 2 * half; + float y[PER]; +#pragma unroll + for (int i = 0; i < PER; i++) { + const float partner = __shfl_xor_sync(0xffffffffu, x[i], (half / PER)); + const int d = d0 + i; + y[i] = x[i]; + if (d < rot) { + const int fi = d % half; + const float c = f32(cos[(size_t)t * half + fi]), s = f32(sin[(size_t)t * half + fi]); + const float r = d < half ? -partner : partner; + y[i] = round_bf16(round_bf16(x[i] * c) + round_bf16(r * s)); + } + } + if (is_q) { + store8(q + ((size_t)t * Hq + h) * DH + d0, y); + float gv[PER]; + load8(src + DH + d0, gv); + store8(gate + ((size_t)t * Hq + h) * DH + d0, gv); + } else { + store8(k + ((size_t)t * Hk + h) * DH + d0, y); + } +} + +// Causal attention, float32 scores and online softmax. A block takes QB queries of +// one head and walks the keys up to its last query in shared-memory tiles. +__global__ void __launch_bounds__(ATTN_THREADS) + attention_kernel(const bf16* __restrict__ q, const bf16* __restrict__ k, const bf16* __restrict__ v, int ldv, + bf16* __restrict__ out, int T, int Hq, int Hk, float scale) { + __shared__ __align__(16) bf16 ks[KB * DH]; + __shared__ __align__(16) bf16 vs[KB * DH]; + const int h = blockIdx.y, hk = h / (Hq / Hk); + const int warp = threadIdx.x / 32, lane = threadIdx.x & 31, d0 = lane * PER; + const int first = blockIdx.x * QB + warp * 2; + + float qv[2][PER], acc[2][PER], m[2], l[2]; +#pragma unroll + for (int r = 0; r < 2; r++) { + const int t = first + r; + if (t < T) { + load8(q + ((size_t)t * Hq + h) * DH + d0, qv[r]); + } else { +#pragma unroll + for (int i = 0; i < PER; i++) qv[r][i] = 0.f; + } +#pragma unroll + for (int i = 0; i < PER; i++) acc[r][i] = 0.f; + m[r] = -INFINITY; + l[r] = 0.f; + } + + const int kv_end = min(T, (int)(blockIdx.x * QB + QB)); + for (int k0 = 0; k0 < kv_end; k0 += KB) { + __syncthreads(); + for (int x = threadIdx.x; x < KB * DH / 8; x += blockDim.x) { + const int j = x / (DH / 8), c = (x % (DH / 8)) * 8, s = k0 + j; + Pack8 kk{}, vv{}; + if (s < T) { + kk = *reinterpret_cast(k + ((size_t)s * Hk + hk) * DH + c); + vv = *reinterpret_cast(v + (size_t)s * ldv + (size_t)hk * DH + c); + } + *reinterpret_cast(ks + j * DH + c) = kk; + *reinterpret_cast(vs + j * DH + c) = vv; + } + __syncthreads(); + const int jn = min(KB, kv_end - k0); + for (int j = 0; j < jn; j++) { + float kv[PER]; + load8(ks + j * DH + d0, kv); + float dot[2] = {0.f, 0.f}; +#pragma unroll + for (int i = 0; i < PER; i++) { + dot[0] = fmaf(qv[0][i], kv[i], dot[0]); + dot[1] = fmaf(qv[1][i], kv[i], dot[1]); + } + dot[0] = warp_sum(dot[0]); + dot[1] = warp_sum(dot[1]); + float vx[PER]; + load8(vs + j * DH + d0, vx); + const int s = k0 + j; +#pragma unroll + for (int r = 0; r < 2; r++) { + if (s > first + r) continue; + const float score = dot[r] * scale; + const float mn = fmaxf(m[r], score); + const float corr = expf(m[r] - mn), p = expf(score - mn); + l[r] = l[r] * corr + p; +#pragma unroll + for (int i = 0; i < PER; i++) acc[r][i] = fmaf(p, vx[i], acc[r][i] * corr); + m[r] = mn; + } + } + } +#pragma unroll + for (int r = 0; r < 2; r++) { + const int t = first + r; + if (t >= T) continue; + float o[PER]; +#pragma unroll + for (int i = 0; i < PER; i++) o[i] = acc[r][i] / l[r]; + store8(out + ((size_t)t * Hq + h) * DH + d0, o); + } +} + +// ---- flash attention ---- + +namespace flash { + +constexpr int D = 256, BM = 64, BN = 32, THREADS = 128; +constexpr int LDS = D + 8; // shared row stride in elements: 528 bytes keeps ldmatrix conflict-free +constexpr int SMEM_BYTES = (BM + 2 * BN) * LDS * 2; + +__global__ void __launch_bounds__(THREADS) + flash_kernel(const bf16* __restrict__ q, const bf16* __restrict__ k, const bf16* __restrict__ v, int ldv, + bf16* __restrict__ out, int T, int Hq, int Hk, float scale_log2) { + extern __shared__ __align__(16) unsigned char smem[]; + bf16* qs = reinterpret_cast(smem); + bf16* ks = qs + BM * LDS; + bf16* vs = ks + BN * LDS; + const int h = blockIdx.y, hk = h / (Hq / Hk); + const int q0 = (gridDim.x - 1 - blockIdx.x) * BM; // the longest blocks first + const int tid = threadIdx.x, warp = tid / 32, lane = tid % 32; + const int g = lane / 4, t = lane % 4; + const int row0 = q0 + warp * 16; // this warp's first query + + for (int c = tid; c < BM * (D / 8); c += THREADS) { + const int r = c / (D / 8), col = (c % (D / 8)) * 8, row = q0 + r; + cp_async16(qs + r * LDS + col, q + ((size_t)min(row, T - 1) * Hq + h) * D + col, row < T); + } + cp_async_commit(); + + float o[D / 8][4]; +#pragma unroll + for (int n = 0; n < D / 8; n++) o[n][0] = o[n][1] = o[n][2] = o[n][3] = 0.f; + float m[2] = {-INFINITY, -INFINITY}, l[2] = {0.f, 0.f}; + + const int kv_end = min(T, q0 + BM); + for (int k0 = 0; k0 < kv_end; k0 += BN) { + for (int c = tid; c < BN * (D / 8); c += THREADS) { + const int r = c / (D / 8), col = (c % (D / 8)) * 8, s = k0 + r; + cp_async16(ks + r * LDS + col, k + ((size_t)min(s, T - 1) * Hk + hk) * D + col, s < T); + } + cp_async_commit(); + for (int c = tid; c < BN * (D / 8); c += THREADS) { + const int r = c / (D / 8), col = (c % (D / 8)) * 8, s = k0 + r; + cp_async16(vs + r * LDS + col, v + (size_t)min(s, T - 1) * ldv + (size_t)hk * D + col, s < T); + } + cp_async_commit(); + cp_async_wait<1>(); // Q and K + __syncthreads(); + + // keys past every query of this warp contribute nothing + const bool active = k0 <= row0 + 15; + float sc[BN / 8][4]; +#pragma unroll + for (int n = 0; n < BN / 8; n++) sc[n][0] = sc[n][1] = sc[n][2] = sc[n][3] = 0.f; + if (active) { +#pragma unroll + for (int kk = 0; kk < D; kk += 16) { + uint32_t a[4]; + load_a(a, qs, LDS, warp * 16, kk, lane); +#pragma unroll + for (int n = 0; n < BN / 8; n += 2) { + uint32_t b[4]; + load_b_nk(b, ks, LDS, kk, n * 8, lane); + mma16816(sc[n], a, b[0], b[1]); + mma16816(sc[n + 1], a, b[2], b[3]); + } + } + } + uint32_t p[BN / 16][4]; + if (active) { + // causal and length mask, then the online softmax in base 2 + float mx[2] = {-INFINITY, -INFINITY}; +#pragma unroll + for (int n = 0; n < BN / 8; n++) { +#pragma unroll + for (int e = 0; e < 4; e++) { + const int key = k0 + n * 8 + 2 * t + (e & 1), row = row0 + g + (e >> 1) * 8; + sc[n][e] = (key <= row && key < T) ? sc[n][e] * scale_log2 : -INFINITY; + mx[e >> 1] = fmaxf(mx[e >> 1], sc[n][e]); + } + } + float alpha[2], base[2]; +#pragma unroll + for (int r = 0; r < 2; r++) { + mx[r] = fmaxf(mx[r], __shfl_xor_sync(0xffffffffu, mx[r], 1)); + mx[r] = fmaxf(mx[r], __shfl_xor_sync(0xffffffffu, mx[r], 2)); + const float mn = fmaxf(m[r], mx[r]); + base[r] = mn == -INFINITY ? 0.f : mn; + alpha[r] = exp2f(m[r] - base[r]); + m[r] = mn; + l[r] *= alpha[r]; + } +#pragma unroll + for (int n = 0; n < BN / 8; n++) { +#pragma unroll + for (int e = 0; e < 4; e++) { + sc[n][e] = exp2f(sc[n][e] - base[e >> 1]); + l[e >> 1] += sc[n][e]; + } + } +#pragma unroll + for (int n = 0; n < D / 8; n++) { + o[n][0] *= alpha[0]; + o[n][1] *= alpha[0]; + o[n][2] *= alpha[1]; + o[n][3] *= alpha[1]; + } + // the score accumulators, two 8-key tiles at a time, are the A fragments of P*V +#pragma unroll + for (int j = 0; j < BN / 16; j++) { + p[j][0] = pack_bf16(sc[2 * j][0], sc[2 * j][1]); + p[j][1] = pack_bf16(sc[2 * j][2], sc[2 * j][3]); + p[j][2] = pack_bf16(sc[2 * j + 1][0], sc[2 * j + 1][1]); + p[j][3] = pack_bf16(sc[2 * j + 1][2], sc[2 * j + 1][3]); + } + } + cp_async_wait<0>(); // V + __syncthreads(); + if (active) { +#pragma unroll + for (int j = 0; j < BN / 16; j++) { +#pragma unroll + for (int n = 0; n < D / 8; n += 2) { + uint32_t b[4]; + load_b_kn(b, vs, LDS, j * 16, n * 8, lane); + mma16816(o[n], p[j], b[0], b[1]); + mma16816(o[n + 1], p[j], b[2], b[3]); + } + } + } + __syncthreads(); // before the next tile overwrites K and V + } + + // the four lanes of a row each summed a quarter of its keys +#pragma unroll + for (int r = 0; r < 2; r++) { + l[r] += __shfl_xor_sync(0xffffffffu, l[r], 1); + l[r] += __shfl_xor_sync(0xffffffffu, l[r], 2); + } + const float inv[2] = {1.f / l[0], 1.f / l[1]}; +#pragma unroll + for (int r = 0; r < 2; r++) { + const int row = row0 + g + r * 8; + if (row >= T) continue; + bf16* dst = out + ((size_t)row * Hq + h) * D + 2 * t; +#pragma unroll + for (int n = 0; n < D / 8; n++) + *reinterpret_cast(dst + n * 8) = pack_bf16(o[n][2 * r] * inv[r], o[n][2 * r + 1] * inv[r]); + } +} + +} // namespace flash + +} // namespace +} // namespace cs1 + +using namespace cs1; + +extern "C" int cs1_attn_prep(const void* qg, const void* kr, int ld, const void* qw, const void* kw, + const void* cos, const void* sin, void* q, void* gate, void* k, int T, int Hq, int Hk, + int Dh, int half, float eps, void* stream) { + if (ld % 8 != 0 || T < 0) return cudaErrorInvalidValue; + // the lane ^ (half / 8) partner exchange needs 2 * half <= 256 and half a multiple of 8 + if (Dh != DH || half % PER != 0 || 2 * half > DH || (half / PER) & ((half / PER) - 1)) return cudaErrorInvalidValue; + const int warps = T * (Hq + Hk); + if (warps == 0) return cudaSuccess; + constexpr int WARPS = 8; + attn_prep_kernel<<<(warps + WARPS - 1) / WARPS, WARPS * 32, 0, static_cast(stream)>>>( + static_cast(qg), static_cast(kr), ld, static_cast(qw), + static_cast(kw), static_cast(cos), static_cast(sin), + static_cast(q), static_cast(gate), static_cast(k), T, Hq, Hk, half, eps); + return cudaGetLastError(); +} + +extern "C" int cs1_attention_simple(const void* q, const void* k, const void* v, int ldv, void* out, int T, int Hq, + int Hk, int Dh, float scale, void* stream) { + if (Dh != DH || Hk <= 0 || Hq % Hk != 0 || ldv % 8 != 0 || ldv < Hk * Dh || T < 0) return cudaErrorInvalidValue; + if (T == 0) return cudaSuccess; + attention_kernel<<(stream)>>>( + static_cast(q), static_cast(k), static_cast(v), ldv, + static_cast(out), T, Hq, Hk, scale); + return cudaGetLastError(); +} + +extern "C" int cs1_attention(const void* q, const void* k, const void* v, int ldv, void* out, int T, int Hq, int Hk, + int Dh, float scale, void* stream) { + if (Dh != flash::D || Hk <= 0 || Hq % Hk != 0 || ldv % 8 != 0 || ldv < Hk * Dh || T < 0) + return cudaErrorInvalidValue; + if (T == 0) return cudaSuccess; + // once per process (for the device current at the first call) + static const cudaError_t configured = cudaFuncSetAttribute( + flash::flash_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, flash::SMEM_BYTES); + if (configured != cudaSuccess) return configured; + constexpr float LOG2E = 1.4426950408889634f; + flash::flash_kernel<<(stream)>>>( + static_cast(q), static_cast(k), static_cast(v), ldv, + static_cast(out), T, Hq, Hk, scale * LOG2E); + return cudaGetLastError(); +} diff --git a/src/backends/cuda/qwen3_5/build.sh b/src/backends/cuda/qwen3_5/build.sh new file mode 100755 index 00000000..ef2dad6b --- /dev/null +++ b/src/backends/cuda/qwen3_5/build.sh @@ -0,0 +1,27 @@ +#!/usr/bin/env bash +# Build libqwen3_5_cuda.so from the kernels in this directory. +# +# src/backends/cuda/qwen3_5/build.sh [compute capability, e.g. 89] +# +# Needs nvcc (from NVCC, CUDA_HOME/bin or PATH) and cuBLASLt; the compute capability +# defaults to CUDA_COMPUTE_CAP, else 89. The library holds machine code for that +# compute capability and PTX that newer GPUs can compile at load time. The CUDA +# runtime is linked statically and cuBLASLt dynamically, with an rpath to the +# toolkit's library directory when there is one next to nvcc. +set -euo pipefail +here=$(cd "$(dirname "$0")" && pwd) +out=${1:?usage: build.sh [compute capability]} +arch=${2:-${CUDA_COMPUTE_CAP:-89}} +nvcc=${NVCC:-} +if [ -z "$nvcc" ]; then + if [ -n "${CUDA_HOME:-}" ]; then nvcc=$CUDA_HOME/bin/nvcc; else nvcc=$(command -v nvcc); fi +fi +link=(-lcublasLt) +if lib=$(cd "$(dirname "$nvcc")/../lib64" 2>/dev/null && pwd); then + link=(-L"$lib" -lcublasLt -Xlinker -rpath -Xlinker "$lib") +fi +mkdir -p "$out" +"$nvcc" -O3 -std=c++17 -gencode "arch=compute_${arch},code=[sm_${arch},compute_${arch}]" \ + -shared -Xcompiler -fPIC -Xcompiler -Wall,-Wextra -I"$here" "$here"/*.cu \ + "${link[@]}" -o "$out/libqwen3_5_cuda.so" +echo "built $out/libqwen3_5_cuda.so for sm_${arch}" diff --git a/src/backends/cuda/qwen3_5/common.cuh b/src/backends/cuda/qwen3_5/common.cuh new file mode 100644 index 00000000..d70c1761 --- /dev/null +++ b/src/backends/cuda/qwen3_5/common.cuh @@ -0,0 +1,62 @@ +// Helpers shared by the Qwen3.5 kernels. +#pragma once + +#include +#include +#include + +namespace cs1 { + +using bf16 = __nv_bfloat16; + +__device__ __forceinline__ float f32(bf16 x) { return __bfloat162float(x); } +__device__ __forceinline__ bf16 to_bf16(float x) { return __float2bfloat16(x); } +// A float rounded through bfloat16, as PyTorch stores the result of each bfloat16 op. +__device__ __forceinline__ float round_bf16(float x) { return __bfloat162float(__float2bfloat16(x)); } + +__device__ __forceinline__ float warp_sum(float x) { +#pragma unroll + for (int o = 16; o > 0; o >>= 1) x += __shfl_xor_sync(0xffffffffu, x, o); + return x; +} + +// Sum over the block; `scratch` holds at least 32 floats. Every thread gets the sum. +__device__ __forceinline__ float block_sum(float x, float* scratch) { + const int lane = threadIdx.x & 31, warp = threadIdx.x >> 5, warps = (blockDim.x + 31) >> 5; + x = warp_sum(x); + if (lane == 0) scratch[warp] = x; + __syncthreads(); + if (warp == 0) { + float t = lane < warps ? scratch[lane] : 0.f; + t = warp_sum(t); + if (lane == 0) scratch[0] = t; + } + __syncthreads(); + const float total = scratch[0]; + __syncthreads(); + return total; +} + +// PyTorch's float32 SiLU and sigmoid on CUDA. +__device__ __forceinline__ float silu(float x) { return x / (1.f + expf(-x)); } +__device__ __forceinline__ float sigmoid(float x) { return 1.f / (1.f + expf(-x)); } + +// 8 bfloat16 values, one 16-byte load or store. +struct alignas(16) Pack8 { + bf16 v[8]; +}; + +__device__ __forceinline__ void load8(const bf16* p, float out[8]) { + const Pack8 pk = *reinterpret_cast(p); +#pragma unroll + for (int i = 0; i < 8; i++) out[i] = f32(pk.v[i]); +} + +__device__ __forceinline__ void store8(bf16* p, const float in[8]) { + Pack8 pk; +#pragma unroll + for (int i = 0; i < 8; i++) pk.v[i] = to_bf16(in[i]); + *reinterpret_cast(p) = pk; +} + +} // namespace cs1 diff --git a/src/backends/cuda/qwen3_5/elementwise.cu b/src/backends/cuda/qwen3_5/elementwise.cu new file mode 100644 index 00000000..3b1c01bd --- /dev/null +++ b/src/backends/cuda/qwen3_5/elementwise.cu @@ -0,0 +1,126 @@ +// Embedding lookup, the Gated DeltaNet conv and gates, and the elementwise +// activations of Qwen3.5. +#include "common.cuh" +#include "ops.h" + +namespace cs1 { +namespace { + +constexpr int THREADS = 256; + +__global__ void embed_kernel(const int32_t* __restrict__ ids, const Pack8* __restrict__ table, + Pack8* __restrict__ out, int packs) { + const size_t t = blockIdx.x; + const size_t id = ids[t]; + for (int i = threadIdx.x; i < packs; i += blockDim.x) out[t * packs + i] = table[id * packs + i]; +} + +// F.conv1d in bfloat16 (float32 accumulation, rounded), then SiLU (rounded again), +// written to three contiguous outputs. +__global__ void gdn_conv_kernel(const bf16* __restrict__ qkv, int ld, const bf16* __restrict__ w, + bf16* __restrict__ q, bf16* __restrict__ k, bf16* __restrict__ v, int T, + int key_dim, int value_dim) { + const int channels = 2 * key_dim + value_dim; + const size_t idx = (size_t)blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= (size_t)T * channels) return; + const int t = idx / channels, c = idx % channels; + float acc = 0.f; +#pragma unroll + for (int j = 0; j < 4; j++) { + const int s = t - 3 + j; + if (s >= 0) acc = fmaf(f32(w[c * 4 + j]), f32(qkv[(size_t)s * ld + c]), acc); + } + const bf16 y = to_bf16(silu(round_bf16(acc))); + if (c < key_dim) + q[(size_t)t * key_dim + c] = y; + else if (c < 2 * key_dim) + k[(size_t)t * key_dim + c - key_dim] = y; + else + v[(size_t)t * value_dim + c - 2 * key_dim] = y; +} + +// beta = sigmoid(b) in bfloat16; g = -exp(A_log) * softplus(a + dt_bias) in float32 +// (F.softplus with threshold 20). +__global__ void gdn_gates_kernel(const bf16* __restrict__ b, const bf16* __restrict__ a, int ld, + const bf16* __restrict__ A_log, const bf16* __restrict__ dt_bias, + bf16* __restrict__ beta, float* __restrict__ g, int n, int H) { + const int i = blockIdx.x * blockDim.x + threadIdx.x; + if (i >= n) return; + const int h = i % H; + const size_t src = (size_t)(i / H) * ld + h; + beta[i] = to_bf16(sigmoid(f32(b[src]))); + const float x = f32(a[src]) + f32(dt_bias[h]); + const float sp = x > 20.f ? x : log1pf(expf(x)); + g[i] = -expf(f32(A_log[h])) * sp; +} + +// attn_output * torch.sigmoid(gate), both bfloat16. +__global__ void sigmoid_gate_kernel(bf16* __restrict__ x, const bf16* __restrict__ gate, size_t n) { + const size_t i = (size_t)blockIdx.x * blockDim.x + threadIdx.x; + if (i >= n) return; + x[i] = to_bf16(f32(x[i]) * round_bf16(sigmoid(f32(gate[i])))); +} + +// act_fn(gate_proj(x)) * up_proj(x), both bfloat16; gate and up are the two halves of +// each row of gate_up. +__global__ void silu_mul_kernel(const bf16* __restrict__ gate_up, int ld, bf16* __restrict__ out, int I, + size_t n) { + const size_t i = (size_t)blockIdx.x * blockDim.x + threadIdx.x; + if (i >= n) return; + const size_t t = i / I, j = i % I; + const bf16* row = gate_up + t * ld; + out[i] = to_bf16(round_bf16(silu(f32(row[j]))) * f32(row[I + j])); +} + +unsigned blocks(size_t n) { return (unsigned)((n + THREADS - 1) / THREADS); } + +} // namespace +} // namespace cs1 + +using namespace cs1; + +extern "C" int cs1_embed(const int32_t* ids, const void* table, void* out, int T, int D, void* stream) { + if (D % 8 != 0) return cudaErrorInvalidValue; + if (T <= 0) return cudaSuccess; + embed_kernel<<(stream)>>>( + ids, static_cast(table), static_cast(out), D / 8); + return cudaGetLastError(); +} + +extern "C" int cs1_gdn_conv(const void* qkv, int ld, const void* w, void* q, void* k, void* v, int T, int key_dim, + int value_dim, void* stream) { + if (T < 0 || key_dim < 0 || value_dim < 0 || ld < 2 * key_dim + value_dim) return cudaErrorInvalidValue; + const size_t n = (size_t)T * (2 * key_dim + value_dim); + if (n == 0) return cudaSuccess; + gdn_conv_kernel<<(stream)>>>( + static_cast(qkv), ld, static_cast(w), static_cast(q), + static_cast(k), static_cast(v), T, key_dim, value_dim); + return cudaGetLastError(); +} + +extern "C" int cs1_gdn_gates(const void* b, const void* a, int ld, const void* A_log, const void* dt_bias, + void* beta, float* g, int T, int H, void* stream) { + if (T < 0 || H < 0 || ld < H) return cudaErrorInvalidValue; + const int n = T * H; + if (n == 0) return cudaSuccess; + gdn_gates_kernel<<(stream)>>>( + static_cast(b), static_cast(a), ld, static_cast(A_log), + static_cast(dt_bias), static_cast(beta), g, n, H); + return cudaGetLastError(); +} + +extern "C" int cs1_sigmoid_gate(void* x, const void* gate, size_t n, void* stream) { + if (n == 0) return cudaSuccess; + sigmoid_gate_kernel<<(stream)>>>( + static_cast(x), static_cast(gate), n); + return cudaGetLastError(); +} + +extern "C" int cs1_silu_mul(const void* gate_up, int ld, void* out, int T, int I, void* stream) { + if (T < 0 || I < 0 || ld < 2 * I) return cudaErrorInvalidValue; + const size_t n = (size_t)T * I; + if (n == 0) return cudaSuccess; + silu_mul_kernel<<(stream)>>>( + static_cast(gate_up), ld, static_cast(out), I, n); + return cudaGetLastError(); +} diff --git a/src/backends/cuda/qwen3_5/gdn_prefill.cu b/src/backends/cuda/qwen3_5/gdn_prefill.cu new file mode 100644 index 00000000..516e51b9 --- /dev/null +++ b/src/backends/cuda/qwen3_5/gdn_prefill.cu @@ -0,0 +1,528 @@ +// Chunked Gated DeltaNet prefill (forward only, batch 1) on tensor cores, sm_80 and later. +// +// The math of torch_chunk_gated_delta_rule in Transformers (chunks of 64), split into +// three kernels the way flash-linear-attention splits its chunked forward pass: +// 1. gdn_chunk_prep, per (chunk, head): L2 norms, the cumulative decays, the pair +// products k.k and q.k (TF32 with float32 accumulation), the triangular inverse +// T = (I + A)^-1 (float32, CUDA cores), and u = T (beta v), w = T (beta exp(cum) k) +// (bfloat16 with float32 accumulation). Results are stored as bfloat16. +// 2. gdn_chunk_state, per (head, 32 value columns), over the chunks in order: keeps +// the state S in float32 registers, stores it as bfloat16 before each chunk, and +// computes v_new = u - w S and S = decay S + kd^T v_new with mma.sync. +// 3. gdn_chunk_out, per (chunk, head): o = qd S + P v_new with mma.sync. +// Transformers computes all of this in float32. Keeping the intermediate results in +// bfloat16, as flash-linear-attention does, makes this kernel less precise than that +// path; tests/kernels.rs checks it against a float64 token-by-token reference. +#include +#include +#include +#include + +#include "mma.cuh" +#include "ops.h" + +using namespace nvcuda; + +namespace { + +using cs1::bf16; +using cs1::warp_sum; + +constexpr int C = 64; // chunk length +constexpr int K = 128; // key head dim +constexpr int V = 128; // value head dim +constexpr int THREADS = 256; +// row strides: multiples of 16 bytes as WMMA needs, and not multiples of 32 floats +constexpr int KP = K + 4; // float +constexpr int CP = C + 4; // float +constexpr int HB = K + 8; // bfloat16 +constexpr int TB = C + 8; // bfloat16 + +using ATf32Row = wmma::fragment; +using BTf32Col = wmma::fragment; +using CTf32 = wmma::fragment; +using ABf16 = wmma::fragment; +using BBf16 = wmma::fragment; +using CBf16 = wmma::fragment; + +template +__device__ __forceinline__ void to_tf32(F& f) { +#pragma unroll + for (int t = 0; t < f.num_elements; t++) f.x[t] = wmma::__float_to_tf32(f.x[t]); +} + +struct Work { + bf16* u; // [H, NCC, V] + bf16* w; // [H, NCC, K] + bf16* qd; // [H, NCC, K] q * scale * exp(cum) + bf16* kd; // [H, NCC, K] k * exp(cum_last - cum) + float* p; // [H, NC, C, C] float32 scratch of q.k + bf16* pb; // [H, NC, C, C] (q.k) exp(cum_i - cum_j), j <= i + float* decay; // [H, NC] exp(cum_last) + bf16* s; // [H, NC, K, V] state before each chunk + bf16* vn; // [H, NCC, V] v_new +}; + +// Byte offsets of the workspace parts, 256-byte aligned. +struct Layout { + size_t u, w, qd, kd, p, pb, decay, s, vn, total; + Layout(int T, int H) { + const size_t NC = (T + C - 1) / C, NCC = NC * C; + size_t at = 0; + auto take = [&](size_t bytes) { + const size_t off = at; + at = (at + bytes + 255) / 256 * 256; + return off; + }; + u = take((size_t)H * NCC * V * 2); + w = take((size_t)H * NCC * K * 2); + qd = take((size_t)H * NCC * K * 2); + kd = take((size_t)H * NCC * K * 2); + p = take((size_t)H * NC * C * C * 4); + pb = take((size_t)H * NC * C * C * 2); + decay = take((size_t)H * NC * 4); + s = take((size_t)H * NC * K * V * 2); + vn = take((size_t)H * NCC * V * 2); + total = at; + } +}; + +// kernel 1 shared memory, bytes +constexpr int R1 = 0; // kn float [C][KP]; later v as bf16 [C][HB] +constexpr int R2 = R1 + C * KP * 4; // qn float [C][KP]; later T float [C][CP] + scratch, then T as bf16 +constexpr int R3 = R2 + C * KP * 4; // A float [C][CP]; later k exp(cum) as bf16 [C][HB] +constexpr int R4 = R3 + C * CP * 4; // cum, beta, exp(cum), exp(cum_last - cum) +constexpr int R5 = R4 + 4 * C * 4; // per-warp 16x16 float staging of u and w +constexpr size_t SMEM1_BYTES = R5 + (THREADS / 32) * 256 * 4; +constexpr int TBF = C * CP * 4 + 3 * 256 * 4; // offset of T as bf16 inside R2 +static_assert(C * HB * 2 <= C * CP * 4, "k exp(cum) as bf16 fits over A"); +static_assert(TBF + C * TB * 2 <= C * KP * 4, "T as bf16 fits in R2"); + +__global__ void __launch_bounds__(THREADS) gdn_chunk_prep( + const __nv_bfloat16* __restrict__ q, const __nv_bfloat16* __restrict__ k, + const __nv_bfloat16* __restrict__ v, const float* __restrict__ g, + const __nv_bfloat16* __restrict__ beta, Work ws, int T, int H, int HK, float scale) { + extern __shared__ __align__(128) unsigned char sm[]; + float* kn = reinterpret_cast(sm + R1); + float* qn = reinterpret_cast(sm + R2); + float* tm = qn; + float* sc = tm + C * CP; + float* am = reinterpret_cast(sm + R3); + float* cum = reinterpret_cast(sm + R4); + float* bet = cum + C; + float* ecum = bet + C; + float* erem = ecum + C; + __nv_bfloat16* vb = reinterpret_cast<__nv_bfloat16*>(sm + R1); + __nv_bfloat16* tb = reinterpret_cast<__nv_bfloat16*>(sm + R2 + TBF); + __nv_bfloat16* kw = reinterpret_cast<__nv_bfloat16*>(sm + R3); + + const int c = blockIdx.x, h = blockIdx.y, NC = gridDim.x, NCC = NC * C; + const int hk = h / (H / HK); + const int tid = threadIdx.x, lane = tid & 31, warp = tid >> 5; + + // 1. load q and k rows of this chunk, L2-normalize (padding rows stay zero) + for (int r = warp; r < C; r += THREADS / 32) { + const int t = c * C + r; + float kv[4], qv[4], ks = 0.f, qs = 0.f; +#pragma unroll + for (int e = 0; e < 4; e++) { + const int d = lane + 32 * e; + float kx = 0.f, qx = 0.f; + if (t < T) { + kx = __bfloat162float(k[((size_t)t * HK + hk) * K + d]); + qx = __bfloat162float(q[((size_t)t * HK + hk) * K + d]); + } + kv[e] = kx; + qv[e] = qx; + ks += kx * kx; + qs += qx * qx; + } + ks = warp_sum(ks); + qs = warp_sum(qs); + const float kinv = rsqrtf(ks + 1e-6f), qinv = rsqrtf(qs + 1e-6f) * scale; +#pragma unroll + for (int e = 0; e < 4; e++) { + const int d = lane + 32 * e; + kn[r * KP + d] = kv[e] * kinv; + qn[r * KP + d] = qv[e] * qinv; + } + } + if (tid < C) { + const int t = c * C + tid; + cum[tid] = t < T ? g[(size_t)t * H + h] : 0.f; + bet[tid] = t < T ? __bfloat162float(beta[(size_t)t * H + h]) : 0.f; + } + __syncthreads(); + if (tid == 0) { + float s = 0.f; + for (int i = 0; i < C; i++) { + s += cum[i]; + cum[i] = s; + } + ws.decay[h * NC + c] = expf(s); + } + __syncthreads(); + if (tid < C) { + ecum[tid] = expf(cum[tid]); + erem[tid] = expf(cum[C - 1] - cum[tid]); + } + __syncthreads(); + + // 2. decayed q and k for the state and output kernels + for (int x = tid; x < C * K; x += THREADS) { + const int i = x / K, d = x % K; + const size_t row = ((size_t)h * NCC + c * C + i) * K + d; + ws.qd[row] = __float2bfloat16(qn[i * KP + d] * ecum[i]); + ws.kd[row] = __float2bfloat16(kn[i * KP + d] * erem[i]); + } + + // 3. pair products on tensor cores: the 10 lower 16x16 tiles of k.k (into A) and of + // q.k (into P, in global memory), then the masks and decays elementwise + float* pout = ws.p + ((size_t)h * NC + c) * C * C; + for (int e = warp; e < 20; e += THREADS / 32) { + const int tile = e % 10; + const int it = tile < 1 ? 0 : tile < 3 ? 1 : tile < 6 ? 2 : 3; + const int jt = tile - it * (it + 1) / 2; + const float* a = (e < 10 ? kn : qn) + it * 16 * KP; + const float* b = kn + jt * 16 * KP; + CTf32 acc; + wmma::fill_fragment(acc, 0.f); +#pragma unroll 4 + for (int k0 = 0; k0 < K; k0 += 8) { + ATf32Row fa; + BTf32Col fb; + wmma::load_matrix_sync(fa, a + k0, KP); + wmma::load_matrix_sync(fb, b + k0, KP); + to_tf32(fa); + to_tf32(fb); + wmma::mma_sync(acc, fa, fb, acc); + } + if (e < 10) + wmma::store_matrix_sync(am + it * 16 * CP + jt * 16, acc, CP, wmma::mem_row_major); + else + wmma::store_matrix_sync(pout + it * 16 * C + jt * 16, acc, C, wmma::mem_row_major); + } + __syncthreads(); + __nv_bfloat16* pb = ws.pb + ((size_t)h * NC + c) * C * C; + for (int x = tid; x < C * C; x += THREADS) { + const int i = x / C, j = x % C; + const float dec = j <= i ? expf(cum[i] - cum[j]) : 0.f; + am[i * CP + j] = j < i ? bet[i] * am[i * CP + j] * dec : 0.f; + pb[i * C + j] = __float2bfloat16(j <= i ? pout[i * C + j] * dec : 0.f); + } + __syncthreads(); + + // 4. T = (I + A)^-1 in 16x16 blocks + for (int x = tid; x < C * CP; x += THREADS) tm[x] = 0.f; + __syncthreads(); + if (tid < C) { + const int base = 16 * (tid / 16), col = tid % 16; + float xs[16]; +#pragma unroll + for (int r = 0; r < 16; r++) { + float acc = r == col ? 1.f : 0.f; +#pragma unroll + for (int j = 0; j < r; j++) acc = fmaf(-am[(base + r) * CP + base + j], xs[j], acc); + xs[r] = acc; + } +#pragma unroll + for (int r = 0; r < 16; r++) tm[(base + r) * CP + base + col] = xs[r]; + } + __syncthreads(); + for (int lv = 1; lv < 4; lv++) { + const int nb = 4 - lv; + for (int x = tid; x < nb * 256; x += THREADS) { + const int bi = lv + x / 256, bj = bi - lv, r = (x % 256) / 16, cc = x % 16; + float acc = 0.f; + for (int kb = bj; kb < bi; kb++) +#pragma unroll + for (int m = 0; m < 16; m++) + acc = fmaf(am[(16 * bi + r) * CP + 16 * kb + m], tm[(16 * kb + m) * CP + 16 * bj + cc], acc); + sc[x] = acc; + } + __syncthreads(); + for (int x = tid; x < nb * 256; x += THREADS) { + const int bi = lv + x / 256, bj = bi - lv, r = (x % 256) / 16, cc = x % 16; + const float* s0 = sc + (x / 256) * 256 + cc; + float acc = 0.f; + for (int m = 0; m <= r; m++) acc = fmaf(tm[(16 * bi + r) * CP + 16 * bi + m], s0[m * 16], acc); + tm[(16 * bi + r) * CP + 16 * bj + cc] = -acc; + } + __syncthreads(); + } + + // 5. bfloat16 operands: k exp(cum) over A, T beta after T, then v over kn; + // u = T (beta v) and w = T (beta exp(cum) k) on tensor cores, stored as bfloat16 + for (int x = tid; x < C * K; x += THREADS) { + const int j = x / K, d = x % K; + kw[j * HB + d] = __float2bfloat16(kn[j * KP + d] * ecum[j]); + } + for (int x = tid; x < C * C; x += THREADS) { + const int i = x / C, j = x % C; + tb[i * TB + j] = __float2bfloat16(tm[i * CP + j] * bet[j]); + } + __syncthreads(); + for (int x = tid; x < C * V; x += THREADS) { + const int j = x / V, d = x % V; + const int t = c * C + j; + vb[j * HB + d] = t < T ? v[((size_t)t * H + h) * V + d] : __float2bfloat16(0.f); + } + __syncthreads(); + float* stage = reinterpret_cast(sm + R5) + warp * 256; + for (int e = warp; e < 64; e += THREADS / 32) { + const bool is_u = e < 32; + const int it = (e % 32) / 8, dt = e % 8; + const __nv_bfloat16* bsrc = is_u ? vb : kw; + CBf16 acc; + wmma::fill_fragment(acc, 0.f); + for (int kb = 0; kb <= it; kb++) { // T is lower triangular + ABf16 fa; + BBf16 fb; + wmma::load_matrix_sync(fa, tb + it * 16 * TB + kb * 16, TB); + wmma::load_matrix_sync(fb, bsrc + kb * 16 * HB + dt * 16, HB); + wmma::mma_sync(acc, fa, fb, acc); + } + wmma::store_matrix_sync(stage, acc, 16, wmma::mem_row_major); + __syncwarp(); + bf16* dst = (is_u ? ws.u : ws.w) + ((size_t)h * NCC + c * C + it * 16) * K + dt * 16; + for (int x = lane; x < 256; x += 32) dst[(x / 16) * K + x % 16] = __float2bfloat16(stage[x]); + __syncwarp(); + } +} + +// ---- kernel 2: the state, chunk by chunk ---- + +constexpr int BVS = 32; // value columns per block +constexpr int ST_THREADS = 128; +constexpr int WS_LD = K + 8; // bfloat16 row stride of staged w and kd +constexpr int SS_LD = BVS + 8; // bfloat16 row stride of the S copy and v_new +constexpr int STAGE = C * WS_LD; // elements of one staged w or kd +constexpr size_t SMEM2_BYTES = (4 * STAGE + K * SS_LD + C * SS_LD) * 2; + +__global__ void __launch_bounds__(ST_THREADS) gdn_chunk_state(Work ws, int NC) { + extern __shared__ __align__(128) unsigned char sm[]; + bf16* wbuf = reinterpret_cast(sm); // [2][C][WS_LD] + bf16* kbuf = wbuf + 2 * STAGE; // [2][C][WS_LD] + bf16* scopy = kbuf + 2 * STAGE; // [K][SS_LD] + bf16* vnew = scopy + K * SS_LD; // [C][SS_LD] + const int h = blockIdx.x, vb0 = blockIdx.y * BVS, NCC = NC * C; + const int tid = threadIdx.x, warp = tid / 32, lane = tid % 32, g = lane / 4, t = lane % 4; + + auto load = [&](int c, int buf) { + const size_t base = ((size_t)h * NCC + (size_t)c * C) * K; + for (int x = tid; x < C * K / 8; x += ST_THREADS) { + const int r = x / (K / 8), col = (x % (K / 8)) * 8; + cs1::cp_async16(wbuf + buf * STAGE + r * WS_LD + col, ws.w + base + (size_t)r * K + col); + cs1::cp_async16(kbuf + buf * STAGE + r * WS_LD + col, ws.kd + base + (size_t)r * K + col); + } + cs1::cp_async_commit(); + }; + + // S rows warp * 32 + mt * 16 + {g, g + 8}, columns nt * 8 + {2t, 2t + 1} + float st[2][4][4]; +#pragma unroll + for (int mt = 0; mt < 2; mt++) +#pragma unroll + for (int nt = 0; nt < 4; nt++) st[mt][nt][0] = st[mt][nt][1] = st[mt][nt][2] = st[mt][nt][3] = 0.f; + + load(0, 0); + for (int c = 0; c < NC; c++) { + const int buf = c & 1; + cs1::cp_async_wait<0>(); + __syncthreads(); // chunk c is staged, and chunk c - 1 is done with the other buffer + if (c + 1 < NC) load(c + 1, buf ^ 1); + // 1. S as bfloat16, to shared memory for w S and to global memory for the output + bf16* sg = ws.s + ((size_t)h * NC + c) * K * V + vb0; +#pragma unroll + for (int mt = 0; mt < 2; mt++) +#pragma unroll + for (int nt = 0; nt < 4; nt++) +#pragma unroll + for (int r = 0; r < 2; r++) { + const int row = warp * 32 + mt * 16 + g + r * 8, col = nt * 8 + 2 * t; + const uint32_t pk = cs1::pack_bf16(st[mt][nt][2 * r], st[mt][nt][2 * r + 1]); + *reinterpret_cast(scopy + row * SS_LD + col) = pk; + *reinterpret_cast(sg + (size_t)row * V + col) = pk; + } + __syncthreads(); + // 2. v_new = u - w S, rows warp * 16 .. + const bf16* wc = wbuf + buf * STAGE; + float acc[4][4]; +#pragma unroll + for (int nt = 0; nt < 4; nt++) acc[nt][0] = acc[nt][1] = acc[nt][2] = acc[nt][3] = 0.f; +#pragma unroll + for (int kk = 0; kk < K; kk += 16) { + uint32_t a[4]; + cs1::load_a(a, wc, WS_LD, warp * 16, kk, lane); +#pragma unroll + for (int nt = 0; nt < 4; nt += 2) { + uint32_t b[4]; + cs1::load_b_kn(b, scopy, SS_LD, kk, nt * 8, lane); + cs1::mma16816(acc[nt], a, b[0], b[1]); + cs1::mma16816(acc[nt + 1], a, b[2], b[3]); + } + } + const size_t row0 = (size_t)h * NCC + (size_t)c * C + warp * 16; +#pragma unroll + for (int nt = 0; nt < 4; nt++) +#pragma unroll + for (int r = 0; r < 2; r++) { + const int row = g + r * 8, col = nt * 8 + 2 * t; + const __nv_bfloat162 uv = + *reinterpret_cast(ws.u + (row0 + row) * V + vb0 + col); + const uint32_t pk = cs1::pack_bf16(__low2float(uv) - acc[nt][2 * r], __high2float(uv) - acc[nt][2 * r + 1]); + *reinterpret_cast(vnew + (warp * 16 + row) * SS_LD + col) = pk; + *reinterpret_cast(ws.vn + (row0 + row) * V + vb0 + col) = pk; + } + __syncthreads(); + // 3. S = decay S + kd^T v_new, rows warp * 32 .. + const float dec = ws.decay[h * NC + c]; +#pragma unroll + for (int mt = 0; mt < 2; mt++) +#pragma unroll + for (int nt = 0; nt < 4; nt++) +#pragma unroll + for (int e = 0; e < 4; e++) st[mt][nt][e] *= dec; + const bf16* kc = kbuf + buf * STAGE; +#pragma unroll + for (int kk = 0; kk < C; kk += 16) { +#pragma unroll + for (int mt = 0; mt < 2; mt++) { + uint32_t a[4]; + cs1::load_a_trans(a, kc, WS_LD, warp * 32 + mt * 16, kk, lane); +#pragma unroll + for (int nt = 0; nt < 4; nt += 2) { + uint32_t b[4]; + cs1::load_b_kn(b, vnew, SS_LD, kk, nt * 8, lane); + cs1::mma16816(st[mt][nt], a, b[0], b[1]); + cs1::mma16816(st[mt][nt + 1], a, b[2], b[3]); + } + } + } + } +} + +// ---- kernel 3: the output, per chunk ---- + +constexpr int OUT_THREADS = 128; +constexpr int O_LD = K + 8; // bfloat16 row stride of qd, S and v_new (all 128 wide) +constexpr int P_LD = C + 8; // bfloat16 row stride of P +constexpr size_t SMEM3_BYTES = (C * O_LD + K * O_LD + C * P_LD + C * O_LD) * 2; +static_assert(K == V, "S, qd and v_new share a row stride"); + +__global__ void __launch_bounds__(OUT_THREADS) gdn_chunk_out(Work ws, bf16* __restrict__ o, int T, int H) { + extern __shared__ __align__(128) unsigned char sm[]; + bf16* qs = reinterpret_cast(sm); // [C][O_LD] + bf16* ss = qs + C * O_LD; // [K][O_LD] + bf16* ps = ss + K * O_LD; // [C][P_LD] + bf16* vs = ps + C * P_LD; // [C][O_LD] + const int c = blockIdx.x, h = blockIdx.y, NC = gridDim.x, NCC = NC * C; + const int tid = threadIdx.x, warp = tid / 32, lane = tid % 32, g = lane / 4, t = lane % 4; + const size_t rows = (size_t)h * NCC + (size_t)c * C; + for (int x = tid; x < C * K / 8; x += OUT_THREADS) { + const int r = x / (K / 8), col = (x % (K / 8)) * 8; + cs1::cp_async16(qs + r * O_LD + col, ws.qd + (rows + r) * K + col); + cs1::cp_async16(vs + r * O_LD + col, ws.vn + (rows + r) * V + col); + } + for (int x = tid; x < K * V / 8; x += OUT_THREADS) { + const int r = x / (V / 8), col = (x % (V / 8)) * 8; + cs1::cp_async16(ss + r * O_LD + col, ws.s + ((size_t)h * NC + c) * K * V + (size_t)r * V + col); + } + for (int x = tid; x < C * C / 8; x += OUT_THREADS) { + const int r = x / (C / 8), col = (x % (C / 8)) * 8; + cs1::cp_async16(ps + r * P_LD + col, ws.pb + ((size_t)h * NC + c) * C * C + r * C + col); + } + cs1::cp_async_commit(); + cs1::cp_async_wait<0>(); + __syncthreads(); + + float acc[V / 8][4]; +#pragma unroll + for (int nt = 0; nt < V / 8; nt++) acc[nt][0] = acc[nt][1] = acc[nt][2] = acc[nt][3] = 0.f; + // qd S +#pragma unroll 2 + for (int kk = 0; kk < K; kk += 16) { + uint32_t a[4]; + cs1::load_a(a, qs, O_LD, warp * 16, kk, lane); +#pragma unroll + for (int nt = 0; nt < V / 8; nt += 2) { + uint32_t b[4]; + cs1::load_b_kn(b, ss, O_LD, kk, nt * 8, lane); + cs1::mma16816(acc[nt], a, b[0], b[1]); + cs1::mma16816(acc[nt + 1], a, b[2], b[3]); + } + } + // P v_new; P is lower triangular, so these rows need positions below (warp + 1) * 16 + for (int kk = 0; kk < (warp + 1) * 16; kk += 16) { + uint32_t a[4]; + cs1::load_a(a, ps, P_LD, warp * 16, kk, lane); +#pragma unroll + for (int nt = 0; nt < V / 8; nt += 2) { + uint32_t b[4]; + cs1::load_b_kn(b, vs, O_LD, kk, nt * 8, lane); + cs1::mma16816(acc[nt], a, b[0], b[1]); + cs1::mma16816(acc[nt + 1], a, b[2], b[3]); + } + } +#pragma unroll + for (int r = 0; r < 2; r++) { + const int tok = c * C + warp * 16 + g + r * 8; + if (tok >= T) continue; + bf16* dst = o + ((size_t)tok * H + h) * V + 2 * t; +#pragma unroll + for (int nt = 0; nt < V / 8; nt++) + *reinterpret_cast(dst + nt * 8) = cs1::pack_bf16(acc[nt][2 * r], acc[nt][2 * r + 1]); + } +} + +Work split(float* workspace, int T, int H) { + const Layout l(T, H); + unsigned char* b = reinterpret_cast(workspace); + Work ws; + ws.u = reinterpret_cast(b + l.u); + ws.w = reinterpret_cast(b + l.w); + ws.qd = reinterpret_cast(b + l.qd); + ws.kd = reinterpret_cast(b + l.kd); + ws.p = reinterpret_cast(b + l.p); + ws.pb = reinterpret_cast(b + l.pb); + ws.decay = reinterpret_cast(b + l.decay); + ws.s = reinterpret_cast(b + l.s); + ws.vn = reinterpret_cast(b + l.vn); + return ws; +} + +} // namespace + +extern "C" { + +size_t cs1_gdn_workspace_floats(int T, int H) { return (Layout(T, H).total + 3) / 4; } + +int cs1_gdn_prefill(const void* q, const void* k, const void* v, const float* g, const void* beta, + void* o, float* workspace, int T, int H, int HK, float scale, void* stream) { + if (T < 0 || HK <= 0 || H % HK != 0) return cudaErrorInvalidValue; + if (T == 0) return cudaSuccess; + // once per process (for the device current at the first call) + static const cudaError_t configured = [] { + cudaError_t e = cudaFuncSetAttribute(gdn_chunk_prep, cudaFuncAttributeMaxDynamicSharedMemorySize, + (int)SMEM1_BYTES); + if (e == cudaSuccess) + e = cudaFuncSetAttribute(gdn_chunk_state, cudaFuncAttributeMaxDynamicSharedMemorySize, + (int)SMEM2_BYTES); + if (e == cudaSuccess) + e = cudaFuncSetAttribute(gdn_chunk_out, cudaFuncAttributeMaxDynamicSharedMemorySize, + (int)SMEM3_BYTES); + return e; + }(); + if (configured != cudaSuccess) return configured; + const int NC = (T + C - 1) / C; + const Work ws = split(workspace, T, H); + cudaStream_t st = static_cast(stream); + gdn_chunk_prep<<>>( + static_cast(q), static_cast(k), + static_cast(v), g, static_cast(beta), ws, T, H, HK, scale); + gdn_chunk_state<<>>(ws, NC); + gdn_chunk_out<<>>(ws, static_cast(o), T, H); + return cudaGetLastError(); +} + +} // extern "C" diff --git a/src/backends/cuda/qwen3_5/gemm.cu b/src/backends/cuda/qwen3_5/gemm.cu new file mode 100644 index 00000000..22f0add6 --- /dev/null +++ b/src/backends/cuda/qwen3_5/gemm.cu @@ -0,0 +1,454 @@ +// bfloat16 GEMMs through cuBLASLt, float32 accumulation. +// +// Row-major y [M, N] = x [M, K] * w [N, K]^T is the column-major product +// y^T [N, M] = (w viewed as [K, N])^T * (x viewed as [K, M]); y's rows may be +// strided (ldy >= N), so one GEMM can fill a slice of a wider buffer. +// +// Algorithms: cs1_gemm_tune times cuBLASLt's candidates for a shape with L2 flushed +// before every call and keeps the fastest: the heuristic's shortlist, or (exhaustive) +// each algorithm id with each tile, stage count, custom option and swizzle it supports +// and split-K factors of 1 to 6, 8, 12 and 16, as far as cuBLASLt accepts them for the +// shape (a first pass of one call each keeps 12 to time properly). +// Many configurations are within noise of each other, so two searches often keep +// different ones with about the same speed. +// It replaces the heuristic's first choice only when it is more than 3% faster, so +// near ties rarely change between runs, and +// cs1_gemm_export / cs1_gemm_import let a caller keep the choices across runs. A shape +// that was not tuned uses the algorithm tuned for the nearest M with the same N, K +// and ldy (the smallest tuned M above it, else the largest below), or the heuristic's +// first choice. Split-K reductions that accumulate into the output in place are +// excluded, since their order, and so the rounding, is not fixed. +#include + +#include +#include +#include +#include +#include +#include +#include + +#include "ops.h" + +namespace { + +struct Plan { + cublasLtMatmulDesc_t op = nullptr; + cublasLtMatrixLayout_t a = nullptr, b = nullptr, c = nullptr; + cublasLtMatmulAlgo_t algo{}; + bool tuned = false; +}; + +using Key = std::tuple; // M, N, K, ldy + +constexpr size_t FLUSH_BYTES = 256u << 20; + +struct Gemm { + cublasLtHandle_t handle = nullptr; + void* workspace = nullptr; + size_t workspace_bytes = 0; + std::map plans; + // tuning only: a buffer larger than L2, a sink for its reads, and two events + void* flush = nullptr; + int* sink = nullptr; + cudaEvent_t e0 = nullptr, e1 = nullptr; +}; + +void release_tuning(Gemm& g) { + if (g.flush) cudaFree(g.flush); + if (g.sink) cudaFree(g.sink); + if (g.e0) cudaEventDestroy(g.e0); + if (g.e1) cudaEventDestroy(g.e1); + g.flush = nullptr; + g.sink = nullptr; + g.e0 = g.e1 = nullptr; +} + +int status(cublasStatus_t s) { return s == CUBLAS_STATUS_SUCCESS ? 0 : 1000 + (int)s; } + +void destroy(Plan& p) { + if (p.a) cublasLtMatrixLayoutDestroy(p.a); + if (p.b) cublasLtMatrixLayoutDestroy(p.b); + if (p.c) cublasLtMatrixLayoutDestroy(p.c); + if (p.op) cublasLtMatmulDescDestroy(p.op); + p = Plan{}; +} + +int describe(int M, int N, int K, int ldy, Plan& p) { + cublasStatus_t s = cublasLtMatmulDescCreate(&p.op, CUBLAS_COMPUTE_32F, CUDA_R_32F); + if (s != CUBLAS_STATUS_SUCCESS) return status(s); + const cublasOperation_t ta = CUBLAS_OP_T, tb = CUBLAS_OP_N; + cublasLtMatmulDescSetAttribute(p.op, CUBLASLT_MATMUL_DESC_TRANSA, &ta, sizeof(ta)); + cublasLtMatmulDescSetAttribute(p.op, CUBLASLT_MATMUL_DESC_TRANSB, &tb, sizeof(tb)); + if ((s = cublasLtMatrixLayoutCreate(&p.a, CUDA_R_16BF, K, N, K)) != CUBLAS_STATUS_SUCCESS) return status(s); + if ((s = cublasLtMatrixLayoutCreate(&p.b, CUDA_R_16BF, K, M, K)) != CUBLAS_STATUS_SUCCESS) return status(s); + if ((s = cublasLtMatrixLayoutCreate(&p.c, CUDA_R_16BF, N, M, ldy)) != CUBLAS_STATUS_SUCCESS) return status(s); + return 0; +} + +int heuristics(Gemm& g, const Plan& p, int want, std::vector& out) { + cublasLtMatmulPreference_t pref; + cublasStatus_t s = cublasLtMatmulPreferenceCreate(&pref); + if (s != CUBLAS_STATUS_SUCCESS) return status(s); + cublasLtMatmulPreferenceSetAttribute(pref, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, &g.workspace_bytes, + sizeof(g.workspace_bytes)); + const uint32_t schemes = CUBLASLT_REDUCTION_SCHEME_MASK & ~CUBLASLT_REDUCTION_SCHEME_INPLACE; + cublasLtMatmulPreferenceSetAttribute(pref, CUBLASLT_MATMUL_PREF_REDUCTION_SCHEME_MASK, &schemes, + sizeof(schemes)); + out.resize(want); + int found = 0; + s = cublasLtMatmulAlgoGetHeuristic(g.handle, p.op, p.a, p.b, p.c, p.c, pref, want, out.data(), &found); + cublasLtMatmulPreferenceDestroy(pref); + if (s != CUBLAS_STATUS_SUCCESS) return status(s); + out.resize(found); + out.erase(std::remove_if(out.begin(), out.end(), + [](const cublasLtMatmulHeuristicResult_t& r) { return r.state != CUBLAS_STATUS_SUCCESS; }), + out.end()); + if (out.empty()) return status(CUBLAS_STATUS_NOT_SUPPORTED); + return 0; +} + +bool usable(Gemm& g, const Plan& p, const cublasLtMatmulAlgo_t& algo) { + cublasLtMatmulHeuristicResult_t r{}; + return cublasLtMatmulAlgoCheck(g.handle, p.op, p.a, p.b, p.c, p.c, &algo, &r) == CUBLAS_STATUS_SUCCESS && + r.workspaceSize <= g.workspace_bytes; +} + +bool reduces_in_place(const cublasLtMatmulAlgo_t& algo) { + uint32_t red = CUBLASLT_REDUCTION_SCHEME_NONE; + size_t n = 0; + cublasLtMatmulAlgoConfigGetAttribute(&algo, CUBLASLT_ALGO_CONFIG_REDUCTION_SCHEME, &red, sizeof red, &n); + return red == CUBLASLT_REDUCTION_SCHEME_INPLACE; +} + +template +std::vector cap_array(const cublasLtMatmulAlgo_t& algo, cublasLtMatmulAlgoCapAttributes_t attr) { + size_t bytes = 0; + cublasLtMatmulAlgoCapGetAttribute(&algo, attr, nullptr, 0, &bytes); + std::vector v(bytes / sizeof(T)); + if (bytes) cublasLtMatmulAlgoCapGetAttribute(&algo, attr, v.data(), bytes, &bytes); + return v; +} + +template +T cap(const cublasLtMatmulAlgo_t& algo, cublasLtMatmulAlgoCapAttributes_t attr) { + T v{}; + size_t n; + cublasLtMatmulAlgoCapGetAttribute(&algo, attr, &v, sizeof v, &n); + return v; +} + +// The configurations cuBLASLt accepts for the shape, within the workspace: each +// algorithm id with each tile, stage count, custom option and swizzle it supports, and +// split-K factors from `splits` (reduced in the compute or the output type, not in +// place). Other attributes stay at their defaults. +std::vector every_config(Gemm& g, const Plan& p) { + int ids[256], nids = 0; + std::vector out; + if (cublasLtMatmulAlgoGetIds(g.handle, CUBLAS_COMPUTE_32F, CUDA_R_32F, CUDA_R_16BF, CUDA_R_16BF, CUDA_R_16BF, + CUDA_R_16BF, 256, ids, &nids) != CUBLAS_STATUS_SUCCESS) + return out; + const int splits[] = {1, 2, 3, 4, 5, 6, 8, 12, 16}; + const uint32_t schemes[] = {CUBLASLT_REDUCTION_SCHEME_NONE, CUBLASLT_REDUCTION_SCHEME_COMPUTE_TYPE, + CUBLASLT_REDUCTION_SCHEME_OUTPUT_TYPE}; + for (int i = 0; i < nids; i++) { + cublasLtMatmulAlgo_t base; + if (cublasLtMatmulAlgoInit(g.handle, CUBLAS_COMPUTE_32F, CUDA_R_32F, CUDA_R_16BF, CUDA_R_16BF, CUDA_R_16BF, + CUDA_R_16BF, ids[i], &base) != CUBLAS_STATUS_SUCCESS) + continue; + auto tiles = cap_array(base, CUBLASLT_ALGO_CAP_TILE_IDS); + auto stages = cap_array(base, CUBLASLT_ALGO_CAP_STAGES_IDS); + if (tiles.empty()) tiles.push_back(CUBLASLT_MATMUL_TILE_UNDEFINED); + if (stages.empty()) stages.push_back(CUBLASLT_MATMUL_STAGES_UNDEFINED); + const int splitk_ok = cap(base, CUBLASLT_ALGO_CAP_SPLITK_SUPPORT); + const uint32_t red_mask = cap(base, CUBLASLT_ALGO_CAP_REDUCTION_SCHEME_MASK); + const int swizzle_ok = cap(base, CUBLASLT_ALGO_CAP_CTA_SWIZZLING_SUPPORT); + const int custom_max = cap(base, CUBLASLT_ALGO_CAP_CUSTOM_OPTION_MAX); + for (uint32_t tile : tiles) + for (uint32_t stage : stages) + for (int custom = 0; custom <= custom_max; custom++) + for (int swz = 0; swz <= swizzle_ok; swz++) + for (int sk : splits) { + if (sk > 1 && !splitk_ok) break; + for (uint32_t red : schemes) { + if ((sk == 1) != (red == CUBLASLT_REDUCTION_SCHEME_NONE)) continue; + if (red != CUBLASLT_REDUCTION_SCHEME_NONE && !(red_mask & red)) continue; + cublasLtMatmulAlgo_t a = base; + cublasLtMatmulAlgoConfigSetAttribute(&a, CUBLASLT_ALGO_CONFIG_TILE_ID, &tile, + sizeof tile); + cublasLtMatmulAlgoConfigSetAttribute(&a, CUBLASLT_ALGO_CONFIG_STAGES_ID, &stage, + sizeof stage); + cublasLtMatmulAlgoConfigSetAttribute(&a, CUBLASLT_ALGO_CONFIG_CUSTOM_OPTION, &custom, + sizeof custom); + cublasLtMatmulAlgoConfigSetAttribute(&a, CUBLASLT_ALGO_CONFIG_CTA_SWIZZLING, &swz, + sizeof swz); + cublasLtMatmulAlgoConfigSetAttribute(&a, CUBLASLT_ALGO_CONFIG_SPLITK_NUM, &sk, + sizeof sk); + cublasLtMatmulAlgoConfigSetAttribute(&a, CUBLASLT_ALGO_CONFIG_REDUCTION_SCHEME, &red, + sizeof red); + if (usable(g, p, a)) out.push_back(a); + } + } + } + return out; +} + +// The plan for a shape, created on first use. +int plan_for(Gemm& g, int M, int N, int K, int ldy, Plan*& out) { + const Key key{M, N, K, ldy}; + auto it = g.plans.find(key); + if (it != g.plans.end()) { + out = &it->second; + return 0; + } + Plan p; + int rc = describe(M, N, K, ldy, p); + if (rc != 0) { + destroy(p); + return rc; + } + // The algorithm tuned for the nearest M: the smallest above, else the largest below. + const Plan* above = nullptr; + const Plan* below = nullptr; + int above_m = 0, below_m = 0; + for (auto& kv : g.plans) { + const auto [m, n, k, l] = kv.first; + if (n != N || k != K || l != ldy || !kv.second.tuned) continue; + if (m > M && (!above || m < above_m)) above = &kv.second, above_m = m; + if (m < M && (!below || m > below_m)) below = &kv.second, below_m = m; + } + if (above && usable(g, p, above->algo)) { + p.algo = above->algo; + } else if (below && usable(g, p, below->algo)) { + p.algo = below->algo; + } else { + std::vector cands; + rc = heuristics(g, p, 1, cands); + if (rc != 0) { + destroy(p); + return rc; + } + p.algo = cands[0].algo; + } + out = &g.plans.emplace(key, p).first->second; + return 0; +} + +// With CUA_S1_GEMM_LOG set, print each tuned choice to stderr. +void log_choice(int M, int N, int K, const cublasLtMatmulAlgo_t& a, float ms, float first_ms, size_t pick) { + static const bool on = std::getenv("CUA_S1_GEMM_LOG") != nullptr; + if (!on) return; + int tile = 0, stages = 0, splitk = 0, inner = 0, id = 0; + size_t n; + cublasLtMatmulAlgoConfigGetAttribute(&a, CUBLASLT_ALGO_CONFIG_ID, &id, sizeof(int), &n); + cublasLtMatmulAlgoConfigGetAttribute(&a, CUBLASLT_ALGO_CONFIG_TILE_ID, &tile, sizeof(int), &n); + cublasLtMatmulAlgoConfigGetAttribute(&a, CUBLASLT_ALGO_CONFIG_STAGES_ID, &stages, sizeof(int), &n); + cublasLtMatmulAlgoConfigGetAttribute(&a, CUBLASLT_ALGO_CONFIG_SPLITK_NUM, &splitk, sizeof(int), &n); + cublasLtMatmulAlgoConfigGetAttribute(&a, CUBLASLT_ALGO_CONFIG_INNER_SHAPE_ID, &inner, sizeof(int), &n); + fprintf(stderr, "gemm %5d x %5d x %5d: candidate %zu, algo %d tile %d stages %d splitK %d inner %d, %.1f us (first %.1f us)\n", + M, N, K, pick, id, tile, stages, splitk, inner, ms * 1e3f, first_ms * 1e3f); +} + +// Read a buffer larger than L2, so the next call finds none of its operands cached. +__global__ void flush_l2(const int4* p, size_t n, int* sink) { + int acc = 0; + for (size_t i = blockIdx.x * (size_t)blockDim.x + threadIdx.x; i < n; i += (size_t)gridDim.x * blockDim.x) + acc ^= p[i].x ^ p[i].w; + if (acc == 0x7fffffff) *sink = acc; +} + +} // namespace + +extern "C" void* cs1_gemm_create(size_t workspace_bytes) { + Gemm* g = new Gemm(); + if (cublasLtCreate(&g->handle) != CUBLAS_STATUS_SUCCESS || + (workspace_bytes > 0 && cudaMalloc(&g->workspace, workspace_bytes) != cudaSuccess)) { + if (g->handle) cublasLtDestroy(g->handle); + delete g; + return nullptr; + } + g->workspace_bytes = workspace_bytes; + return g; +} + +extern "C" void cs1_gemm_destroy(void* gemm) { + Gemm* g = static_cast(gemm); + if (!g) return; + for (auto& kv : g->plans) destroy(kv.second); + release_tuning(*g); + if (g->workspace) cudaFree(g->workspace); + cublasLtDestroy(g->handle); + delete g; +} + +extern "C" int cs1_gemm_tune(void* gemm, const void* x, const void* w, void* y, int M, int N, int K, int ldy, + int exhaustive, void* stream) { + Gemm* g = static_cast(gemm); + if (!g || M <= 0 || N <= 0 || K <= 0 || ldy < N) return cudaErrorInvalidValue; + const Key key{M, N, K, ldy}; + auto it = g->plans.find(key); + if (it != g->plans.end() && it->second.tuned) return 0; + if (it != g->plans.end()) { + destroy(it->second); + g->plans.erase(it); + } + Plan p; + int rc = describe(M, N, K, ldy, p); + std::vector shortlist; + if (rc == 0) rc = heuristics(*g, p, 16, shortlist); + if (rc != 0) { + destroy(p); + return rc; + } + // candidates: the heuristic's shortlist first, then (exhaustive) those of every_config + std::vector cands; + for (auto& r : shortlist) cands.push_back(r.algo); + if (exhaustive) { + for (auto& a : every_config(*g, p)) + if (std::none_of(cands.begin(), cands.end(), + [&](const cublasLtMatmulAlgo_t& c) { return std::memcmp(&c, &a, sizeof a) == 0; })) + cands.push_back(a); + } + cudaStream_t st = static_cast(stream); + if (!g->flush) { + if (cudaMalloc(&g->flush, FLUSH_BYTES) != cudaSuccess || cudaMalloc(&g->sink, sizeof(int)) != cudaSuccess || + cudaMemsetAsync(g->flush, 0, FLUSH_BYTES, st) != cudaSuccess || cudaEventCreate(&g->e0) != cudaSuccess || + cudaEventCreate(&g->e1) != cudaSuccess) { + release_tuning(*g); + destroy(p); + cudaGetLastError(); + return (int)cudaErrorMemoryAllocation; + } + } + const float alpha = 1.f, beta = 0.f; + // median of `reps` calls, each after an L2 flush; a huge value if the call fails + auto time = [&](const cublasLtMatmulAlgo_t& algo, int reps) { + if (cublasLtMatmul(g->handle, p.op, &alpha, w, p.a, x, p.b, &beta, y, p.c, y, p.c, &algo, g->workspace, + g->workspace_bytes, st) != CUBLAS_STATUS_SUCCESS) { + cudaGetLastError(); + return 1e30f; + } + std::vector times; + for (int r = 0; r < reps; r++) { + flush_l2<<<1024, 256, 0, st>>>(static_cast(g->flush), FLUSH_BYTES / sizeof(int4), g->sink); + cudaEventRecord(g->e0, st); + cublasLtMatmul(g->handle, p.op, &alpha, w, p.a, x, p.b, &beta, y, p.c, y, p.c, &algo, g->workspace, + g->workspace_bytes, st); + cudaEventRecord(g->e1, st); + if (cudaEventSynchronize(g->e1) != cudaSuccess) return 1e30f; + float ms = 0.f; + cudaEventElapsedTime(&ms, g->e0, g->e1); + times.push_back(ms); + } + std::sort(times.begin(), times.end()); + return times[reps / 2]; + }; + // with many candidates, one timed call each picks the 12 to time properly + std::vector keep; + if (cands.size() > 16) { + std::vector> quick; + for (size_t i = 0; i < cands.size(); i++) quick.push_back({time(cands[i], 1), i}); + std::sort(quick.begin(), quick.end()); + keep.push_back(0); // the heuristic's first choice, the baseline + for (size_t j = 0; j < quick.size() && keep.size() < 13; j++) + if (quick[j].second != 0 && quick[j].first < 1e30f) keep.push_back(quick[j].second); + } else { + for (size_t i = 0; i < cands.size(); i++) keep.push_back(i); + } + // nine timed calls, or three for shapes that take over 2 ms + const int reps = time(cands[0], 1) > 2.f ? 3 : 9; + std::vector median(cands.size(), 1e30f); + for (size_t i : keep) median[i] = time(cands[i], reps); + // the heuristic's first working choice, unless another is more than 3% faster + size_t pick = 0; + while (pick < median.size() && median[pick] >= 1e30f) pick++; + if (pick == median.size()) { + destroy(p); + return status(CUBLAS_STATUS_NOT_SUPPORTED); + } + const size_t first = pick; + const size_t fastest = std::min_element(median.begin(), median.end()) - median.begin(); + if (median[fastest] < 0.97f * median[pick]) pick = fastest; + p.algo = cands[pick]; + p.tuned = true; + log_choice(M, N, K, p.algo, median[pick], median[first], pick); + g->plans.emplace(key, p); + return (int)cudaGetLastError(); +} + +extern "C" void cs1_gemm_tune_done(void* gemm) { + if (gemm) release_tuning(*static_cast(gemm)); +} + +extern "C" size_t cs1_gemm_export(void* gemm, Cs1GemmPlan* out, size_t cap) { + Gemm* g = static_cast(gemm); + if (!g) return 0; + size_t n = 0; + for (auto& kv : g->plans) { + if (!kv.second.tuned) continue; + if (n < cap) { + const auto [m, nn, k, l] = kv.first; + out[n] = Cs1GemmPlan{m, nn, k, l, {}}; + static_assert(sizeof(cublasLtMatmulAlgo_t) == sizeof(out[n].algo), "algo layout"); + std::memcpy(out[n].algo, &kv.second.algo, sizeof(out[n].algo)); + } + n++; + } + return n; +} + +extern "C" int cs1_gemm_import(void* gemm, const Cs1GemmPlan* plans, size_t n) { + Gemm* g = static_cast(gemm); + if (!g) return cudaErrorInvalidValue; + // check every plan before using any, so that a rejected file changes nothing + std::vector> checked; + int rc = 0; + for (size_t i = 0; i < n && rc == 0; i++) { + const Cs1GemmPlan& r = plans[i]; + if (r.m <= 0 || r.n <= 0 || r.k <= 0 || r.ldy < r.n) { + rc = cudaErrorInvalidValue; + break; + } + Plan p; + rc = describe(r.m, r.n, r.k, r.ldy, p); + if (rc == 0) { + std::memcpy(&p.algo, r.algo, sizeof(p.algo)); + if (reduces_in_place(p.algo) || !usable(*g, p, p.algo)) rc = status(CUBLAS_STATUS_NOT_SUPPORTED); + } + if (rc != 0) { + destroy(p); + break; + } + p.tuned = true; + checked.emplace_back(Key{r.m, r.n, r.k, r.ldy}, p); + } + if (rc != 0) { + for (auto& kv : checked) destroy(kv.second); + return rc; + } + for (auto& [key, p] : checked) { + auto it = g->plans.find(key); + if (it != g->plans.end()) { + destroy(it->second); + it->second = p; + } else { + g->plans.emplace(key, p); + } + } + return 0; +} + +extern "C" size_t cs1_gemm_version(void) { return cublasLtGetVersion(); } + +extern "C" int cs1_gemm(void* gemm, const void* x, const void* w, void* y, int M, int N, int K, int ldy, + void* stream) { + Gemm* g = static_cast(gemm); + if (!g || M < 0 || N <= 0 || K <= 0 || ldy < N) return cudaErrorInvalidValue; + if (M == 0) return cudaSuccess; + Plan* p = nullptr; + const int rc = plan_for(*g, M, N, K, ldy, p); + if (rc != 0) return rc; + const float alpha = 1.f, beta = 0.f; + return status(cublasLtMatmul(g->handle, p->op, &alpha, w, p->a, x, p->b, &beta, y, p->c, y, p->c, &p->algo, + g->workspace, g->workspace_bytes, static_cast(stream))); +} diff --git a/src/backends/cuda/qwen3_5/mma.cuh b/src/backends/cuda/qwen3_5/mma.cuh new file mode 100644 index 00000000..e865fb75 --- /dev/null +++ b/src/backends/cuda/qwen3_5/mma.cuh @@ -0,0 +1,74 @@ +// Tensor-core building blocks for sm_80 and later: mma.sync m16n8k16 on bfloat16 with +// float32 accumulation, ldmatrix, and cp.async. +// +// Fragment layouts (PTX ISA, mma.m16n8k16): with g = lane / 4 and t = lane % 4, an +// accumulator holds rows g (elements 0, 1) and g + 8 (elements 2, 3) at columns 2t and +// 2t + 1 of its 8-column tile; the A operand registers are (row g, cols 2t..), +// (row g + 8, cols 2t..), (row g, cols 8 + 2t..) and (row g + 8, cols 8 + 2t..). +#pragma once + +#include "common.cuh" + +namespace cs1 { + +__device__ __forceinline__ uint32_t smem_addr(const void* p) { + return static_cast(__cvta_generic_to_shared(p)); +} + +// 16-byte asynchronous copy from global to shared memory; zero-fills when !valid. +__device__ __forceinline__ void cp_async16(void* dst, const void* src, bool valid = true) { + asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;\n" ::"r"(smem_addr(dst)), "l"(src), + "r"(valid ? 16 : 0)); +} +__device__ __forceinline__ void cp_async_commit() { asm volatile("cp.async.commit_group;\n" ::); } +template +__device__ __forceinline__ void cp_async_wait() { + asm volatile("cp.async.wait_group %0;\n" ::"n"(N)); +} + +__device__ __forceinline__ void ldmatrix_x4(uint32_t (&r)[4], const bf16* p) { + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];\n" + : "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3]) + : "r"(smem_addr(p))); +} +__device__ __forceinline__ void ldmatrix_x4_trans(uint32_t (&r)[4], const bf16* p) { + asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.shared.b16 {%0,%1,%2,%3}, [%4];\n" + : "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3]) + : "r"(smem_addr(p))); +} + +// A fragment of the 16x16 tile at (row0, col0) of a row-major shared matrix. +__device__ __forceinline__ void load_a(uint32_t (&a)[4], const bf16* m, int ld, int row0, int col0, int lane) { + ldmatrix_x4(a, m + (row0 + (lane % 8) + ((lane / 8) % 2) * 8) * ld + col0 + (lane / 16) * 8); +} +// A fragment of the 16x16 tile at (row0, col0) of the transpose of a row-major shared +// matrix: rows of A are columns of m. +__device__ __forceinline__ void load_a_trans(uint32_t (&a)[4], const bf16* m, int ld, int row0, int col0, + int lane) { + ldmatrix_x4_trans(a, m + (col0 + (lane % 8) + (lane / 16) * 8) * ld + row0 + ((lane / 8) % 2) * 8); +} +// B fragments of two 16x8 tiles (k0.., n0..) and (k0.., n0 + 8..) of a row-major +// shared matrix whose rows are k: b[0], b[1] and b[2], b[3]. +__device__ __forceinline__ void load_b_kn(uint32_t (&b)[4], const bf16* m, int ld, int k0, int n0, int lane) { + ldmatrix_x4_trans(b, m + (k0 + (lane % 8) + ((lane / 8) % 2) * 8) * ld + n0 + (lane / 16) * 8); +} +// The same when the shared matrix is stored with rows n (each row contiguous in k). +__device__ __forceinline__ void load_b_nk(uint32_t (&b)[4], const bf16* m, int ld, int k0, int n0, int lane) { + ldmatrix_x4(b, m + (n0 + (lane % 8) + (lane / 16) * 8) * ld + k0 + ((lane / 8) % 2) * 8); +} + +// d += a * b for a 16x16 (row-major) by 16x8 (column-major) tile. +__device__ __forceinline__ void mma16816(float (&d)[4], const uint32_t (&a)[4], uint32_t b0, uint32_t b1) { + asm volatile( + "mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, " + "{%0,%1,%2,%3};\n" + : "+f"(d[0]), "+f"(d[1]), "+f"(d[2]), "+f"(d[3]) + : "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b0), "r"(b1)); +} + +__device__ __forceinline__ uint32_t pack_bf16(float lo, float hi) { + const __nv_bfloat162 v = __floats2bfloat162_rn(lo, hi); + return *reinterpret_cast(&v); +} + +} // namespace cs1 diff --git a/src/backends/cuda/qwen3_5/norm.cu b/src/backends/cuda/qwen3_5/norm.cu new file mode 100644 index 00000000..38d6f98a --- /dev/null +++ b/src/backends/cuda/qwen3_5/norm.cu @@ -0,0 +1,106 @@ +// RMSNorm variants of Qwen3.5. +#include "common.cuh" +#include "ops.h" + +namespace cs1 { +namespace { + +constexpr int NORM_THREADS = 256; + +// Qwen3_5RMSNorm: float32 normalization, times (1 + w), rounded once to bfloat16. +__global__ void __launch_bounds__(NORM_THREADS) + rms_norm_kernel(const bf16* __restrict__ x, const bf16* __restrict__ w, bf16* __restrict__ out, int D, + float eps) { + __shared__ float scratch[32]; + const size_t row = blockIdx.x; + x += row * D; + out += row * D; + float ss = 0.f; + for (int i = threadIdx.x; i < D; i += blockDim.x) { + const float v = f32(x[i]); + ss += v * v; + } + const float inv = rsqrtf(block_sum(ss, scratch) / D + eps); + for (int i = threadIdx.x; i < D; i += blockDim.x) out[i] = to_bf16(f32(x[i]) * inv * (1.f + f32(w[i]))); +} + +// The residual add of a decoder layer (bfloat16 + bfloat16, rounded) fused with the +// following Qwen3_5RMSNorm. Each thread rereads only the elements it wrote. +__global__ void __launch_bounds__(NORM_THREADS) + add_rms_norm_kernel(bf16* __restrict__ residual, const bf16* __restrict__ delta, const bf16* __restrict__ w, + bf16* __restrict__ out, int D, float eps) { + __shared__ float scratch[32]; + const size_t row = blockIdx.x; + residual += row * D; + delta += row * D; + out += row * D; + float ss = 0.f; + for (int i = threadIdx.x; i < D; i += blockDim.x) { + const bf16 r = to_bf16(f32(residual[i]) + f32(delta[i])); + residual[i] = r; + const float v = f32(r); + ss += v * v; + } + const float inv = rsqrtf(block_sum(ss, scratch) / D + eps); + for (int i = threadIdx.x; i < D; i += blockDim.x) + out[i] = to_bf16(f32(residual[i]) * inv * (1.f + f32(w[i]))); +} + +// Qwen3_5RMSNormGated with D = 128, one warp per row: the normalized value is +// rounded to bfloat16, multiplied by w in bfloat16, then by silu(z) in float32. +__global__ void gated_rms_norm_kernel(const bf16* __restrict__ x, const bf16* __restrict__ z, int ldz, + const bf16* __restrict__ w, bf16* __restrict__ out, int rows, int H, + float eps) { + constexpr int D = 128, PER = D / 32; + const int row = blockIdx.x * (blockDim.x / 32) + threadIdx.x / 32, lane = threadIdx.x & 31; + if (row >= rows) return; + const size_t base = (size_t)row * D + lane * PER; + const size_t zbase = (size_t)(row / H) * ldz + (row % H) * D + lane * PER; + float v[PER]; + float ss = 0.f; +#pragma unroll + for (int i = 0; i < PER; i++) { + v[i] = f32(x[base + i]); + ss += v[i] * v[i]; + } + const float inv = rsqrtf(warp_sum(ss) / D + eps); +#pragma unroll + for (int i = 0; i < PER; i++) { + const float normed = round_bf16(v[i] * inv); + const float scaled = round_bf16(f32(w[lane * PER + i]) * normed); + out[base + i] = to_bf16(scaled * silu(f32(z[zbase + i]))); + } +} + +} // namespace +} // namespace cs1 + +using namespace cs1; + +extern "C" int cs1_rms_norm(const void* x, const void* w, void* out, int rows, int D, float eps, void* stream) { + if (rows <= 0) return cudaSuccess; + rms_norm_kernel<<(stream)>>>( + static_cast(x), static_cast(w), static_cast(out), D, eps); + return cudaGetLastError(); +} + +extern "C" int cs1_add_rms_norm(void* residual, const void* delta, const void* w, void* out, int rows, int D, + float eps, void* stream) { + if (rows <= 0) return cudaSuccess; + add_rms_norm_kernel<<(stream)>>>( + static_cast(residual), static_cast(delta), static_cast(w), + static_cast(out), D, eps); + return cudaGetLastError(); +} + +extern "C" int cs1_gated_rms_norm(const void* x, const void* z, int ldz, const void* w, void* out, int T, int H, + int D, float eps, void* stream) { + if (D != 128 || ldz < H * D) return cudaErrorInvalidValue; + const int rows = T * H; + if (rows <= 0) return cudaSuccess; + constexpr int WARPS = 8; + gated_rms_norm_kernel<<<(rows + WARPS - 1) / WARPS, WARPS * 32, 0, static_cast(stream)>>>( + static_cast(x), static_cast(z), ldz, static_cast(w), + static_cast(out), rows, H, eps); + return cudaGetLastError(); +} diff --git a/src/backends/cuda/qwen3_5/ops.h b/src/backends/cuda/qwen3_5/ops.h new file mode 100644 index 00000000..70ce39ca --- /dev/null +++ b/src/backends/cuda/qwen3_5/ops.h @@ -0,0 +1,125 @@ +// C interface of libqwen3_5_cuda.so: the Qwen3.5 prefill operations and the few +// CUDA runtime calls a caller needs, so that it can load this one library at run time. +// +// Tensors are row-major and bfloat16 unless noted. Every operation queues work on +// `stream` (a cudaStream_t) and returns a cudaError_t, or 1000 + a cublasStatus_t for +// the GEMMs. The norm, elementwise and attention-prep operations round to bfloat16 +// where the Transformers reference (modeling_qwen3_5.py) does; attention and the +// Gated DeltaNet prefill keep some intermediate results in bfloat16, as FlashAttention +// and flash-linear-attention do (see attention.cu and gdn_prefill.cu). +#pragma once + +#include +#include + +// Bumped whenever a signature below changes. +#define CS1_ABI_VERSION 1 + +#ifdef __cplusplus +extern "C" { +#endif + +// ---- runtime ---- + +uint32_t cs1_abi_version(void); +const char* cs1_error_string(int code); +int cs1_set_device(int device); +// Name, compute capability (major * 10 + minor) and SM count of the current device. +int cs1_device_info(char* name, size_t cap, int* compute_capability, int* sms); +int cs1_malloc(void** ptr, size_t bytes); +int cs1_free(void* ptr); +int cs1_stream_create(void** stream); +int cs1_stream_sync(void* stream); +// Copy and wait for the copy. +int cs1_upload(void* dst, const void* src, size_t bytes, void* stream); +int cs1_download(void* dst, const void* src, size_t bytes, void* stream); +// Capture the work queued on `stream` between begin and end into an executable graph. +int cs1_graph_begin(void* stream); +int cs1_graph_end(void* stream, void** exec); +int cs1_graph_launch(void* exec, void* stream); +int cs1_graph_destroy(void* exec); + +// ---- operations ---- + +// out[t] = table[ids[t]], rows of D. +int cs1_embed(const int32_t* ids, const void* table, void* out, int T, int D, void* stream); + +// Zero-centred RMSNorm over rows of D: out = x / rms(x) * (1 + w), in float32. +int cs1_rms_norm(const void* x, const void* w, void* out, int rows, int D, float eps, void* stream); + +// residual = residual + delta (rounded to bfloat16), then out = cs1_rms_norm(residual). +int cs1_add_rms_norm(void* residual, const void* delta, const void* w, void* out, int rows, int D, + float eps, void* stream); + +// Gated RMSNorm of the Gated DeltaNet output x [T, H, D] (D = 128), with z [T, H*D] +// in rows of ldz: out = (w * (x / rms(x))) * silu(z). +int cs1_gated_rms_norm(const void* x, const void* z, int ldz, const void* w, void* out, int T, int H, + int D, float eps, void* stream); + +// Depthwise causal conv1d (kernel 4, no bias) and SiLU over qkv [T, ld], split into +// q [T, key_dim], k [T, key_dim] and v [T, value_dim]. w is [key_dim*2 + value_dim, 4]. +int cs1_gdn_conv(const void* qkv, int ld, const void* w, void* q, void* k, void* v, int T, int key_dim, + int value_dim, void* stream); + +// beta = sigmoid(b) (bfloat16) and g = -exp(A_log) * softplus(a + dt_bias) (float32), [T, H]; +// b and a are [T, H] in rows of ld. +int cs1_gdn_gates(const void* b, const void* a, int ld, const void* A_log, const void* dt_bias, + void* beta, float* g, int T, int H, void* stream); + +// Chunked gated delta rule, q and k L2-normalized inside, q scaled by `scale`. +// q, k [T, HK, 128], v [T, H, 128], g float [T, H], beta [T, H], o [T, H, 128]. +size_t cs1_gdn_workspace_floats(int T, int H); +int cs1_gdn_prefill(const void* q, const void* k, const void* v, const float* g, const void* beta, + void* o, float* workspace, int T, int H, int HK, float scale, void* stream); + +// Attention inputs: q and gate from qg [T, Hq, 2*Dh], k from kr [T, Hk, Dh], both in rows +// of ld; per-head zero-centred RMSNorm, then rotary embedding on the first 2*half dims +// using cos/sin [T, half] (bfloat16). Writes q [T, Hq, Dh], gate [T, Hq*Dh], k [T, Hk, Dh]. +int cs1_attn_prep(const void* qg, const void* kr, int ld, const void* qw, const void* kw, + const void* cos, const void* sin, void* q, void* gate, void* k, int T, int Hq, int Hk, + int Dh, int half, float eps, void* stream); + +// Causal attention with grouped KV heads, Dh = 256: q [T, Hq, Dh], k [T, Hk, Dh], v +// [T, Hk, Dh] in rows of ldv; out [T, Hq, Dh]. +int cs1_attention(const void* q, const void* k, const void* v, int ldv, void* out, int T, int Hq, + int Hk, int Dh, float scale, void* stream); +// The same in float32 on CUDA cores, two queries per warp: slow, kept for checks. +int cs1_attention_simple(const void* q, const void* k, const void* v, int ldv, void* out, int T, + int Hq, int Hk, int Dh, float scale, void* stream); + +// x = x * sigmoid(gate), n elements. +int cs1_sigmoid_gate(void* x, const void* gate, size_t n, void* stream); + +// out [T, I] = silu(gate) * up, from gate_up [T, 2*I] (gate first) in rows of ld. +int cs1_silu_mul(const void* gate_up, int ld, void* out, int T, int I, void* stream); + +// y [M, N] (rows of ldy) = x [M, K] * w [N, K]^T through cuBLASLt, float32 accumulation. +// cs1_gemm_tune picks the algorithm for one shape by timing, among the heuristic's +// shortlist or (exhaustive) a wider enumeration (see gemm.cu); it must not +// run during stream capture. cs1_gemm_tune_done frees the buffers tuning used. +// A tuned algorithm for one shape; `algo` holds a cublasLtMatmulAlgo_t. +typedef struct { + int32_t m, n, k, ldy; + uint64_t algo[8]; +} Cs1GemmPlan; + +void* cs1_gemm_create(size_t workspace_bytes); +void cs1_gemm_destroy(void* gemm); +int cs1_gemm_tune(void* gemm, const void* x, const void* w, void* y, int M, int N, int K, int ldy, + int exhaustive, void* stream); +void cs1_gemm_tune_done(void* gemm); +// Copy up to `cap` tuned plans to `out`; returns how many there are. +size_t cs1_gemm_export(void* gemm, Cs1GemmPlan* out, size_t cap); +// Use these plans (from cs1_gemm_export, possibly of an earlier run): all of them, or +// none if one fails cuBLASLt's check on this device or reduces split-K in place. +// The check does not tell whether a plan was tuned on this GPU and cuBLASLt version; +// the caller keeps that with the plans. +int cs1_gemm_import(void* gemm, const Cs1GemmPlan* plans, size_t n); +// cublasLtGetVersion(). +size_t cs1_gemm_version(void); +int cs1_gemm(void* gemm, const void* x, const void* w, void* y, int M, int N, int K, int ldy, + void* stream); + +#ifdef __cplusplus +} +#endif diff --git a/src/backends/cuda/qwen3_5/runtime.cu b/src/backends/cuda/qwen3_5/runtime.cu new file mode 100644 index 00000000..deaacc26 --- /dev/null +++ b/src/backends/cuda/qwen3_5/runtime.cu @@ -0,0 +1,69 @@ +// The CUDA runtime calls the Rust side needs, so that it loads one library +// (libqwen3_5_cuda.so) and never links CUDA itself. +#include + +#include + +#include "ops.h" + +extern "C" { + +uint32_t cs1_abi_version(void) { return CS1_ABI_VERSION; } + +const char* cs1_error_string(int code) { return cudaGetErrorString(static_cast(code)); } + +int cs1_set_device(int device) { return cudaSetDevice(device); } + +int cs1_device_info(char* name, size_t cap, int* compute_capability, int* sms) { + int device = 0; + cudaDeviceProp prop; + cudaError_t e = cudaGetDevice(&device); + if (e == cudaSuccess) e = cudaGetDeviceProperties(&prop, device); + if (e != cudaSuccess) return e; + if (cap > 0) std::snprintf(name, cap, "%s", prop.name); + *compute_capability = prop.major * 10 + prop.minor; + *sms = prop.multiProcessorCount; + return cudaSuccess; +} + +int cs1_malloc(void** ptr, size_t bytes) { return cudaMalloc(ptr, bytes); } + +int cs1_free(void* ptr) { return cudaFree(ptr); } + +int cs1_stream_create(void** stream) { + return cudaStreamCreateWithFlags(reinterpret_cast(stream), cudaStreamNonBlocking); +} + +int cs1_stream_sync(void* stream) { return cudaStreamSynchronize(static_cast(stream)); } + +int cs1_upload(void* dst, const void* src, size_t bytes, void* stream) { + const cudaStream_t st = static_cast(stream); + const cudaError_t e = cudaMemcpyAsync(dst, src, bytes, cudaMemcpyHostToDevice, st); + return e != cudaSuccess ? e : cudaStreamSynchronize(st); +} + +int cs1_download(void* dst, const void* src, size_t bytes, void* stream) { + const cudaStream_t st = static_cast(stream); + const cudaError_t e = cudaMemcpyAsync(dst, src, bytes, cudaMemcpyDeviceToHost, st); + return e != cudaSuccess ? e : cudaStreamSynchronize(st); +} + +int cs1_graph_begin(void* stream) { + return cudaStreamBeginCapture(static_cast(stream), cudaStreamCaptureModeThreadLocal); +} + +int cs1_graph_end(void* stream, void** exec) { + cudaGraph_t graph = nullptr; + cudaError_t e = cudaStreamEndCapture(static_cast(stream), &graph); + if (e == cudaSuccess) e = cudaGraphInstantiate(reinterpret_cast(exec), graph, 0); + if (graph) cudaGraphDestroy(graph); + return e; +} + +int cs1_graph_launch(void* exec, void* stream) { + return cudaGraphLaunch(static_cast(exec), static_cast(stream)); +} + +int cs1_graph_destroy(void* exec) { return cudaGraphExecDestroy(static_cast(exec)); } + +} // extern "C" diff --git a/src/models/cua_s1/README.md b/src/models/cua_s1/README.md index f0712d26..4212fae1 100644 --- a/src/models/cua_s1/README.md +++ b/src/models/cua_s1/README.md @@ -2,7 +2,7 @@ This directory owns Cua-S1 4B 0.2 ([#10](https://github.com/ThinkFlowLab/system1-omni/issues/10)): request mapping, prompt construction, adapter selection, execution, and the answer-letter readout. This page records the pinned upstream revisions, the inference contract an implementation must match, and how its outputs will be compared with the upstream reference. -Status: a reference worker for the `text` adapter is in [`text/`](text/). It loads the model directly through Hugging Face Transformers and PEFT and serves `/v1/systemone`; setup and checks are in [`recipe/cua_s1/text.md`](../../../recipe/cua_s1/text.md). A worker for the `multimodal` adapter is proposed in [#12](https://github.com/ThinkFlowLab/system1-omni/pull/12). +Status: a reference worker for the `text` adapter is in [`text/`](text/). It loads the model directly through Hugging Face Transformers and PEFT and serves `/v1/systemone`; setup and checks are in [`recipe/cua_s1/text.md`](../../../recipe/cua_s1/text.md). A native worker for the `text` adapter is in [`native/`](native/): Rust, with the Qwen3.5 forward pass on the CUDA kernels in [`src/backends/cuda/qwen3_5/`](../../backends/cuda/qwen3_5/); see [`recipe/cua_s1/native.md`](../../../recipe/cua_s1/native.md). A worker for the `multimodal` adapter is proposed in [#12](https://github.com/ThinkFlowLab/system1-omni/pull/12). | Path | Contents | | --- | --- | @@ -10,6 +10,7 @@ Status: a reference worker for the `text` adapter is in [`text/`](text/). It loa | `text/engine.py` | Model and adapter loading and the answer-letter readout. | | `text/server.py` | The HTTP worker (`GET /health`, `POST /v1/systemone`). | | `text/adapter.py` | Finds and checks the local `text` adapter and its downloaded revision. | +| `native/` | The native worker for the `text` adapter (Rust crate `omni-cua-s1-native`): the same request handling and answers as `text/`, the forward pass on `src/backends/cuda/qwen3_5/`. | | `tests/cua_s1/test_text_*.py` (repository root) | Tests that need neither weights nor a GPU, and tokenizer checks. The fixed input set is `tests/cua_s1/data/text_inputs.json`. | ## Pinned revisions @@ -105,7 +106,7 @@ The status is `422` when a well-formed request cannot be answered: | Worker prompt vs upstream `build_prompt` | Same inputs | Identical token ids | | Worker vs upstream `FourBModel` | Same GPU and reference environment, bfloat16, adapter not merged, full logits, one unpadded prompt per forward pass | Identical fp32 probabilities | | Through the frontend vs direct to the worker | Same worker | Identical status, content type and body bytes | -| Native engine vs fp32 worker (later) | Same GPU; the engine runs in bfloat16; the fp32 worker runs with TF32 disabled | (1) Over the whole input set, the largest per-option probability difference is at most twice the bfloat16 worker's largest difference from the fp32 worker, plus 0.01. (2) The top option matches wherever the fp32 worker's top-two margin is at least 0.05. | +| Native engine vs fp32 worker | Same GPU; the engine runs in bfloat16; the fp32 worker runs with TF32 disabled | (1) Over the whole input set, the largest per-option probability difference is at most twice the bfloat16 worker's largest difference from the fp32 worker, plus 0.01. (2) The top option matches wherever the fp32 worker's top-two margin is at least 0.05. | The bfloat16 worker's own difference from the fp32 worker is reported next to each native-engine result. diff --git a/src/models/cua_s1/native/Cargo.toml b/src/models/cua_s1/native/Cargo.toml new file mode 100644 index 00000000..35cf3a8e --- /dev/null +++ b/src/models/cua_s1/native/Cargo.toml @@ -0,0 +1,26 @@ +[package] +name = "omni-cua-s1-native" +version = "0.1.0" +edition = "2024" +publish = false +description = "Native /v1/systemone worker for Cua-S1 4B 0.2 (text adapter)" + +[[bin]] +name = "omni-cua-s1-native" +path = "src/main.rs" + +[dependencies] +anyhow = "1.0.100" +axum = "0.8.8" +clap = { version = "4.5.54", features = ["derive", "env"] } +half = "2.7.1" +http-body-util = "0.1.3" +# the CUDA kernels live in libqwen3_5_cuda.so, loaded at run time +libloading = "0.8" +memmap2 = "0.9.9" +safetensors = "0.8.0" +serde_json = { version = "1.0.149", features = ["preserve_order"] } +sha2 = "0.10.9" +# the onig regex backend, as in the Python tokenizers wheel +tokenizers = { version = "=0.22.2", default-features = false, features = ["onig"] } +tokio = { version = "1.49.0", features = ["macros", "net", "rt-multi-thread", "sync"] } diff --git a/src/models/cua_s1/native/README.md b/src/models/cua_s1/native/README.md new file mode 100644 index 00000000..eb247a5d --- /dev/null +++ b/src/models/cua_s1/native/README.md @@ -0,0 +1,48 @@ +# Cua-S1 4B 0.2 native text worker + +A `/v1/systemone` worker for the `text` adapter in Rust. It answers every request the way the reference worker in [`../text/`](../text/) does (same validation, error bodies, prompt token ids and answer format) and runs the Qwen3.5-4B forward pass on the CUDA kernels in [`src/backends/cuda/qwen3_5/`](../../../backends/cuda/qwen3_5/), with no Python and no PyTorch. Setup, launch and checks are in [`recipe/cua_s1/native.md`](../../../../recipe/cua_s1/native.md). + +| File | Contents | +| --- | --- | +| `src/pyjson.rs` | JSON parsing and output that behave like CPython 3.12's `json` module and `repr`, so request errors and answers match the Python worker byte for byte. | +| `src/contract.rs` | The request mapping, prompt text, confidence and answers of `../text/contract.py`. | +| `src/server.rs` | The HTTP worker (`GET /health`, `POST /v1/systemone`), with the reference worker's limits, status codes and error bodies. | +| `src/engine.rs` | Tokenization, the letter rows of the output projection, and scoring. | +| `src/model.rs` | The Qwen3.5 text model: weights, buffers, the layer loop, CUDA graphs and GEMM tuning. | +| `src/cuda.rs` | Loading `libqwen3_5_cuda.so` and the calls into it. | +| `tests/kernels.rs` | GPU checks of the attention and Gated DeltaNet kernels (ignored unless asked for; they need `CUA_S1_CUDA_LIB`). | +| `THIRD_PARTY_NOTICES.md` | The license of the prompt text and fixed values that `src/contract.rs` copies from trycua/cua. | + +## How a question is answered + +1. The request is parsed and mapped as in the reference worker, and each question's prompt is tokenized with the `tokenizer.json` exported next to the merged weights (the worker checks its SHA-256 against `cua_s1_export.json`). +2. One forward pass runs over the prompt: bfloat16 weights, with the `text` adapter merged into them by `recipe/cua_s1/export_text_merged.py`. The operations follow the Transformers implementation and round to bfloat16 where it does, except inside attention and the Gated DeltaNet prefill, which keep some intermediate results in bfloat16 as FlashAttention and flash-linear-attention do. +3. The final-norm hidden state at the last position is multiplied by the 26 letter rows of the output projection (float32 with float64 accumulation), and a softmax over the question's letters gives the option probabilities. + +The whole path from request to answer is in this crate. `src/backends/cuda/qwen3_5/` provides the operations: RMSNorm variants, the Gated DeltaNet convolution, gates and chunked prefill, rotary embedding and attention, and bfloat16 GEMMs through cuBLASLt. + +## CUDA library, graphs and GEMM plans + +The kernels are built into `libqwen3_5_cuda.so` by `src/backends/cuda/qwen3_5/build.sh` and loaded when the worker starts (`--cuda-lib`, by default next to the executable), so building the crate needs no CUDA toolkit and the workspace checks run anywhere. + +Prompts up to `--graph-max-tokens` (2048) run as a CUDA graph captured for their exact length on first use; the 128 most recently used lengths keep theirs. There is no padding, and a graph queues the same kernels with the same GEMM algorithms as the eager pass, so both give bitwise identical results. Longer prompts run eagerly. + +Projections that share an input run as one GEMM (`gate_proj` and `up_proj`; `in_proj_qkv`, `in_proj_z`, `in_proj_b` and `in_proj_a`; `q_proj`, `k_proj` and `v_proj`), with their weights stored back to back. GEMM algorithms are tuned at startup for a set of prompt lengths, by timing cuBLASLt's candidates. For prompts up to `--graph-max-tokens`, `--gemm-search` enumerates far more configurations (each algorithm with its tiles, stage counts, swizzles and several split-K factors), times each once and times the 12 fastest properly. `--gemm-plans` keeps the choices in a file, so that later starts reuse them and give the same results. The file records the GPU, the cuBLASLt version, the workspace size and `--graph-max-tokens`; a file that does not match is refused, and removing it tunes again. + +## Known differences from the reference worker + +Request bodies are handled the same way. The nesting limits copy what the Python worker does on CPython 3.12.13 (parsing fails for arrays nested more than 9,990 deep, and the check after parsing past 969 nested calls). In Python both come from recursion limits, so they move with the interpreter's call stack, and some values need more of it: close to the limit, `NaN` inside 9,985 arrays is "nested too deeply" in Python and "NaN is not valid JSON" here. The status is 400 either way. + +Outside the request body the HTTP stacks differ (Starlette and h11 in Python, axum and hyper here): + +- A trailing slash (`/v1/systemone/`) gets a 307 redirect from Starlette and a 404 here; a percent-encoded path is decoded by uvicorn and not here. +- FastAPI also serves `/docs`, `/redoc` and `/openapi.json`; they are not served here. +- `HEAD /health` returns 200 here and 405 from FastAPI. +- The HTTP parsers reject different malformed requests (control bytes in header values, a missing `Host`, a `Content-Length` too large to parse), and those rejections are not JSON. A body that fails partway through reading gets a JSON 400 here. +- `GET /health` also reports `"mode": "native"`, and its `dtype` is always `bfloat16`. + +The probabilities are not bitwise identical to the reference worker's: the adapter is merged, the kernels differ, and the GEMM algorithms depend on the prompt length and the GPU. They are checked against the float32 worker under the tolerance in [`../README.md`](../README.md#validation). + +## Tests + +`cargo test -p omni-cua-s1-native` runs the request-handling tests. Three more are ignored unless asked for with `-- --ignored`: a check of float formatting against Python, on a file that `tests/make_float_vectors.py` writes (`CUA_S1_FLOAT_VECTORS`), and the kernel checks in `tests/kernels.rs`, which need a GPU and `CUA_S1_CUDA_LIB` pointing to a built `libqwen3_5_cuda.so`. diff --git a/src/models/cua_s1/native/THIRD_PARTY_NOTICES.md b/src/models/cua_s1/native/THIRD_PARTY_NOTICES.md new file mode 100644 index 00000000..32a974d6 --- /dev/null +++ b/src/models/cua_s1/native/THIRD_PARTY_NOTICES.md @@ -0,0 +1,23 @@ +The system message, prompt layout and fixed values in `src/contract.rs` come from [trycua/cua](https://github.com/trycua/cua) at `0e75660ce4c2edda519e0c795fa3ad98abf4e76f` under the following license. + +MIT License + +Copyright (c) 2025 Cua AI, Inc. + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/src/models/cua_s1/native/src/contract.rs b/src/models/cua_s1/native/src/contract.rs new file mode 100644 index 00000000..4336689a --- /dev/null +++ b/src/models/cua_s1/native/src/contract.rs @@ -0,0 +1,439 @@ +//! Request mapping, prompt construction and answers for Cua-S1 4B 0.2, ported from +//! `src/models/cua_s1/text/contract.py` so the two workers answer alike, down to the +//! error messages. + +use std::fmt::Write as _; + +use crate::pyjson::{self, PyStr, Value, dumps, float_repr, repr, repr_str, write_json_str}; + +pub const MODEL_NAME: &str = "cua-s1-4b-0.2"; +pub const ADAPTER_REPO: &str = "cua-ai/cua-s1-4b-0.2"; +pub const ADAPTER_REVISION: &str = "16818868b0cc7813808aae4e87b417657046ab79"; +pub const BASE_REPO: &str = "Qwen/Qwen3.5-4B"; +pub const BASE_REVISION: &str = "851bf6e806efd8d0a36b00ddf55e13ccb7b8cd0a"; + +pub const LETTERS: &str = "ABCDEFGHIJKLMNOPQRSTUVWXYZ"; +pub const MAX_OPTIONS: usize = 26; + +// The system message, the user message layout and the fixed values below are copied +// from trycua/cua at 0e75660ce4c2edda519e0c795fa3ad98abf4e76f: +// `libs/cua-s1/python/src/cua_s1/four_b.py` (SYSTEM_PROMPT, build_prompt, +// _describe_option) and `libs/cua-driver/examples/jev-use/python/decision_models.py` +// (S1DecisionModel.score). MIT License, Copyright (c) 2025 Cua AI, Inc.; the full +// notice is in THIRD_PARTY_NOTICES.md. +pub const SYSTEM_PROMPT: &str = "You are a one-pass computer-use decision model. You are shown the \ +current state of a screen and a fixed, closed list of candidate \ +(element, action) options, each given a single letter. Choose exactly \ +one option: the single best next action to take. Answer with ONLY that \ +option's letter -- no words, no punctuation, no explanation."; +pub const APP: &str = "Cua Driver"; +pub const TASK_FAMILY: &str = "closed-candidate decision"; +pub const ROLE: &str = "Decision"; +pub const ACTION: &str = "select"; + +/// A request the worker rejects, with the HTTP status to return. +#[derive(Debug, Clone, PartialEq)] +pub struct RequestError { + pub status: u16, + pub message: String, +} + +impl RequestError { + pub fn new(status: u16, message: impl Into) -> Self { + Self { + status, + message: message.into(), + } + } + + fn unprocessable(message: impl Into) -> Self { + Self::new(422, message) + } +} + +impl From for RequestError { + fn from(e: pyjson::JsonError) -> Self { + RequestError::new(400, e.message()) + } +} + +/// One `choice` question mapped onto the prompt fields. +#[derive(Debug, Clone, PartialEq)] +pub struct Question { + pub name: String, + pub goal: String, + pub keys: Vec, + pub labels: Vec, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct Request { + pub state: String, + pub questions: Vec, +} + +pub fn parse_body(raw: &[u8]) -> Result, RequestError> { + Ok(pyjson::parse(raw)?) +} + +fn text(s: &PyStr) -> &str { + s.as_str().expect("parse rejects lone surrogates") +} + +/// `state` or `instructions` as prompt text: a string as is, anything else as +/// `json.dumps(value, ensure_ascii=False)`. +fn as_text(value: &Value) -> String { + match value { + Value::Str(s) => text(s).to_string(), + other => dumps(other, false), + } +} + +/// An option label escaped the way upstream's chooser does: +/// `json.dumps(value, ensure_ascii=False)[1:-1]`. +fn escape_label(value: &str) -> String { + let mut out = String::new(); + write_json_str(value, &mut out); + out[1..out.len() - 1].to_string() +} + +fn check_json_value(value: &Value, place: &str, allow_null: bool) -> Result<(), RequestError> { + match value { + Value::Null if !allow_null => Err(RequestError::unprocessable(format!( + "{place} must not be null" + ))), + Value::Bool(_) | Value::Int(_) | Value::Float(_) => Err(RequestError::unprocessable( + format!("{place} must be a string, an object or an array"), + )), + _ => Ok(()), + } +} + +fn get<'a>(pairs: &'a [(PyStr, Value)], key: &str) -> Option<&'a Value> { + pairs.iter().find(|(k, _)| k == key).map(|(_, v)| v) +} + +/// Validate a `/v1/systemone` body and map it onto prompt fields. +pub fn map_request(body: &[(PyStr, Value)], max_questions: usize) -> Result { + match get(body, "model") { + Some(Value::Str(s)) if s == MODEL_NAME => {} + _ => { + return Err(RequestError::unprocessable(format!( + "'model' must be {}", + repr_str(&PyStr::new(MODEL_NAME)) + ))); + } + } + + let state_value = + get(body, "state").ok_or_else(|| RequestError::unprocessable("'state' is required"))?; + check_json_value(state_value, "'state'", false)?; + let empty = match state_value { + Value::Str(s) => s == "", + Value::Object(p) => p.is_empty(), + Value::Array(a) => a.is_empty(), + _ => false, + }; + if empty { + return Err(RequestError::unprocessable("'state' must not be empty")); + } + let state = as_text(state_value); + + let questions = match get(body, "questions") { + Some(Value::Object(q)) if !q.is_empty() => q, + _ => { + return Err(RequestError::unprocessable( + "'questions' must be a non-empty object", + )); + } + }; + if questions.len() > max_questions { + return Err(RequestError::new( + 413, + format!("too many questions ({} > {max_questions})", questions.len()), + )); + } + + // Every question type is checked before the per-question checks, so a `score` + // or `noul` question anywhere rejects the whole request with that reason. + for (name, question) in questions { + let Value::Object(q) = question else { + return Err(RequestError::unprocessable(format!( + "question {} must be an object", + repr_str(name) + ))); + }; + let missing = Value::Null; + let kind = get(q, "type").unwrap_or(&missing); + match kind { + Value::Str(s) if s == "score" || s == "noul" => { + return Err(RequestError::unprocessable(format!( + "question {}: type {} is not supported; Cua-S1 4B 0.2 answers 'choice' questions only", + repr_str(name), + repr(kind) + ))); + } + Value::Str(s) if s == "choice" => {} + _ => { + return Err(RequestError::unprocessable(format!( + "question {}: unknown type {}", + repr_str(name), + repr(kind) + ))); + } + } + } + + let mut mapped = Vec::with_capacity(questions.len()); + for (name, question) in questions { + let Value::Object(q) = question else { + unreachable!("checked above") + }; + let place = format!("question {}", repr_str(name)); + let instructions = get(q, "instructions").ok_or_else(|| { + RequestError::unprocessable(format!("{place}: 'instructions' is required")) + })?; + check_json_value(instructions, &format!("{place}: 'instructions'"), true)?; + let goal = match instructions { + Value::Null => String::new(), + other => as_text(other), + }; + + let criteria = match get(q, "criteria") { + Some(Value::Object(c)) => c, + _ => { + return Err(RequestError::unprocessable(format!( + "{place}: 'criteria' must be an object" + ))); + } + }; + if criteria.is_empty() { + return Err(RequestError::unprocessable(format!( + "{place}: 'criteria' must have at least one option" + ))); + } + if criteria.len() > MAX_OPTIONS { + return Err(RequestError::unprocessable(format!( + "{place}: {} options; at most {MAX_OPTIONS} are supported", + criteria.len() + ))); + } + let mut keys = Vec::with_capacity(criteria.len()); + let mut labels = Vec::with_capacity(criteria.len()); + for (key, value) in criteria { + check_json_value(value, &format!("{place}: option {}", repr_str(key)), true)?; + let label = match value { + Value::Null => text(key).to_string(), + other => as_text(other), + }; + keys.push(text(key).to_string()); + labels.push(escape_label(&label)); + } + mapped.push(Question { + name: text(name).to_string(), + goal, + keys, + labels, + }); + } + Ok(Request { + state, + questions: mapped, + }) +} + +/// The user message for one question, matching upstream `build_prompt` (text). +pub fn user_message(state: &str, question: &Question) -> String { + let mut user = String::new(); + if !question.goal.is_empty() { + write!(user, "Goal: {}\n\n", question.goal).unwrap(); + } + write!(user, "App: {APP}\nTask family: {TASK_FAMILY}\n\n").unwrap(); + write!(user, "Accessibility tree:\n{state}\n\n").unwrap(); + user.push_str("Options:\n"); + for (i, (letter, label)) in LETTERS.chars().zip(&question.labels).enumerate() { + if i > 0 { + user.push('\n'); + } + write!(user, "{letter}. {ROLE} \"{label}\" -> {ACTION}").unwrap(); + } + user.push_str("\n\nAnswer with a single letter."); + user +} + +/// The prompt text the Qwen3.5 chat template renders for the system and user messages +/// with `add_generation_prompt=True` (thinking left on). The template trims message +/// content, which changes nothing here: the system prompt is fixed, and the user +/// message starts with "Goal: " or "App: " and ends with "letter.". +pub fn chat_text(state: &str, question: &Question) -> String { + format!( + "<|im_start|>system\n{SYSTEM_PROMPT}<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n\n", + user_message(state, question) + ) +} + +/// CPython 3.12's `sum()` over floats (Neumaier compensated summation). +fn py_sum(values: impl IntoIterator) -> f64 { + let (mut total, mut c) = (0.0f64, 0.0f64); + for x in values { + let t = total + x; + if total.abs() >= x.abs() { + c += (total - t) + x; + } else { + c += (x - t) + total; + } + total = t; + } + if c != 0.0 && c.is_finite() { + total += c; + } + total +} + +/// Normalized entropy, `1 - H(p) / ln(n)`, as the LAYA worker reports it. +pub fn confidence(probabilities: &[f64]) -> f64 { + let n = probabilities.len(); + if n < 2 { + return 1.0; + } + let entropy = -py_sum(probabilities.iter().map(|&p| p * p.clamp(1e-12, 1.0).ln())); + (1.0 - entropy / (n as f64).ln()).clamp(0.0, 1.0) +} + +/// One choice answer, already serialized the way the Python worker's response is +/// (`json.dumps(..., ensure_ascii=False, separators=(",", ":"))`). Ties go to the +/// earliest option. +pub fn answer_json(question: &Question, probabilities: &[f32]) -> Result { + let p: Vec = probabilities.iter().map(|&x| x as f64).collect(); + if p.len() != question.keys.len() || !p.iter().all(|x| x.is_finite() && (0.0..=1.0).contains(x)) + { + return Err(format!("model returned invalid probabilities: {p:?}")); + } + let total = py_sum(p.iter().copied()); + // math.isclose(total, 1.0, abs_tol=1e-5) + if (total - 1.0).abs() > f64::max(1e-9 * total.abs().max(1.0), 1e-5) { + return Err(format!("model probabilities do not sum to one: {p:?}")); + } + let mut best = 0; + for (i, &x) in p.iter().enumerate() { + if x > p[best] { + best = i; + } + } + let mut out = String::from("{\"type\":\"choice\",\"choice\":"); + write_json_str(&question.keys[best], &mut out); + out.push_str(",\"probabilities\":{"); + for (i, (key, x)) in question.keys.iter().zip(&p).enumerate() { + if i > 0 { + out.push(','); + } + write_json_str(key, &mut out); + out.push(':'); + out.push_str(&float_repr(*x)); + } + out.push_str("},\"confidence\":"); + out.push_str(&float_repr(confidence(&p))); + out.push('}'); + Ok(out) +} + +pub fn model_identity(revision: &str) -> String { + format!("{ADAPTER_REPO}@{revision}:text") +} + +/// `{"detail": message}` as the Python worker serializes it. +pub fn detail_json(message: &str) -> String { + let mut out = String::from("{\"detail\":"); + write_json_str(message, &mut out); + out.push('}'); + out +} + +pub const WARMUP_BODY: &str = r#"{"model": "cua-s1-4b-0.2", "state": "Dialog: 'Update installed.' Button: OK", "questions": {"warmup": {"type": "choice", "instructions": "Close the dialog.", "criteria": {"ok": "Click OK", "wait": "Wait"}}}}"#; + +#[cfg(test)] +mod tests { + use super::*; + + fn map(body: &str) -> Result { + map_request(&parse_body(body.as_bytes())?, 64) + } + + fn detail(body: &str) -> (u16, String) { + let e = map(body).unwrap_err(); + (e.status, e.message) + } + + const OK: &str = r#"{"model": "cua-s1-4b-0.2", "state": "S", "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "A", "b": null}}}}"#; + + #[test] + fn maps_a_request() { + let r = map(OK).unwrap(); + assert_eq!(r.state, "S"); + assert_eq!(r.questions[0].keys, ["a", "b"]); + assert_eq!(r.questions[0].labels, ["A", "b"]); + let text = chat_text(&r.state, &r.questions[0]); + assert!(text.ends_with( + "Options:\nA. Decision \"A\" -> select\nB. Decision \"b\" -> select\n\nAnswer with a single letter.<|im_end|>\n<|im_start|>assistant\n\n" + )); + assert!(text.contains("<|im_start|>user\nGoal: go\n\nApp: Cua Driver\n")); + } + + #[test] + fn errors_match_python() { + let cases: &[(&str, u16, &str)] = &[ + (r#"{"state": "S"}"#, 422, "'model' must be 'cua-s1-4b-0.2'"), + (r#"{"model": "cua-s1-4b-0.2"}"#, 422, "'state' is required"), + ( + r#"{"model": "cua-s1-4b-0.2", "state": null}"#, + 422, + "'state' must not be null", + ), + ( + r#"{"model": "cua-s1-4b-0.2", "state": 1.5}"#, + 422, + "'state' must be a string, an object or an array", + ), + ( + r#"{"model": "cua-s1-4b-0.2", "state": {}}"#, + 422, + "'state' must not be empty", + ), + ( + r#"{"model": "cua-s1-4b-0.2", "state": "S", "questions": []}"#, + 422, + "'questions' must be a non-empty object", + ), + ( + r#"{"model": "cua-s1-4b-0.2", "state": "S", "questions": {"q": {"type": "choice"}, "r": {"type": "noul"}}}"#, + 422, + "question 'r': type 'noul' is not supported; Cua-S1 4B 0.2 answers 'choice' questions only", + ), + ( + r#"{"model": "cua-s1-4b-0.2", "state": "S", "questions": {"it's": {"type": [1, {"a": null}]}}}"#, + 422, + "question \"it's\": unknown type [1, {'a': None}]", + ), + ( + r#"{"model": "cua-s1-4b-0.2", "state": "S", "questions": {"q": {"type": "choice"}}}"#, + 422, + "question 'q': 'instructions' is required", + ), + ( + r#"{"model": "cua-s1-4b-0.2", "state": "S", "questions": {"q": {"type": "choice", "instructions": null, "criteria": {"a": true}}}}"#, + 422, + "question 'q': option 'a' must be a string, an object or an array", + ), + ]; + for (body, status, message) in cases { + assert_eq!(detail(body), (*status, message.to_string()), "{body}"); + } + } + + #[test] + fn confidence_matches_python() { + // values from the Python worker + let p = [0.00247262348420918f64, 0.9975274205207825]; + assert_eq!(float_repr(confidence(&p)), "0.9750249548256825"); + } +} diff --git a/src/models/cua_s1/native/src/cuda.rs b/src/models/cua_s1/native/src/cuda.rs new file mode 100644 index 00000000..d6e8c47e --- /dev/null +++ b/src/models/cua_s1/native/src/cuda.rs @@ -0,0 +1,335 @@ +//! The CUDA side, loaded at run time from `libqwen3_5_cuda.so` (built by +//! `src/backends/cuda/qwen3_5/build.sh`): the Qwen3.5 operations and the few CUDA +//! runtime calls the model needs. Building this crate needs no CUDA toolkit. + +use std::ffi::{CStr, c_char, c_int, c_void}; +use std::path::{Path, PathBuf}; +use std::sync::OnceLock; + +use anyhow::{Context, Result, bail, ensure}; + +/// `CS1_ABI_VERSION` in ops.h. +const ABI_VERSION: u32 = 1; +pub const LIBRARY: &str = "libqwen3_5_cuda.so"; + +/// A `cudaStream_t`. +#[repr(transparent)] +#[derive(Clone, Copy)] +pub struct Stream(*mut c_void); + +// SAFETY: a stream handle may be used from any thread; the model queues work on it +// from one thread at a time. +unsafe impl Send for Stream {} + +/// A tuned GEMM algorithm for one shape (`Cs1GemmPlan` in ops.h). +#[repr(C)] +#[derive(Clone, Copy, Default)] +pub struct GemmPlan { + pub m: i32, + pub n: i32, + pub k: i32, + pub ldy: i32, + pub algo: [u64; 8], +} + +macro_rules! api { + ($($name:ident($($arg:ident: $ty:ty),* $(,)?) $(-> $ret:ty)?;)*) => { + /// The functions of the library, as declared in ops.h. + pub struct Api { + _lib: libloading::Library, + $(pub $name: unsafe extern "C" fn($($ty),*) $(-> $ret)?,)* + } + + impl Api { + fn resolve(lib: libloading::Library) -> Result { + $( + // SAFETY: the signature is the one ops.h declares for this symbol. + let $name = unsafe { + lib.get:: $ret)?>( + concat!(stringify!($name), "\0").as_bytes(), + ) + .map(|f| *f) + } + .with_context(|| format!("{} has no {}", LIBRARY, stringify!($name)))?; + )* + Ok(Self { _lib: lib, $($name,)* }) + } + } + }; +} + +api! { + cs1_abi_version() -> u32; + cs1_error_string(code: c_int) -> *const c_char; + cs1_set_device(device: c_int) -> c_int; + cs1_device_info(name: *mut c_char, cap: usize, compute_capability: *mut c_int, sms: *mut c_int) -> c_int; + cs1_malloc(ptr: *mut *mut c_void, bytes: usize) -> c_int; + cs1_free(ptr: *mut c_void) -> c_int; + cs1_stream_create(stream: *mut Stream) -> c_int; + cs1_stream_sync(stream: Stream) -> c_int; + cs1_upload(dst: *mut c_void, src: *const c_void, bytes: usize, stream: Stream) -> c_int; + cs1_download(dst: *mut c_void, src: *const c_void, bytes: usize, stream: Stream) -> c_int; + cs1_graph_begin(stream: Stream) -> c_int; + cs1_graph_end(stream: Stream, exec: *mut *mut c_void) -> c_int; + cs1_graph_launch(exec: *mut c_void, stream: Stream) -> c_int; + cs1_graph_destroy(exec: *mut c_void) -> c_int; + cs1_embed(ids: *const i32, table: *const c_void, out: *mut c_void, t: c_int, d: c_int, stream: Stream) -> c_int; + cs1_rms_norm( + x: *const c_void, w: *const c_void, out: *mut c_void, rows: c_int, d: c_int, eps: f32, stream: Stream, + ) -> c_int; + cs1_add_rms_norm( + residual: *mut c_void, delta: *const c_void, w: *const c_void, out: *mut c_void, rows: c_int, d: c_int, + eps: f32, stream: Stream, + ) -> c_int; + cs1_gated_rms_norm( + x: *const c_void, z: *const c_void, ldz: c_int, w: *const c_void, out: *mut c_void, t: c_int, h: c_int, + d: c_int, eps: f32, stream: Stream, + ) -> c_int; + cs1_gdn_conv( + qkv: *const c_void, ld: c_int, w: *const c_void, q: *mut c_void, k: *mut c_void, v: *mut c_void, t: c_int, + key_dim: c_int, value_dim: c_int, stream: Stream, + ) -> c_int; + cs1_gdn_gates( + b: *const c_void, a: *const c_void, ld: c_int, a_log: *const c_void, dt_bias: *const c_void, + beta: *mut c_void, g: *mut f32, t: c_int, h: c_int, stream: Stream, + ) -> c_int; + cs1_gdn_workspace_floats(t: c_int, h: c_int) -> usize; + cs1_gdn_prefill( + q: *const c_void, k: *const c_void, v: *const c_void, g: *const f32, beta: *const c_void, o: *mut c_void, + workspace: *mut f32, t: c_int, h: c_int, hk: c_int, scale: f32, stream: Stream, + ) -> c_int; + cs1_attn_prep( + qg: *const c_void, kr: *const c_void, ld: c_int, qw: *const c_void, kw: *const c_void, cos: *const c_void, + sin: *const c_void, q: *mut c_void, gate: *mut c_void, k: *mut c_void, t: c_int, hq: c_int, hk: c_int, + dh: c_int, half: c_int, eps: f32, stream: Stream, + ) -> c_int; + cs1_attention( + q: *const c_void, k: *const c_void, v: *const c_void, ldv: c_int, out: *mut c_void, t: c_int, hq: c_int, + hk: c_int, dh: c_int, scale: f32, stream: Stream, + ) -> c_int; + cs1_attention_simple( + q: *const c_void, k: *const c_void, v: *const c_void, ldv: c_int, out: *mut c_void, t: c_int, hq: c_int, + hk: c_int, dh: c_int, scale: f32, stream: Stream, + ) -> c_int; + cs1_sigmoid_gate(x: *mut c_void, gate: *const c_void, n: usize, stream: Stream) -> c_int; + cs1_silu_mul(gate_up: *const c_void, ld: c_int, out: *mut c_void, t: c_int, i: c_int, stream: Stream) -> c_int; + cs1_gemm_create(workspace_bytes: usize) -> *mut c_void; + cs1_gemm_destroy(gemm: *mut c_void); + cs1_gemm_tune( + gemm: *mut c_void, x: *const c_void, w: *const c_void, y: *mut c_void, m: c_int, n: c_int, k: c_int, + ldy: c_int, exhaustive: c_int, stream: Stream, + ) -> c_int; + cs1_gemm_tune_done(gemm: *mut c_void); + cs1_gemm_export(gemm: *mut c_void, out: *mut GemmPlan, cap: usize) -> usize; + cs1_gemm_import(gemm: *mut c_void, plans: *const GemmPlan, n: usize) -> c_int; + cs1_gemm_version() -> usize; + cs1_gemm( + gemm: *mut c_void, x: *const c_void, w: *const c_void, y: *mut c_void, m: c_int, n: c_int, k: c_int, + ldy: c_int, stream: Stream, + ) -> c_int; +} + +static API: OnceLock = OnceLock::new(); + +/// The library next to the running executable. +pub fn default_library() -> Result { + let exe = std::env::current_exe()?; + Ok(exe + .parent() + .context("the executable has no directory")? + .join(LIBRARY)) +} + +/// Load the library (once per process) and check its ABI version. +pub fn load(path: &Path) -> Result<&'static Api> { + if let Some(api) = API.get() { + return Ok(api); + } + // SAFETY: loading runs the library's initializers; it is the library these + // sources build. + let lib = unsafe { libloading::Library::new(path) }.with_context(|| { + format!( + "loading {} (build it with src/backends/cuda/qwen3_5/build.sh)", + path.display() + ) + })?; + let api = Api::resolve(lib)?; + // SAFETY: takes no arguments. + let abi = unsafe { (api.cs1_abi_version)() }; + ensure!( + abi == ABI_VERSION, + "{} has ABI version {abi}, this build needs {ABI_VERSION}; rebuild it", + path.display() + ); + Ok(API.get_or_init(|| api)) +} + +/// Name, compute capability (major * 10 + minor) and SM count of the current device. +pub fn device_info() -> Result<(String, i32, i32)> { + let mut name = [0 as c_char; 256]; + let (mut cc, mut sms) = (0, 0); + // SAFETY: `name` has room for 256 bytes, of which the library writes a terminated + // string; the two integers are valid for writes. + check( + unsafe { (api().cs1_device_info)(name.as_mut_ptr(), name.len(), &mut cc, &mut sms) }, + "reading the device properties", + )?; + // SAFETY: terminated by the library. + let name = unsafe { CStr::from_ptr(name.as_ptr()) }; + Ok((name.to_string_lossy().into_owned(), cc, sms)) +} + +/// The loaded library; `load` must have succeeded before. +pub fn api() -> &'static Api { + API.get().expect("the CUDA library is not loaded") +} + +/// Turn a return code of the library into an error; codes from 1000 up are cuBLAS +/// statuses (see gemm.cu). +pub fn check(code: c_int, what: &str) -> Result<()> { + if code == 0 { + return Ok(()); + } + if code >= 1000 { + bail!("{what}: cuBLAS status {}", code - 1000); + } + // SAFETY: cudaGetErrorString returns a static string for any code. + let msg = unsafe { CStr::from_ptr((api().cs1_error_string)(code)) }; + bail!("{what}: {} ({code})", msg.to_string_lossy()); +} + +pub fn set_device(device: i32) -> Result<()> { + // SAFETY: plain runtime call. + check(unsafe { (api().cs1_set_device)(device) }, "cudaSetDevice") +} + +/// A device allocation, freed on drop. +pub struct DeviceBuffer { + ptr: *mut c_void, + bytes: usize, +} + +// SAFETY: the pointer is a device address; access is serialized by the owner. +unsafe impl Send for DeviceBuffer {} +unsafe impl Sync for DeviceBuffer {} + +impl DeviceBuffer { + pub fn new(bytes: usize) -> Result { + let mut ptr = std::ptr::null_mut(); + if bytes > 0 { + // SAFETY: `ptr` is a valid out-pointer. + check(unsafe { (api().cs1_malloc)(&mut ptr, bytes) }, "cudaMalloc")?; + } + Ok(Self { ptr, bytes }) + } + + pub fn bytes(&self) -> usize { + self.bytes + } + + /// The device address `offset` bytes into the buffer. + pub fn at(&self, offset: usize) -> *mut c_void { + debug_assert!(offset <= self.bytes); + self.ptr.wrapping_byte_add(offset) + } +} + +impl Drop for DeviceBuffer { + fn drop(&mut self) { + if !self.ptr.is_null() { + // SAFETY: allocated by cs1_malloc and not freed before. + unsafe { (api().cs1_free)(self.ptr) }; + } + } +} + +pub fn new_stream() -> Result { + let mut stream = Stream(std::ptr::null_mut()); + // SAFETY: `stream` is a valid out-pointer. + check( + unsafe { (api().cs1_stream_create)(&mut stream) }, + "cudaStreamCreateWithFlags", + )?; + Ok(stream) +} + +pub fn synchronize(stream: Stream) -> Result<()> { + // SAFETY: a Stream only comes from new_stream. + check( + unsafe { (api().cs1_stream_sync)(stream) }, + "cudaStreamSynchronize", + ) +} + +/// Copy host bytes to `dst` and wait for the copy. +/// +/// # Safety +/// `dst` must be a device allocation with room for `src.len()` bytes. +pub unsafe fn upload(dst: *mut c_void, src: &[u8], stream: Stream) -> Result<()> { + // SAFETY: see above; the library waits for the copy before returning. + check( + unsafe { (api().cs1_upload)(dst, src.as_ptr().cast(), src.len(), stream) }, + "copy to device", + ) +} + +/// Copy `dst.len()` bytes from `src` to the host, after the work queued before it. +/// +/// # Safety +/// `src` must be a device allocation holding at least `dst.len()` bytes. +pub unsafe fn download(dst: &mut [u8], src: *const c_void, stream: Stream) -> Result<()> { + // SAFETY: see above; the library waits for the copy before returning. + check( + unsafe { (api().cs1_download)(dst.as_mut_ptr().cast(), src, dst.len(), stream) }, + "copy to host", + ) +} + +/// An instantiated CUDA graph, destroyed on drop. +pub struct Graph { + exec: *mut c_void, +} + +// SAFETY: the executable graph is only launched by its owner, one launch at a time. +unsafe impl Send for Graph {} + +impl Graph { + /// Capture the work `record` queues on `stream` (nothing runs) and instantiate it. + pub fn capture(stream: Stream, record: impl FnOnce() -> Result<()>) -> Result { + // SAFETY: plain runtime calls on a stream from new_stream; the capture is + // always ended, also when `record` fails. + unsafe { + check((api().cs1_graph_begin)(stream), "cudaStreamBeginCapture")?; + let recorded = record(); + let mut exec = std::ptr::null_mut(); + let ended = check( + (api().cs1_graph_end)(stream, &mut exec), + "capturing a CUDA graph", + ); + match recorded.and(ended) { + Ok(()) => Ok(Graph { exec }), + Err(e) => { + if !exec.is_null() { + (api().cs1_graph_destroy)(exec); + } + Err(e) + } + } + } + } + + pub fn launch(&self, stream: Stream) -> Result<()> { + // SAFETY: an instantiated graph whose buffers outlive it (see Model). + check( + unsafe { (api().cs1_graph_launch)(self.exec, stream) }, + "cudaGraphLaunch", + ) + } +} + +impl Drop for Graph { + fn drop(&mut self) { + // SAFETY: instantiated by capture and not destroyed before. + unsafe { (api().cs1_graph_destroy)(self.exec) }; + } +} diff --git a/src/models/cua_s1/native/src/engine.rs b/src/models/cua_s1/native/src/engine.rs new file mode 100644 index 00000000..600b0d4b --- /dev/null +++ b/src/models/cua_s1/native/src/engine.rs @@ -0,0 +1,278 @@ +//! The model side: prompt tokenization, and one prefill-only forward pass per +//! question through the native Qwen3.5 model, scored with the 26 letter rows of +//! the output projection. + +use std::path::Path; +use std::sync::{Arc, Mutex, MutexGuard}; +use std::time::Instant; + +use anyhow::{Context, Result, bail, ensure}; +use serde_json::Value as Json; +use sha2::Digest; +use tokenizers::Tokenizer; + +use crate::contract::{self, LETTERS, Question}; +use crate::model::Model; +pub use crate::model::{Mode, Options}; + +/// What `cua_s1_export.json` records about a merged checkpoint. +#[derive(Debug, Clone)] +pub struct Provenance { + pub base_revision: String, + pub adapter_revision: String, +} + +/// Read and check the export record written next to the merged weights. +pub fn provenance(dir: &Path) -> Result { + let path = dir.join("cua_s1_export.json"); + let info: Json = serde_json::from_str(&std::fs::read_to_string(&path).with_context(|| { + format!( + "{} is missing; export the merged checkpoint first", + path.display() + ) + })?)?; + ensure!( + info["format"] == "cua-s1-text-merged/1", + "{}: unknown format {}", + path.display(), + info["format"] + ); + ensure!( + info["base"]["repo"] == contract::BASE_REPO, + "base is not {}", + contract::BASE_REPO + ); + ensure!( + info["adapter"]["repo"] == contract::ADAPTER_REPO && info["adapter"]["subfolder"] == "text", + "adapter is not the `text` adapter of {}", + contract::ADAPTER_REPO + ); + // Transformers 5.17 tokenizes Qwen3.5 with the rule saved in this file, not the + // one in the base repo's tokenizer.json, so the file must be the exported one. + let want = info["tokenizer"]["sha256"] + .as_str() + .context("cua_s1_export.json does not record the tokenizer's sha256")?; + let got = format!( + "{:x}", + sha2::Sha256::digest(std::fs::read(dir.join("tokenizer.json"))?) + ); + ensure!( + got == want, + "tokenizer.json (sha256 {got}) is not the one exported with the weights ({want})" + ); + let rev = |v: &Json| v.as_str().map(str::to_string).context("revision missing"); + Ok(Provenance { + base_revision: rev(&info["base"]["revision"])?, + adapter_revision: rev(&info["adapter"]["revision"])?, + }) +} + +/// Chat text and token ids for a question; needs only `tokenizer.json`. +pub struct Prompter { + tokenizer: Tokenizer, + pub letter_ids: Vec, +} + +impl Prompter { + /// Checks `tokenizer.json` against `cua_s1_export.json` first (see `provenance`). + pub fn load(dir: &Path) -> Result { + provenance(dir)?; + let tokenizer = Tokenizer::from_file(dir.join("tokenizer.json")) + .map_err(|e| anyhow::anyhow!("tokenizer.json: {e}"))?; + let mut letter_ids = Vec::with_capacity(LETTERS.len()); + for letter in LETTERS.chars() { + let enc = tokenizer + .encode(letter.to_string(), false) + .map_err(|e| anyhow::anyhow!(e))?; + ensure!( + enc.get_ids().len() == 1, + "letter {letter} is not a single token" + ); + letter_ids.push(enc.get_ids()[0]); + } + Ok(Self { + tokenizer, + letter_ids, + }) + } + + pub fn encode(&self, state: &str, question: &Question) -> Result> { + let text = contract::chat_text(state, question); + let enc = self + .tokenizer + .encode(text, false) + .map_err(|e| anyhow::anyhow!(e))?; + Ok(enc.get_ids().to_vec()) + } +} + +/// Where the output projection can live in a Qwen3.5 text checkpoint; with tied +/// weights (as in Qwen3.5-4B) only the embedding is stored. +const HEAD_NAMES: &[&str] = &[ + "lm_head.weight", + "language_model.lm_head.weight", + "model.embed_tokens.weight", + "model.language_model.embed_tokens.weight", + "language_model.model.embed_tokens.weight", +]; + +/// The letter rows of the output projection, as float32, read straight from the +/// safetensors files. +fn letter_rows(dir: &Path, letter_ids: &[u32]) -> Result<(Vec, usize)> { + let index_path = dir.join("model.safetensors.index.json"); + let (file, name) = if index_path.exists() { + let index: Json = serde_json::from_str(&std::fs::read_to_string(&index_path)?)?; + let map = &index["weight_map"]; + let name = HEAD_NAMES + .iter() + .find(|n| map.get(**n).is_some()) + .with_context(|| { + format!( + "no output projection or embedding in {}", + index_path.display() + ) + })?; + let file = map[*name] + .as_str() + .context("weight_map entry is not a file name")?; + (dir.join(file), Some(name.to_string())) + } else { + (dir.join("model.safetensors"), None) + }; + let file = std::fs::File::open(&file)?; + // SAFETY: the checkpoint is not modified while the worker runs. + let mmap = unsafe { memmap2::Mmap::map(&file)? }; + let st = safetensors::SafeTensors::deserialize(&mmap)?; + let name = match name { + Some(n) => n, + None => { + let names = st.names(); + HEAD_NAMES + .iter() + .find(|n| names.iter().any(|m| m == *n)) + .context("no output projection or embedding in model.safetensors")? + .to_string() + } + }; + let view = st.tensor(&name)?; + ensure!( + view.dtype() == safetensors::Dtype::BF16 && view.shape().len() == 2, + "{name}: expected a 2-D bfloat16 tensor, got {:?} {:?}", + view.dtype(), + view.shape() + ); + let (vocab, hidden) = (view.shape()[0], view.shape()[1]); + let data = view.data(); + let mut rows = Vec::with_capacity(letter_ids.len() * hidden); + for &id in letter_ids { + let id = id as usize; + ensure!(id < vocab, "letter id {id} outside the vocabulary"); + let row = &data[id * hidden * 2..(id + 1) * hidden * 2]; + rows.extend( + row.chunks_exact(2) + .map(|b| half::bf16::from_le_bytes([b[0], b[1]]).to_f32()), + ); + } + Ok((rows, hidden)) +} + +pub struct Engine { + pub prompter: Prompter, + model: Arc>, + letters: Vec, + hidden: usize, + pub load_seconds: f64, + pub device: String, +} + +impl Engine { + /// Load the CUDA library and the model, and prepare CUDA graphs (see `Options`). + pub async fn load(dir: &Path, opts: &Options) -> Result { + let started = Instant::now(); + let prompter = Prompter::load(dir)?; + let (letters, hidden) = letter_rows(dir, &prompter.letter_ids)?; + let dir = dir.to_path_buf(); + let opts = opts.clone(); + let model = tokio::task::spawn_blocking(move || Model::load(&dir, &opts)).await??; + ensure!( + model.cfg.hidden == hidden, + "hidden size {} does not match the head ({hidden})", + model.cfg.hidden + ); + Ok(Self { + prompter, + model: Arc::new(Mutex::new(model)), + letters, + hidden, + load_seconds: started.elapsed().as_secs_f64(), + device: "cuda".to_string(), + }) + } + + fn lock(&self) -> Result> { + self.model + .lock() + .map_err(|_| anyhow::anyhow!("model lock poisoned")) + } + + /// The longest prompt that runs as a CUDA graph (0: none). + pub fn graph_max_tokens(&self) -> usize { + self.lock().map(|m| m.graph_max_tokens()).unwrap_or(0) + } + + /// The final-norm hidden state at the last position. + async fn last_hidden(&self, ids: Vec, mode: Mode) -> Result> { + let model = self.model.clone(); + tokio::task::spawn_blocking(move || { + let mut model = model + .lock() + .map_err(|_| anyhow::anyhow!("model lock poisoned"))?; + model.forward(&ids, mode) + }) + .await? + } + + /// One forward pass on the calling thread, for timing. + pub fn last_hidden_blocking(&self, ids: &[u32]) -> Result> { + self.lock()?.forward(ids, Mode::Auto) + } + + /// Option probabilities for one prompt: the final-norm hidden state at the last + /// position times the letter rows, in float32 with float64 accumulation, then a + /// softmax over the first `n_options` letters. + pub async fn score(&self, ids: Vec, n_options: usize) -> Result> { + self.score_mode(ids, n_options, Mode::Auto).await + } + + /// As `score`, run eagerly or from a graph. + pub async fn score_mode( + &self, + ids: Vec, + n_options: usize, + mode: Mode, + ) -> Result> { + let last = self.last_hidden(ids, mode).await?; + if last.len() != self.hidden { + bail!( + "hidden size {} does not match the head ({})", + last.len(), + self.hidden + ); + } + let logits: Vec = self + .letters + .chunks_exact(self.hidden) + .take(n_options) + .map(|w| { + w.iter() + .zip(&last) + .map(|(&a, &b)| a as f64 * b as f64) + .sum::() as f32 + }) + .collect(); + let max = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max) as f64; + let exps: Vec = logits.iter().map(|&l| (l as f64 - max).exp()).collect(); + let total: f64 = exps.iter().sum(); + Ok(exps.iter().map(|e| (e / total) as f32).collect()) + } +} diff --git a/src/models/cua_s1/native/src/lib.rs b/src/models/cua_s1/native/src/lib.rs new file mode 100644 index 00000000..c9e478bb --- /dev/null +++ b/src/models/cua_s1/native/src/lib.rs @@ -0,0 +1,11 @@ +//! A native `/v1/systemone` worker for Cua-S1 4B 0.2 (`text` adapter): request +//! handling, tokenization and scoring in Rust, the Qwen3.5 forward pass on the CUDA +//! kernels of `src/backends/cuda/qwen3_5`, loaded at run time. + +pub mod contract; +pub mod cuda; +pub mod engine; +pub mod model; +pub mod printable; +pub mod pyjson; +pub mod server; diff --git a/src/models/cua_s1/native/src/main.rs b/src/models/cua_s1/native/src/main.rs new file mode 100644 index 00000000..81240c64 --- /dev/null +++ b/src/models/cua_s1/native/src/main.rs @@ -0,0 +1,257 @@ +//! Cua-S1 4B 0.2 (`text` adapter) `/v1/systemone` worker on native CUDA kernels. +//! +//! omni-cua-s1-native --model [--port 8000] +//! +//! See recipe/cua_s1/native.md for building the CUDA library, exporting the merged +//! checkpoint and checking the worker. + +use std::io::{BufRead, Write}; +use std::path::{Path, PathBuf}; +use std::sync::Arc; +use std::time::Instant; + +use anyhow::{Result, ensure}; +use clap::Parser; +use clap::builder::RangedU64ValueParser; + +use omni_cua_s1_native::contract::{self, map_request, parse_body}; +use omni_cua_s1_native::cuda; +use omni_cua_s1_native::engine::{self, Engine, Mode, Options, Prompter}; +use omni_cua_s1_native::server::{self, App, DecideError, Limits}; + +#[derive(Parser)] +#[command(about = "Cua-S1 4B 0.2 text worker on native CUDA kernels")] +struct Args { + /// Merged text checkpoint: Qwen/Qwen3.5-4B with the `text` adapter merged, plus + /// the cua_s1_export.json that recipe/cua_s1/export_text_merged.py writes. + #[arg(long, env = "CUA_S1_MODEL")] + model: PathBuf, + /// libqwen3_5_cuda.so, built by src/backends/cuda/qwen3_5/build.sh [default: next + /// to this executable]. + #[arg(long, env = "CUA_S1_CUDA_LIB")] + cuda_lib: Option, + #[arg(long, env = "CUA_S1_HOST", default_value = "127.0.0.1")] + host: String, + #[arg(long, env = "CUA_S1_PORT", default_value_t = 8000)] + port: u16, + #[arg(long, env = "CUA_S1_MAX_BODY_BYTES", default_value_t = 4 << 20)] + max_body_bytes: usize, + #[arg(long, env = "CUA_S1_MAX_QUESTIONS", default_value_t = 64)] + max_questions: usize, + /// Per question; 0 disables the check. + #[arg(long, env = "CUA_S1_MAX_PROMPT_TOKENS", default_value_t = 16384)] + max_prompt_tokens: usize, + /// Prompts up to this many tokens run as a CUDA graph captured for their length + /// on first use; longer ones run eagerly. 0 runs everything eagerly, with + /// cuBLASLt's first-choice GEMM algorithms and no tuning. + #[arg(long, env = "CUA_S1_GRAPH_MAX_TOKENS", default_value_t = 2048)] + graph_max_tokens: usize, + /// How many prompt lengths keep their captured graph. + #[arg(long, env = "CUA_S1_GRAPH_CACHE", default_value_t = 128, + value_parser = RangedU64ValueParser::::new().range(1..))] + graph_cache: usize, + /// GEMM algorithm choices: read from this file if it exists, else tuned at + /// startup and written to it, so later starts make the same choices. A file tuned + /// on another GPU or cuBLASLt version, or for another --graph-max-tokens, is + /// refused. + #[arg(long, env = "CUA_S1_GEMM_PLANS")] + gemm_plans: Option, + /// When tuning, time far more cuBLASLt configurations for prompts up to + /// --graph-max-tokens (each algorithm with its tiles, stage counts, swizzles and + /// several split-K factors) instead of the heuristic's shortlist. Takes about a + /// minute; use it with --gemm-plans so that it runs once. + #[arg(long, env = "CUA_S1_GEMM_SEARCH")] + gemm_search: bool, + /// Read request bodies on stdin (one JSON string per line), print the prompt + /// token ids or the rejection for each, and exit. Loads only the tokenizer. + #[arg(long)] + encode_only: bool, + /// Score every question of the request bodies in this JSON file (name -> body), + /// eagerly and as served, print one JSON line per run, and exit. + #[arg(long)] + score_all: Option, + /// Time the forward pass over the token ids in each of these JSON files (lists + /// of integers), without HTTP, and exit. + #[arg(long, num_args = 1..)] + bench: Vec, + #[arg(long, default_value_t = 50, value_parser = RangedU64ValueParser::::new().range(1..))] + bench_repeat: usize, + /// Idle time before each timed forward pass, as between separate requests. + #[arg(long, default_value_t = 0)] + bench_gap_ms: u64, +} + +impl Args { + fn options(&self) -> Result { + ensure!( + self.graph_max_tokens > 0 || (self.gemm_plans.is_none() && !self.gemm_search), + "--gemm-plans and --gemm-search need --graph-max-tokens above 0" + ); + Ok(Options { + library: match &self.cuda_lib { + Some(path) => path.clone(), + None => cuda::default_library()?, + }, + graph_max_tokens: self.graph_max_tokens, + graph_cache: self.graph_cache, + gemm_plans: self.gemm_plans.clone(), + gemm_search: self.gemm_search, + }) + } +} + +fn encode_only(args: &Args) -> Result<()> { + let prompter = Prompter::load(&args.model)?; + let mut out = std::io::stdout().lock(); + for line in std::io::stdin().lock().lines() { + let body: String = serde_json::from_str(&line?)?; + let result = parse_body(body.as_bytes()) + .and_then(|b| map_request(&b, args.max_questions)) + .map_err(DecideError::Request) + .and_then(|r| { + let ids = server::encode_all(&prompter, &r, args.max_prompt_tokens)?; + Ok((r, ids)) + }); + let record = match result { + Ok((request, ids)) => serde_json::json!({ + "status": 200, + "questions": request.questions.iter().map(|q| &q.name).collect::>(), + "ids": ids, + }), + Err(DecideError::Request(e)) => { + serde_json::json!({"status": e.status, "detail": e.message}) + } + Err(DecideError::Internal(e)) => return Err(e), + }; + writeln!(out, "{record}")?; + } + Ok(()) +} + +async fn score_all(args: &Args, inputs: &Path) -> Result<()> { + let cases: serde_json::Map = + serde_json::from_str(&std::fs::read_to_string(inputs)?)?; + let engine = Engine::load(&args.model, &args.options()?).await?; + let graph_max = engine.graph_max_tokens(); + let mut out = std::io::stdout().lock(); + for (case, body) in &cases { + let raw = serde_json::to_vec(body)?; + let request = parse_body(&raw) + .and_then(|b| map_request(&b, args.max_questions)) + .map_err(|e| anyhow::anyhow!("{case}: {}", e.message))?; + for question in &request.questions { + let ids = engine.prompter.encode(&request.state, question)?; + for (mode, label) in [(Mode::Eager, "eager"), (Mode::Auto, "served")] { + let probs = engine + .score_mode(ids.clone(), question.keys.len(), mode) + .await?; + let probabilities: serde_json::Map = question + .keys + .iter() + .zip(&probs) + .map(|(k, p)| (k.clone(), serde_json::json!(p))) + .collect(); + writeln!( + out, + "{}", + serde_json::json!({ + "case": case, + "question": question.name, + "tokens": ids.len(), + "mode": label, + "graph": mode == Mode::Auto && ids.len() <= graph_max, + "probabilities": probabilities, + }) + )?; + } + } + } + Ok(()) +} + +async fn bench(args: &Args) -> Result<()> { + let engine = Engine::load(&args.model, &args.options()?).await?; + println!( + "loaded in {:.1} s, graphs up to {} tokens", + engine.load_seconds, + engine.graph_max_tokens() + ); + for file in &args.bench { + let ids: Vec = serde_json::from_str(&std::fs::read_to_string(file)?)?; + let mut times = Vec::with_capacity(args.bench_repeat); + for i in 0..args.bench_repeat + 3 { + std::thread::sleep(std::time::Duration::from_millis(args.bench_gap_ms)); + let started = Instant::now(); + engine.last_hidden_blocking(&ids)?; + if i >= 3 { + times.push(started.elapsed().as_secs_f64() * 1e3); + } + } + times.sort_by(f64::total_cmp); + let at = |q: f64| times[((times.len() - 1) as f64 * q).round() as usize]; + println!( + "{}: {} tokens, forward p50 {:.2} ms p95 {:.2} ms min {:.2} ms", + file.display(), + ids.len(), + at(0.5), + at(0.95), + times[0] + ); + } + Ok(()) +} + +#[tokio::main] +async fn main() -> Result<()> { + let args = Args::parse(); + if args.encode_only { + return encode_only(&args); + } + if let Some(inputs) = &args.score_all { + return score_all(&args, inputs).await; + } + if !args.bench.is_empty() { + return bench(&args).await; + } + let provenance = engine::provenance(&args.model)?; + if provenance.adapter_revision != contract::ADAPTER_REVISION { + eprintln!( + "warning: adapter revision {} is not the pinned {}", + provenance.adapter_revision, + contract::ADAPTER_REVISION + ); + } + if provenance.base_revision != contract::BASE_REVISION { + eprintln!( + "warning: base revision {} is not the pinned {}", + provenance.base_revision, + contract::BASE_REVISION + ); + } + let engine = Engine::load(&args.model, &args.options()?).await?; + println!( + "loaded in {:.1} s on {} (bfloat16, graphs up to {} tokens)", + engine.load_seconds, + engine.device, + engine.graph_max_tokens() + ); + let api_key = std::env::var_os("CUA_S1_API_KEY").map(|k| k.into_encoded_bytes()); + let limits = Limits { + max_body_bytes: args.max_body_bytes, + max_questions: args.max_questions, + max_prompt_tokens: args.max_prompt_tokens, + }; + let app = Arc::new(App::new( + engine, + limits, + api_key, + &provenance.adapter_revision, + )); + let started = Instant::now(); + server::warmup(&app).await?; + println!("warmed up in {:.1} s", started.elapsed().as_secs_f64()); + let listener = tokio::net::TcpListener::bind((args.host.as_str(), args.port)).await?; + println!("listening on {}:{}", args.host, args.port); + axum::serve(listener, server::router(app)).await?; + Ok(()) +} diff --git a/src/models/cua_s1/native/src/model.rs b/src/models/cua_s1/native/src/model.rs new file mode 100644 index 00000000..625810ff --- /dev/null +++ b/src/models/cua_s1/native/src/model.rs @@ -0,0 +1,1159 @@ +//! The Qwen3.5 text model (the language model of Qwen/Qwen3.5-4B), prefill only: one +//! forward pass over a prompt, returning the final-norm hidden state of the last +//! position. The layer loop, buffers, CUDA graphs and kernel choice live here; the +//! operations are the CUDA kernels in `src/backends/cuda/qwen3_5`. +//! +//! Prompts up to `graph_max_tokens` run as a CUDA graph captured for their exact +//! length on first use (the most recent `graph_cache` lengths are kept); longer ones +//! run eagerly. A graph queues the same kernels with the same GEMM algorithms as the +//! eager pass of that length, so both give bitwise identical results. +//! +//! GEMM algorithms are tuned at startup for the lengths in TUNE_ROWS; other lengths +//! use the nearest tuned one (see gemm.cu). The choices can be saved to a file and +//! reused, so that restarts do not change them; the file records the GPU, the +//! cuBLASLt version and the tuned lengths, and one that does not match is refused. +//! +//! The order of operations follows `modeling_qwen3_5.py`, and so do the points where +//! it rounds to bfloat16, except inside attention and the Gated DeltaNet prefill (see +//! their kernels). Text prompts use one position per +//! token, so the multimodal rotary sections all get the same position and the +//! rotary embedding is the plain one. + +use std::collections::{HashMap, VecDeque}; +use std::ffi::c_void; +use std::path::{Path, PathBuf}; + +use anyhow::{Context, Result, bail, ensure}; +use serde_json::Value as Json; + +use crate::cuda::{self, DeviceBuffer, Graph, Stream, check}; + +const ALIGN: usize = 256; +const BF16: usize = 2; +const F32: usize = 4; +const GEMM_WORKSPACE: usize = 32 << 20; + +/// What the GEMM plans depend on besides the shapes: the GPU, cuBLASLt, the +/// workspace, and the longest prompt tuned in the graph range. +fn plan_setup(graph_max_tokens: usize) -> Result { + let (gpu, compute_capability, sms) = cuda::device_info()?; + // SAFETY: takes no arguments. + let cublaslt = unsafe { (cuda::api().cs1_gemm_version)() }; + Ok(serde_json::json!({ + "gpu": gpu, + "compute_capability": compute_capability, + "sms": sms, + "cublaslt": cublaslt, + "workspace_bytes": GEMM_WORKSPACE, + "graph_max_tokens": graph_max_tokens, + })) +} + +/// Prompt lengths the GEMM algorithms are tuned for. Past these, cuBLASLt's first +/// choice for long prompts is an older, half-rate tensor-core kernel on sm_89. +pub const TUNE_ROWS: &[usize] = &[ + 64, 96, 128, 160, 192, 224, 256, 320, 384, 448, 512, 640, 768, 1024, 1536, 2048, 4096, 8192, + 16384, +]; + +/// Startup options of the model. +#[derive(Debug, Clone)] +pub struct Options { + /// libqwen3_5_cuda.so (see src/backends/cuda/qwen3_5/build.sh). + pub library: PathBuf, + /// Longest prompt that runs as a CUDA graph; 0 runs everything eagerly and skips + /// GEMM tuning. + pub graph_max_tokens: usize, + /// How many prompt lengths keep their captured graph. + pub graph_cache: usize, + /// Tuned GEMM algorithms: read from this file if it exists, else tuned and + /// written to it. + pub gemm_plans: Option, + /// Tune the GEMMs of graph-length prompts over far more cuBLASLt configurations + /// than the heuristic's shortlist (see gemm.cu; about a minute). Longer prompts + /// keep the shortlist, whose choices did better in whole forward passes. + pub gemm_search: bool, +} + +#[derive(Debug, Clone)] +pub struct Config { + pub hidden: usize, + pub intermediate: usize, + pub eps: f32, + pub full_attention: Vec, + pub heads: usize, + pub kv_heads: usize, + pub head_dim: usize, + /// Half the number of rotary dims (rotate_half pairs dim i with dim i + half). + pub rotary_half: usize, + pub rope_theta: f64, + pub lin_k_heads: usize, + pub lin_v_heads: usize, + pub lin_k_dim: usize, + pub lin_v_dim: usize, +} + +impl Config { + pub fn load(dir: &Path) -> Result { + let path = dir.join("config.json"); + let root: Json = serde_json::from_str( + &std::fs::read_to_string(&path).with_context(|| format!("{}", path.display()))?, + )?; + let c = root.get("text_config").unwrap_or(&root); + let int = |k: &str| { + c[k].as_u64() + .map(|v| v as usize) + .with_context(|| format!("config.json: `{k}` is missing")) + }; + let rope = &c["rope_parameters"]; + let partial = rope["partial_rotary_factor"] + .as_f64() + .or(c["partial_rotary_factor"].as_f64()) + .unwrap_or(1.0); + let head_dim = int("head_dim")?; + let full_attention = c["layer_types"] + .as_array() + .context("config.json: `layer_types` is missing")? + .iter() + .map(|t| match t.as_str() { + Some("full_attention") => Ok(true), + Some("linear_attention") => Ok(false), + other => bail!("unknown layer type {other:?}"), + }) + .collect::>>()?; + let cfg = Config { + hidden: int("hidden_size")?, + intermediate: int("intermediate_size")?, + eps: c["rms_norm_eps"].as_f64().context("rms_norm_eps")? as f32, + heads: int("num_attention_heads")?, + kv_heads: int("num_key_value_heads")?, + head_dim, + rotary_half: (head_dim as f64 * partial) as usize / 2, + rope_theta: rope["rope_theta"] + .as_f64() + .or(c["rope_theta"].as_f64()) + .context("rope_theta")?, + lin_k_heads: int("linear_num_key_heads")?, + lin_v_heads: int("linear_num_value_heads")?, + lin_k_dim: int("linear_key_head_dim")?, + lin_v_dim: int("linear_value_head_dim")?, + full_attention, + }; + // What the kernels implement. + ensure!( + cfg.full_attention.len() == int("num_hidden_layers")?, + "layer_types does not match num_hidden_layers" + ); + ensure!(c["hidden_act"] == "silu", "hidden_act is not silu"); + ensure!( + c["attn_output_gate"].as_bool().unwrap_or(true), + "attention without the output gate" + ); + ensure!( + c["attention_bias"].as_bool() != Some(true), + "attention with bias" + ); + ensure!( + rope["rope_type"].as_str().unwrap_or("default") == "default", + "rope type {}", + rope["rope_type"] + ); + ensure!(int("linear_conv_kernel_dim")? == 4, "conv kernel is not 4"); + ensure!( + cfg.head_dim == 256 && cfg.lin_k_dim == 128 && cfg.lin_v_dim == 128, + "head dims {} / {} / {}", + cfg.head_dim, + cfg.lin_k_dim, + cfg.lin_v_dim + ); + ensure!(cfg.rotary_half == 32, "{} rotary dims", 2 * cfg.rotary_half); + ensure!( + cfg.kv_heads > 0 && cfg.heads.is_multiple_of(cfg.kv_heads), + "attention heads" + ); + ensure!( + cfg.lin_k_heads > 0 && cfg.lin_v_heads.is_multiple_of(cfg.lin_k_heads), + "linear attention heads" + ); + ensure!(cfg.hidden.is_multiple_of(8), "hidden size"); + Ok(cfg) + } + + fn key_dim(&self) -> usize { + self.lin_k_heads * self.lin_k_dim + } + + fn value_dim(&self) -> usize { + self.lin_v_heads * self.lin_v_dim + } +} + +/// A weight in the device arena. +#[derive(Clone)] +struct Tensor { + ptr: *const c_void, + shape: Vec, +} + +impl Tensor { + fn bytes(&self) -> usize { + self.shape.iter().product::() * BF16 + } +} + +/// Projections that run as one GEMM, in the order their rows are stacked. +const GROUPS: &[&str] = &[ + "linear_attn.in_proj_qkv.weight", + "linear_attn.in_proj_z.weight", + "linear_attn.in_proj_b.weight", + "linear_attn.in_proj_a.weight", + "self_attn.q_proj.weight", + "self_attn.k_proj.weight", + "self_attn.v_proj.weight", + "mlp.gate_proj.weight", + "mlp.up_proj.weight", +]; + +/// Upload order: by layer, and inside a layer the GROUPS members first and in order, +/// so each group's matrices sit back to back and form one [sum N, K] matrix. +fn upload_order(name: &str) -> (usize, usize, String) { + if let Some(tail) = name.strip_prefix("layers.") + && let Some((layer, rest)) = tail.split_once('.') + && let Ok(layer) = layer.parse::() + { + let rank = GROUPS + .iter() + .position(|g| *g == rest) + .unwrap_or(GROUPS.len()); + return (layer, rank, rest.to_string()); + } + (usize::MAX, 0, name.to_string()) +} + +struct Weights { + _arena: DeviceBuffer, + tensors: HashMap, + prefix: String, +} + +impl Weights { + /// Upload every bfloat16 tensor of the language model into one allocation. + fn load(dir: &Path, stream: Stream) -> Result { + let index = dir.join("model.safetensors.index.json"); + let mut files: Vec = if index.exists() { + let index: Json = serde_json::from_str(&std::fs::read_to_string(&index)?)?; + index["weight_map"] + .as_object() + .context("weight_map")? + .values() + .filter_map(|v| v.as_str().map(str::to_string)) + .collect() + } else { + vec!["model.safetensors".to_string()] + }; + files.sort(); + files.dedup(); + let maps = files + .iter() + .map(|f| { + let file = std::fs::File::open(dir.join(f)).with_context(|| f.clone())?; + // SAFETY: the checkpoint is not modified while it is loaded. + Ok(unsafe { memmap2::Mmap::map(&file)? }) + }) + .collect::>>()?; + let sts = maps + .iter() + .map(|m| safetensors::SafeTensors::deserialize(m).map_err(anyhow::Error::from)) + .collect::>>()?; + let names: Vec<(usize, String)> = sts + .iter() + .enumerate() + .flat_map(|(i, st)| st.names().into_iter().map(move |n| (i, n.to_string()))) + .collect(); + let prefix = ["model.language_model.", "model."] + .into_iter() + .find(|p| { + names + .iter() + .any(|(_, n)| *n == format!("{p}embed_tokens.weight")) + }) + .context("no embed_tokens.weight in the checkpoint")? + .to_string(); + let mut ours: Vec<(usize, String)> = names + .into_iter() + .filter(|(_, n)| n.starts_with(&prefix)) + .collect(); + ours.sort_by_key(|(_, n)| upload_order(&n[prefix.len()..])); + let mut plan = Vec::new(); + let mut total = 0usize; + for (i, name) in ours { + let view = sts[i].tensor(&name)?; + ensure!( + view.dtype() == safetensors::Dtype::BF16, + "{name} is {:?}, not bfloat16", + view.dtype() + ); + plan.push((i, name, total)); + total = (total + view.data().len()).next_multiple_of(ALIGN); + } + let arena = DeviceBuffer::new(total)?; + let mut tensors = HashMap::new(); + for (i, name, offset) in plan { + let view = sts[i].tensor(&name)?; + // SAFETY: the arena has room for every planned tensor at its offset. + unsafe { cuda::upload(arena.at(offset), view.data(), stream)? }; + tensors.insert( + name[prefix.len()..].to_string(), + Tensor { + ptr: arena.at(offset), + shape: view.shape().to_vec(), + }, + ); + } + Ok(Self { + _arena: arena, + tensors, + prefix, + }) + } + + fn get(&self, name: &str, shape: &[usize]) -> Result { + let t = self + .tensors + .get(name) + .with_context(|| format!("{}{name} is missing", self.prefix))?; + ensure!( + t.shape == shape, + "{}{name}: shape {:?}, expected {:?}", + self.prefix, + t.shape, + shape + ); + Ok(t.clone()) + } + + /// The row-stacked matrix of tensors that were uploaded back to back. + fn stacked(&self, parts: &[Tensor]) -> Result { + let k = parts[0].shape[1]; + let mut rows = 0; + for (i, p) in parts.iter().enumerate() { + ensure!( + p.shape.len() == 2 && p.shape[1] == k, + "stacked shapes differ" + ); + if i > 0 { + let prev = &parts[i - 1]; + ensure!( + p.ptr == prev.ptr.wrapping_byte_add(prev.bytes()), + "stacked weights are not contiguous" + ); + } + rows += p.shape[0]; + } + Ok(Tensor { + ptr: parts[0].ptr, + shape: vec![rows, k], + }) + } +} + +struct LinearAttention { + /// in_proj_qkv | in_proj_z | in_proj_b | in_proj_a + in_proj: Tensor, + conv: Tensor, + a_log: Tensor, + dt_bias: Tensor, + norm: Tensor, + out: Tensor, +} + +struct FullAttention { + /// q_proj (query and gate per head) | k_proj | v_proj + qkv: Tensor, + o: Tensor, + q_norm: Tensor, + k_norm: Tensor, +} + +enum Mixer { + Linear(LinearAttention), + Full(FullAttention), +} + +struct Layer { + input_norm: Tensor, + post_norm: Tensor, + mixer: Mixer, + /// gate_proj | up_proj + gate_up: Tensor, + down: Tensor, +} + +/// Row widths of the stacked projection outputs. +struct Widths { + conv: usize, + gdn_in: usize, + attn_q: usize, + attn_in: usize, +} + +impl Widths { + fn of(cfg: &Config) -> Self { + let conv = 2 * cfg.key_dim() + cfg.value_dim(); + let attn_q = cfg.heads * cfg.head_dim * 2; + Self { + conv, + gdn_in: conv + cfg.value_dim() + 2 * cfg.lin_v_heads, + attn_q, + attn_in: attn_q + 2 * cfg.kv_heads * cfg.head_dim, + } + } +} + +/// Per-request buffers for up to `cap` tokens, as byte offsets into one allocation, +/// plus the rotary tables for positions below `cap`. +struct Scratch { + cap: usize, + buf: DeviceBuffer, + ids: usize, + res: usize, + x: usize, + delta: usize, + gdn_in: usize, + beta: usize, + g: usize, + lq: usize, + lk: usize, + lv: usize, + lo: usize, + ln: usize, + workspace: usize, + attn_in: usize, + aq: usize, + agate: usize, + ak: usize, + ao: usize, + gate_up: usize, + act: usize, + cos: usize, + sin: usize, +} + +impl Scratch { + fn new(cfg: &Config, cap: usize, stream: Stream) -> Result { + let (h, kd, vd, hv) = (cfg.hidden, cfg.key_dim(), cfg.value_dim(), cfg.lin_v_heads); + let (hq, hk, hd) = (cfg.heads, cfg.kv_heads, cfg.head_dim); + let w = Widths::of(cfg); + let mut next = 0usize; + let mut take = |bytes: usize| { + let off = next; + next = (off + bytes).next_multiple_of(ALIGN); + off + }; + // SAFETY: pure function of its arguments. + let ws_floats = unsafe { (cuda::api().cs1_gdn_workspace_floats)(cap as i32, hv as i32) }; + let offsets = [ + take(cap * 4), + take(cap * h * BF16), + take(cap * h * BF16), + take(cap * h * BF16), + take(cap * w.gdn_in * BF16), + take(cap * hv * BF16), + take(cap * hv * F32), + take(cap * kd * BF16), + take(cap * kd * BF16), + take(cap * vd * BF16), + take(cap * vd * BF16), + take(cap * vd * BF16), + take(ws_floats * F32), + take(cap * w.attn_in * BF16), + take(cap * hq * hd * BF16), + take(cap * hq * hd * BF16), + take(cap * hk * hd * BF16), + take(cap * hq * hd * BF16), + take(cap * 2 * cfg.intermediate * BF16), + take(cap * cfg.intermediate * BF16), + take(cap * cfg.rotary_half * BF16), + take(cap * cfg.rotary_half * BF16), + ]; + let buf = DeviceBuffer::new(next)?; + let [ + ids, + res, + x, + delta, + gdn_in, + beta, + g, + lq, + lk, + lv, + lo, + ln, + workspace, + attn_in, + aq, + agate, + ak, + ao, + gate_up, + act, + cos, + sin, + ] = offsets; + // Rotary tables close to how Qwen3_5TextRotaryEmbedding builds them: inv_freq and + // freqs = inv_freq * position in float32, cos and sin rounded to bfloat16. Here + // cos and sin are taken in float64 on the host rather than in float32 on the + // GPU, so a few of the rounded values can differ by one bfloat16 step. + let half = cfg.rotary_half; + let inv: Vec = (0..half) + .map(|i| 1.0f32 / (cfg.rope_theta as f32).powf((2 * i) as f32 / (2 * half) as f32)) + .collect(); + let mut cos_t = Vec::with_capacity(cap * half * BF16); + let mut sin_t = Vec::with_capacity(cap * half * BF16); + for pos in 0..cap { + for &f in &inv { + let freq = (f * pos as f32) as f64; + cos_t.extend(half::bf16::from_f32(freq.cos() as f32).to_le_bytes()); + sin_t.extend(half::bf16::from_f32(freq.sin() as f32).to_le_bytes()); + } + } + // SAFETY: both tables were laid out for cap * rotary_half bfloat16 values. + unsafe { + cuda::upload(buf.at(cos), &cos_t, stream)?; + cuda::upload(buf.at(sin), &sin_t, stream)?; + } + Ok(Self { + cap, + buf, + ids, + res, + x, + delta, + gdn_in, + beta, + g, + lq, + lk, + lv, + lo, + ln, + workspace, + attn_in, + aq, + agate, + ak, + ao, + gate_up, + act, + cos, + sin, + }) + } + + fn at(&self, offset: usize) -> *mut c_void { + self.buf.at(offset) + } +} + +/// How to run one forward pass. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Mode { + /// From the graph for this length when graphs are on and it fits, else eagerly. + Auto, + Eager, +} + +pub struct Model { + pub cfg: Config, + _weights: Weights, + embed: Tensor, + final_norm: Tensor, + layers: Vec, + stream: Stream, + gemm: *mut c_void, + /// Graphs by prompt length, least recently used first in `graph_lru`; they point + /// into `graph_scratch`, which is never reallocated. + graphs: HashMap, + graph_lru: VecDeque, + graph_scratch: Option, + graph_cache: usize, + /// For prompts longer than `graph_scratch` holds; grows as needed. + eager_scratch: Option, +} + +// SAFETY: the raw pointers are device addresses and a cuBLASLt handle owned by the +// model; the engine runs one forward pass at a time behind a mutex. +unsafe impl Send for Model {} + +impl Drop for Model { + fn drop(&mut self) { + self.graphs.clear(); + // SAFETY: created by cs1_gemm_create and not destroyed before. + unsafe { (cuda::api().cs1_gemm_destroy)(self.gemm) }; + } +} + +impl Model { + /// Load the weights, then tune the GEMMs (or read their plans) for graph use. + pub fn load(dir: &Path, opts: &Options) -> Result { + let cfg = Config::load(dir)?; + cuda::load(&opts.library)?; + cuda::set_device(0)?; + let stream = cuda::new_stream()?; + let weights = Weights::load(dir, stream)?; + let (h, kd, vd) = (cfg.hidden, cfg.key_dim(), cfg.value_dim()); + let embed = weights + .tensors + .get("embed_tokens.weight") + .context("embed_tokens.weight is missing")? + .clone(); + ensure!( + embed.shape.len() == 2 && embed.shape[1] == h, + "embed_tokens.weight shape {:?}", + embed.shape + ); + let final_norm = weights.get("norm.weight", &[h])?; + let mut layers = Vec::with_capacity(cfg.full_attention.len()); + for (i, &full) in cfg.full_attention.iter().enumerate() { + let w = |n: &str, s: &[usize]| weights.get(&format!("layers.{i}.{n}"), s); + let mixer = if full { + let (hq, hk, hd) = (cfg.heads, cfg.kv_heads, cfg.head_dim); + Mixer::Full(FullAttention { + qkv: weights.stacked(&[ + w("self_attn.q_proj.weight", &[hq * hd * 2, h])?, + w("self_attn.k_proj.weight", &[hk * hd, h])?, + w("self_attn.v_proj.weight", &[hk * hd, h])?, + ])?, + o: w("self_attn.o_proj.weight", &[h, hq * hd])?, + q_norm: w("self_attn.q_norm.weight", &[hd])?, + k_norm: w("self_attn.k_norm.weight", &[hd])?, + }) + } else { + let hv = cfg.lin_v_heads; + Mixer::Linear(LinearAttention { + in_proj: weights.stacked(&[ + w("linear_attn.in_proj_qkv.weight", &[2 * kd + vd, h])?, + w("linear_attn.in_proj_z.weight", &[vd, h])?, + w("linear_attn.in_proj_b.weight", &[hv, h])?, + w("linear_attn.in_proj_a.weight", &[hv, h])?, + ])?, + conv: w("linear_attn.conv1d.weight", &[2 * kd + vd, 1, 4])?, + a_log: w("linear_attn.A_log", &[hv])?, + dt_bias: w("linear_attn.dt_bias", &[hv])?, + norm: w("linear_attn.norm.weight", &[cfg.lin_v_dim])?, + out: w("linear_attn.out_proj.weight", &[h, vd])?, + }) + }; + layers.push(Layer { + input_norm: w("input_layernorm.weight", &[h])?, + post_norm: w("post_attention_layernorm.weight", &[h])?, + mixer, + gate_up: weights.stacked(&[ + w("mlp.gate_proj.weight", &[cfg.intermediate, h])?, + w("mlp.up_proj.weight", &[cfg.intermediate, h])?, + ])?, + down: w("mlp.down_proj.weight", &[h, cfg.intermediate])?, + }); + } + // SAFETY: plain allocation; checked for null below. + let gemm = unsafe { (cuda::api().cs1_gemm_create)(GEMM_WORKSPACE) }; + ensure!(!gemm.is_null(), "cuBLASLt setup failed"); + let mut model = Self { + cfg, + _weights: weights, + embed, + final_norm, + layers, + stream, + gemm, + graphs: HashMap::new(), + graph_lru: VecDeque::new(), + graph_scratch: None, + graph_cache: opts.graph_cache.max(1), + eager_scratch: None, + }; + if opts.graph_max_tokens > 0 { + model.prepare_graphs( + opts.graph_max_tokens, + opts.gemm_plans.as_deref(), + opts.gemm_search, + )?; + } + Ok(model) + } + + /// The longest prompt that runs as a graph (0: none). + pub fn graph_max_tokens(&self) -> usize { + self.graph_scratch.as_ref().map_or(0, |s| s.cap) + } + + /// Allocate the graph buffers, then read or tune the GEMM plans. + fn prepare_graphs(&mut self, max: usize, plans: Option<&Path>, search: bool) -> Result<()> { + let s = Scratch::new(&self.cfg, max, self.stream)?; + // one eager pass over the longest prompt fills every buffer and sets up the kernels + let zeros = vec![0u8; max * 4]; + // SAFETY: the ids buffer holds `max` int32 values. + unsafe { cuda::upload(s.at(s.ids), &zeros, self.stream)? }; + self.run(&s, max)?; + cuda::synchronize(self.stream)?; + let setup = plan_setup(max)?; + let loaded = match plans { + Some(path) if path.exists() => { + let n = self.import_plans(path, &setup).with_context(|| { + format!( + "GEMM plans in {} not used; remove the file, or pass another \ + --gemm-plans, to tune again", + path.display() + ) + })?; + eprintln!("GEMM plans: {n} read from {}", path.display()); + true + } + _ => false, + }; + if !loaded { + let mut rows: Vec = TUNE_ROWS.iter().copied().filter(|&m| m < max).collect(); + rows.push(max); + for m in rows { + self.tune(&s, m, search)?; + } + // longer prompts run eagerly; tune those lengths in a temporary buffer + let long: Vec = TUNE_ROWS.iter().copied().filter(|&m| m > max).collect(); + if let Some(&cap) = long.last() { + let big = Scratch::new(&self.cfg, cap, self.stream)?; + for m in long { + self.tune(&big, m, false)?; + } + } + // SAFETY: frees only the tuning buffers. + unsafe { (cuda::api().cs1_gemm_tune_done)(self.gemm) }; + if let Some(path) = plans { + let n = self.export_plans(path, &setup, search)?; + eprintln!("GEMM plans: {n} written to {}", path.display()); + } + } + self.graph_scratch = Some(s); + Ok(()) + } + + fn export_plans(&self, path: &Path, setup: &Json, search: bool) -> Result { + // SAFETY: a null buffer with capacity 0 only counts. + let n = unsafe { (cuda::api().cs1_gemm_export)(self.gemm, std::ptr::null_mut(), 0) }; + let mut plans = vec![cuda::GemmPlan::default(); n]; + // SAFETY: `plans` has room for n records. + unsafe { (cuda::api().cs1_gemm_export)(self.gemm, plans.as_mut_ptr(), n) }; + let records: Vec = plans + .iter() + .map(|p| { + serde_json::json!({ + "m": p.m, "n": p.n, "k": p.k, "ldy": p.ldy, + "algo": p.algo.iter().map(|w| format!("{w:016x}")).collect::>(), + }) + }) + .collect(); + let doc = serde_json::json!({ + "format": "cua-s1-native-gemm-plans/1", + "setup": setup, + "search": search, + "plans": records, + }); + std::fs::write(path, serde_json::to_string_pretty(&doc)? + "\n") + .with_context(|| format!("{}", path.display()))?; + Ok(n) + } + + fn import_plans(&self, path: &Path, setup: &Json) -> Result { + let doc: Json = serde_json::from_str(&std::fs::read_to_string(path)?)?; + ensure!( + doc["format"] == "cua-s1-native-gemm-plans/1", + "unknown format {}", + doc["format"] + ); + ensure!( + doc["setup"].is_object(), + "the file does not record the GPU and cuBLASLt version it was tuned for" + ); + ensure!( + doc["setup"] == *setup, + "tuned for {}, this run is {setup}", + doc["setup"] + ); + let int = |v: &Json| v.as_i64().map(|x| x as i32).context("bad plan field"); + let plans = doc["plans"] + .as_array() + .context("no plans")? + .iter() + .map(|p| { + let words = p["algo"].as_array().context("bad algo")?; + ensure!(words.len() == 8, "bad algo"); + let mut algo = [0u64; 8]; + for (dst, w) in algo.iter_mut().zip(words) { + *dst = u64::from_str_radix(w.as_str().context("bad algo")?, 16)?; + } + Ok(cuda::GemmPlan { + m: int(&p["m"])?, + n: int(&p["n"])?, + k: int(&p["k"])?, + ldy: int(&p["ldy"])?, + algo, + }) + }) + .collect::>>()?; + // SAFETY: `plans` holds plans.len() records. + check( + unsafe { (cuda::api().cs1_gemm_import)(self.gemm, plans.as_ptr(), plans.len()) }, + "importing GEMM plans", + )?; + Ok(plans.len()) + } + + /// Pick the GEMM algorithms for `m` rows, timing the first layer of each kind. + fn tune(&self, s: &Scratch, m: usize, exhaustive: bool) -> Result<()> { + let mut jobs: Vec<(usize, &Tensor, usize)> = Vec::new(); + let layer = &self.layers[0]; + jobs.push((s.x, &layer.gate_up, s.gate_up)); + jobs.push((s.act, &layer.down, s.delta)); + if let Some(la) = self.layers.iter().find_map(|l| match &l.mixer { + Mixer::Linear(la) => Some(la), + Mixer::Full(_) => None, + }) { + jobs.push((s.x, &la.in_proj, s.gdn_in)); + jobs.push((s.ln, &la.out, s.delta)); + } + if let Some(fa) = self.layers.iter().find_map(|l| match &l.mixer { + Mixer::Full(fa) => Some(fa), + Mixer::Linear(_) => None, + }) { + jobs.push((s.x, &fa.qkv, s.attn_in)); + jobs.push((s.ao, &fa.o, s.delta)); + } + for (x, w, y) in jobs { + let (n, k) = (w.shape[0] as i32, w.shape[1] as i32); + // SAFETY: x and y are scratch buffers sized for `cap` >= m rows of this shape. + check( + unsafe { + (cuda::api().cs1_gemm_tune)( + self.gemm, + s.at(x), + w.ptr, + s.at(y), + m as i32, + n, + k, + n, + exhaustive.into(), + self.stream, + ) + }, + "gemm tuning", + )?; + } + Ok(()) + } + + fn gemm(&self, s: &Scratch, x: usize, w: &Tensor, y: usize, m: usize) -> Result<()> { + let (n, k) = (w.shape[0] as i32, w.shape[1] as i32); + // SAFETY: x and y are scratch buffers sized for m rows of w's shape. + check( + unsafe { + (cuda::api().cs1_gemm)( + self.gemm, + s.at(x), + w.ptr, + s.at(y), + m as i32, + n, + k, + n, + self.stream, + ) + }, + "gemm", + ) + } + + /// The final-norm hidden state at the last position, as float32. + pub fn forward(&mut self, ids: &[u32], mode: Mode) -> Result> { + let t = ids.len(); + ensure!(t > 0, "empty prompt"); + let (vocab, h) = (self.embed.shape[0], self.cfg.hidden); + ensure!( + ids.iter().all(|&i| (i as usize) < vocab), + "token id outside the vocabulary" + ); + cuda::set_device(0)?; + let in_graph_scratch = self.graph_scratch.as_ref().is_some_and(|s| t <= s.cap); + if !in_graph_scratch && self.eager_scratch.as_ref().is_none_or(|s| t > s.cap) { + self.eager_scratch = None; + self.eager_scratch = Some(Scratch::new( + &self.cfg, + t.next_multiple_of(1024), + self.stream, + )?); + } + let use_graph = mode == Mode::Auto && in_graph_scratch; + let s = if in_graph_scratch { + self.graph_scratch.as_ref() + } else { + self.eager_scratch.as_ref() + } + .unwrap(); + let ids32: Vec = ids.iter().flat_map(|&i| (i as i32).to_le_bytes()).collect(); + // SAFETY: the ids buffer holds at least t int32 values. + unsafe { cuda::upload(s.at(s.ids), &ids32, self.stream)? }; + if use_graph { + self.graph_for(t)?; + self.graphs[&t].launch(self.stream)?; + } else { + self.run(s, t)?; + } + let s = if in_graph_scratch { + self.graph_scratch.as_ref() + } else { + self.eager_scratch.as_ref() + } + .unwrap(); + let mut last = vec![0u8; h * BF16]; + // SAFETY: x holds at least t rows of the hidden size. + unsafe { cuda::download(&mut last, s.at(s.x + (t - 1) * h * BF16), self.stream)? }; + Ok(last + .chunks_exact(2) + .map(|b| half::bf16::from_le_bytes([b[0], b[1]]).to_f32()) + .collect()) + } + + /// Make sure a graph for `t` tokens exists (capturing it if needed) and mark it + /// as the most recently used, dropping the least recently used beyond the cache. + fn graph_for(&mut self, t: usize) -> Result<()> { + if self.graphs.contains_key(&t) { + self.graph_lru.retain(|&x| x != t); + } else { + while self.graphs.len() >= self.graph_cache { + let Some(old) = self.graph_lru.pop_front() else { + break; + }; + self.graphs.remove(&old); + } + let s = self.graph_scratch.as_ref().context("no graph buffers")?; + let graph = Graph::capture(self.stream, || self.run(s, t))?; + self.graphs.insert(t, graph); + } + self.graph_lru.push_back(t); + Ok(()) + } + + /// Queue one forward pass over the first `t` ids in `s` (nothing else is queued, + /// so it can be captured). The final-norm hidden states end up in `s.x`. + fn run(&self, s: &Scratch, t: usize) -> Result<()> { + let cfg = &self.cfg; + let st = self.stream; + let (ti, hi, eps) = (t as i32, cfg.hidden as i32, cfg.eps); + let (kd, vd, hv) = (cfg.key_dim(), cfg.value_dim(), cfg.lin_v_heads); + let (hq, hk, hd) = (cfg.heads as i32, cfg.kv_heads as i32, cfg.head_dim as i32); + let w = Widths::of(cfg); + let p = |off: usize| s.at(off); + // SAFETY (every kernel call below): the pointers are weights in the arena or + // scratch buffers laid out for at least t tokens with the widths used here. + unsafe { + check( + (cuda::api().cs1_embed)(p(s.ids).cast(), self.embed.ptr, p(s.res), ti, hi, st), + "embed", + )?; + check( + (cuda::api().cs1_rms_norm)( + p(s.res), + self.layers[0].input_norm.ptr, + p(s.x), + ti, + hi, + eps, + st, + ), + "input norm", + )?; + } + for (i, layer) in self.layers.iter().enumerate() { + match &layer.mixer { + Mixer::Linear(la) => { + self.gemm(s, s.x, &la.in_proj, s.gdn_in, t)?; + let ld = w.gdn_in as i32; + let z = s.gdn_in + w.conv * BF16; + let b = z + vd * BF16; + let a = b + hv * BF16; + unsafe { + check( + (cuda::api().cs1_gdn_conv)( + p(s.gdn_in), + ld, + la.conv.ptr, + p(s.lq), + p(s.lk), + p(s.lv), + ti, + kd as i32, + vd as i32, + st, + ), + "gdn conv", + )?; + check( + (cuda::api().cs1_gdn_gates)( + p(b), + p(a), + ld, + la.a_log.ptr, + la.dt_bias.ptr, + p(s.beta), + p(s.g).cast(), + ti, + hv as i32, + st, + ), + "gdn gates", + )?; + check( + (cuda::api().cs1_gdn_prefill)( + p(s.lq), + p(s.lk), + p(s.lv), + p(s.g).cast(), + p(s.beta), + p(s.lo), + p(s.workspace).cast(), + ti, + hv as i32, + cfg.lin_k_heads as i32, + (cfg.lin_k_dim as f32).powf(-0.5), + st, + ), + "gdn prefill", + )?; + check( + (cuda::api().cs1_gated_rms_norm)( + p(s.lo), + p(z), + ld, + la.norm.ptr, + p(s.ln), + ti, + hv as i32, + cfg.lin_v_dim as i32, + eps, + st, + ), + "gated norm", + )?; + } + self.gemm(s, s.ln, &la.out, s.delta, t)?; + } + Mixer::Full(fa) => { + self.gemm(s, s.x, &fa.qkv, s.attn_in, t)?; + let ld = w.attn_in as i32; + let k = s.attn_in + w.attn_q * BF16; + let v = k + cfg.kv_heads * cfg.head_dim * BF16; + unsafe { + check( + (cuda::api().cs1_attn_prep)( + p(s.attn_in), + p(k), + ld, + fa.q_norm.ptr, + fa.k_norm.ptr, + p(s.cos), + p(s.sin), + p(s.aq), + p(s.agate), + p(s.ak), + ti, + hq, + hk, + hd, + cfg.rotary_half as i32, + eps, + st, + ), + "attention prep", + )?; + check( + (cuda::api().cs1_attention)( + p(s.aq), + p(s.ak), + p(v), + ld, + p(s.ao), + ti, + hq, + hk, + hd, + (cfg.head_dim as f32).powf(-0.5), + st, + ), + "attention", + )?; + check( + (cuda::api().cs1_sigmoid_gate)( + p(s.ao), + p(s.agate), + t * cfg.heads * cfg.head_dim, + st, + ), + "attention gate", + )?; + } + self.gemm(s, s.ao, &fa.o, s.delta, t)?; + } + } + unsafe { + check( + (cuda::api().cs1_add_rms_norm)( + p(s.res), + p(s.delta), + layer.post_norm.ptr, + p(s.x), + ti, + hi, + eps, + st, + ), + "post-attention norm", + )?; + } + self.gemm(s, s.x, &layer.gate_up, s.gate_up, t)?; + unsafe { + check( + (cuda::api().cs1_silu_mul)( + p(s.gate_up), + (2 * cfg.intermediate) as i32, + p(s.act), + ti, + cfg.intermediate as i32, + st, + ), + "silu mul", + )?; + } + self.gemm(s, s.act, &layer.down, s.delta, t)?; + let next = self + .layers + .get(i + 1) + .map_or(&self.final_norm, |l| &l.input_norm); + unsafe { + check( + (cuda::api().cs1_add_rms_norm)( + p(s.res), + p(s.delta), + next.ptr, + p(s.x), + ti, + hi, + eps, + st, + ), + "input norm", + )?; + } + } + Ok(()) + } +} diff --git a/src/models/cua_s1/native/src/printable.rs b/src/models/cua_s1/native/src/printable.rs new file mode 100644 index 00000000..9bcd2f9a --- /dev/null +++ b/src/models/cua_s1/native/src/printable.rs @@ -0,0 +1,717 @@ +//! Generated from Python 3.12.13 (Unicode 15.0.0): code points >= 0x80 for which +//! `str.isprintable()` is false, as inclusive ranges. Python's `repr` escapes these. +//! Regenerate with tests/make_printable.py. + +pub const NON_PRINTABLE: &[(u32, u32)] = &[ + (0x80, 0xa0), + (0xad, 0xad), + (0x378, 0x379), + (0x380, 0x383), + (0x38b, 0x38b), + (0x38d, 0x38d), + (0x3a2, 0x3a2), + (0x530, 0x530), + (0x557, 0x558), + (0x58b, 0x58c), + (0x590, 0x590), + (0x5c8, 0x5cf), + (0x5eb, 0x5ee), + (0x5f5, 0x605), + (0x61c, 0x61c), + (0x6dd, 0x6dd), + (0x70e, 0x70f), + (0x74b, 0x74c), + (0x7b2, 0x7bf), + (0x7fb, 0x7fc), + (0x82e, 0x82f), + (0x83f, 0x83f), + (0x85c, 0x85d), + (0x85f, 0x85f), + (0x86b, 0x86f), + (0x88f, 0x897), + (0x8e2, 0x8e2), + (0x984, 0x984), + (0x98d, 0x98e), + (0x991, 0x992), + (0x9a9, 0x9a9), + (0x9b1, 0x9b1), + (0x9b3, 0x9b5), + (0x9ba, 0x9bb), + (0x9c5, 0x9c6), + (0x9c9, 0x9ca), + (0x9cf, 0x9d6), + (0x9d8, 0x9db), + (0x9de, 0x9de), + (0x9e4, 0x9e5), + (0x9ff, 0xa00), + (0xa04, 0xa04), + (0xa0b, 0xa0e), + (0xa11, 0xa12), + (0xa29, 0xa29), + (0xa31, 0xa31), + (0xa34, 0xa34), + (0xa37, 0xa37), + (0xa3a, 0xa3b), + (0xa3d, 0xa3d), + (0xa43, 0xa46), + (0xa49, 0xa4a), + (0xa4e, 0xa50), + (0xa52, 0xa58), + (0xa5d, 0xa5d), + (0xa5f, 0xa65), + (0xa77, 0xa80), + (0xa84, 0xa84), + (0xa8e, 0xa8e), + (0xa92, 0xa92), + (0xaa9, 0xaa9), + (0xab1, 0xab1), + (0xab4, 0xab4), + (0xaba, 0xabb), + (0xac6, 0xac6), + (0xaca, 0xaca), + (0xace, 0xacf), + (0xad1, 0xadf), + (0xae4, 0xae5), + (0xaf2, 0xaf8), + (0xb00, 0xb00), + (0xb04, 0xb04), + (0xb0d, 0xb0e), + (0xb11, 0xb12), + (0xb29, 0xb29), + (0xb31, 0xb31), + (0xb34, 0xb34), + (0xb3a, 0xb3b), + (0xb45, 0xb46), + (0xb49, 0xb4a), + (0xb4e, 0xb54), + (0xb58, 0xb5b), + (0xb5e, 0xb5e), + (0xb64, 0xb65), + (0xb78, 0xb81), + (0xb84, 0xb84), + (0xb8b, 0xb8d), + (0xb91, 0xb91), + (0xb96, 0xb98), + (0xb9b, 0xb9b), + (0xb9d, 0xb9d), + (0xba0, 0xba2), + (0xba5, 0xba7), + (0xbab, 0xbad), + (0xbba, 0xbbd), + (0xbc3, 0xbc5), + (0xbc9, 0xbc9), + (0xbce, 0xbcf), + (0xbd1, 0xbd6), + (0xbd8, 0xbe5), + (0xbfb, 0xbff), + (0xc0d, 0xc0d), + (0xc11, 0xc11), + (0xc29, 0xc29), + (0xc3a, 0xc3b), + (0xc45, 0xc45), + (0xc49, 0xc49), + (0xc4e, 0xc54), + (0xc57, 0xc57), + (0xc5b, 0xc5c), + (0xc5e, 0xc5f), + (0xc64, 0xc65), + (0xc70, 0xc76), + (0xc8d, 0xc8d), + (0xc91, 0xc91), + (0xca9, 0xca9), + (0xcb4, 0xcb4), + (0xcba, 0xcbb), + (0xcc5, 0xcc5), + (0xcc9, 0xcc9), + (0xcce, 0xcd4), + (0xcd7, 0xcdc), + (0xcdf, 0xcdf), + (0xce4, 0xce5), + (0xcf0, 0xcf0), + (0xcf4, 0xcff), + (0xd0d, 0xd0d), + (0xd11, 0xd11), + (0xd45, 0xd45), + (0xd49, 0xd49), + (0xd50, 0xd53), + (0xd64, 0xd65), + (0xd80, 0xd80), + (0xd84, 0xd84), + (0xd97, 0xd99), + (0xdb2, 0xdb2), + (0xdbc, 0xdbc), + (0xdbe, 0xdbf), + (0xdc7, 0xdc9), + (0xdcb, 0xdce), + (0xdd5, 0xdd5), + (0xdd7, 0xdd7), + (0xde0, 0xde5), + (0xdf0, 0xdf1), + (0xdf5, 0xe00), + (0xe3b, 0xe3e), + (0xe5c, 0xe80), + (0xe83, 0xe83), + (0xe85, 0xe85), + (0xe8b, 0xe8b), + (0xea4, 0xea4), + (0xea6, 0xea6), + (0xebe, 0xebf), + (0xec5, 0xec5), + (0xec7, 0xec7), + (0xecf, 0xecf), + (0xeda, 0xedb), + (0xee0, 0xeff), + (0xf48, 0xf48), + (0xf6d, 0xf70), + (0xf98, 0xf98), + (0xfbd, 0xfbd), + (0xfcd, 0xfcd), + (0xfdb, 0xfff), + (0x10c6, 0x10c6), + (0x10c8, 0x10cc), + (0x10ce, 0x10cf), + (0x1249, 0x1249), + (0x124e, 0x124f), + (0x1257, 0x1257), + (0x1259, 0x1259), + (0x125e, 0x125f), + (0x1289, 0x1289), + (0x128e, 0x128f), + (0x12b1, 0x12b1), + (0x12b6, 0x12b7), + (0x12bf, 0x12bf), + (0x12c1, 0x12c1), + (0x12c6, 0x12c7), + (0x12d7, 0x12d7), + (0x1311, 0x1311), + (0x1316, 0x1317), + (0x135b, 0x135c), + (0x137d, 0x137f), + (0x139a, 0x139f), + (0x13f6, 0x13f7), + (0x13fe, 0x13ff), + (0x1680, 0x1680), + (0x169d, 0x169f), + (0x16f9, 0x16ff), + (0x1716, 0x171e), + (0x1737, 0x173f), + (0x1754, 0x175f), + (0x176d, 0x176d), + (0x1771, 0x1771), + (0x1774, 0x177f), + (0x17de, 0x17df), + (0x17ea, 0x17ef), + (0x17fa, 0x17ff), + (0x180e, 0x180e), + (0x181a, 0x181f), + (0x1879, 0x187f), + (0x18ab, 0x18af), + (0x18f6, 0x18ff), + (0x191f, 0x191f), + (0x192c, 0x192f), + (0x193c, 0x193f), + (0x1941, 0x1943), + (0x196e, 0x196f), + (0x1975, 0x197f), + (0x19ac, 0x19af), + (0x19ca, 0x19cf), + (0x19db, 0x19dd), + (0x1a1c, 0x1a1d), + (0x1a5f, 0x1a5f), + (0x1a7d, 0x1a7e), + (0x1a8a, 0x1a8f), + (0x1a9a, 0x1a9f), + (0x1aae, 0x1aaf), + (0x1acf, 0x1aff), + (0x1b4d, 0x1b4f), + (0x1b7f, 0x1b7f), + (0x1bf4, 0x1bfb), + (0x1c38, 0x1c3a), + (0x1c4a, 0x1c4c), + (0x1c89, 0x1c8f), + (0x1cbb, 0x1cbc), + (0x1cc8, 0x1ccf), + (0x1cfb, 0x1cff), + (0x1f16, 0x1f17), + (0x1f1e, 0x1f1f), + (0x1f46, 0x1f47), + (0x1f4e, 0x1f4f), + (0x1f58, 0x1f58), + (0x1f5a, 0x1f5a), + (0x1f5c, 0x1f5c), + (0x1f5e, 0x1f5e), + (0x1f7e, 0x1f7f), + (0x1fb5, 0x1fb5), + (0x1fc5, 0x1fc5), + (0x1fd4, 0x1fd5), + (0x1fdc, 0x1fdc), + (0x1ff0, 0x1ff1), + (0x1ff5, 0x1ff5), + (0x1fff, 0x200f), + (0x2028, 0x202f), + (0x205f, 0x206f), + (0x2072, 0x2073), + (0x208f, 0x208f), + (0x209d, 0x209f), + (0x20c1, 0x20cf), + (0x20f1, 0x20ff), + (0x218c, 0x218f), + (0x2427, 0x243f), + (0x244b, 0x245f), + (0x2b74, 0x2b75), + (0x2b96, 0x2b96), + (0x2cf4, 0x2cf8), + (0x2d26, 0x2d26), + (0x2d28, 0x2d2c), + (0x2d2e, 0x2d2f), + (0x2d68, 0x2d6e), + (0x2d71, 0x2d7e), + (0x2d97, 0x2d9f), + (0x2da7, 0x2da7), + (0x2daf, 0x2daf), + (0x2db7, 0x2db7), + (0x2dbf, 0x2dbf), + (0x2dc7, 0x2dc7), + (0x2dcf, 0x2dcf), + (0x2dd7, 0x2dd7), + (0x2ddf, 0x2ddf), + (0x2e5e, 0x2e7f), + (0x2e9a, 0x2e9a), + (0x2ef4, 0x2eff), + (0x2fd6, 0x2fef), + (0x2ffc, 0x3000), + (0x3040, 0x3040), + (0x3097, 0x3098), + (0x3100, 0x3104), + (0x3130, 0x3130), + (0x318f, 0x318f), + (0x31e4, 0x31ef), + (0x321f, 0x321f), + (0xa48d, 0xa48f), + (0xa4c7, 0xa4cf), + (0xa62c, 0xa63f), + (0xa6f8, 0xa6ff), + (0xa7cb, 0xa7cf), + (0xa7d2, 0xa7d2), + (0xa7d4, 0xa7d4), + (0xa7da, 0xa7f1), + (0xa82d, 0xa82f), + (0xa83a, 0xa83f), + (0xa878, 0xa87f), + (0xa8c6, 0xa8cd), + (0xa8da, 0xa8df), + (0xa954, 0xa95e), + (0xa97d, 0xa97f), + (0xa9ce, 0xa9ce), + (0xa9da, 0xa9dd), + (0xa9ff, 0xa9ff), + (0xaa37, 0xaa3f), + (0xaa4e, 0xaa4f), + (0xaa5a, 0xaa5b), + (0xaac3, 0xaada), + (0xaaf7, 0xab00), + (0xab07, 0xab08), + (0xab0f, 0xab10), + (0xab17, 0xab1f), + (0xab27, 0xab27), + (0xab2f, 0xab2f), + (0xab6c, 0xab6f), + (0xabee, 0xabef), + (0xabfa, 0xabff), + (0xd7a4, 0xd7af), + (0xd7c7, 0xd7ca), + (0xd7fc, 0xf8ff), + (0xfa6e, 0xfa6f), + (0xfada, 0xfaff), + (0xfb07, 0xfb12), + (0xfb18, 0xfb1c), + (0xfb37, 0xfb37), + (0xfb3d, 0xfb3d), + (0xfb3f, 0xfb3f), + (0xfb42, 0xfb42), + (0xfb45, 0xfb45), + (0xfbc3, 0xfbd2), + (0xfd90, 0xfd91), + (0xfdc8, 0xfdce), + (0xfdd0, 0xfdef), + (0xfe1a, 0xfe1f), + (0xfe53, 0xfe53), + (0xfe67, 0xfe67), + (0xfe6c, 0xfe6f), + (0xfe75, 0xfe75), + (0xfefd, 0xff00), + (0xffbf, 0xffc1), + (0xffc8, 0xffc9), + (0xffd0, 0xffd1), + (0xffd8, 0xffd9), + (0xffdd, 0xffdf), + (0xffe7, 0xffe7), + (0xffef, 0xfffb), + (0xfffe, 0xffff), + (0x1000c, 0x1000c), + (0x10027, 0x10027), + (0x1003b, 0x1003b), + (0x1003e, 0x1003e), + (0x1004e, 0x1004f), + (0x1005e, 0x1007f), + (0x100fb, 0x100ff), + (0x10103, 0x10106), + (0x10134, 0x10136), + (0x1018f, 0x1018f), + (0x1019d, 0x1019f), + (0x101a1, 0x101cf), + (0x101fe, 0x1027f), + (0x1029d, 0x1029f), + (0x102d1, 0x102df), + (0x102fc, 0x102ff), + (0x10324, 0x1032c), + (0x1034b, 0x1034f), + (0x1037b, 0x1037f), + (0x1039e, 0x1039e), + (0x103c4, 0x103c7), + (0x103d6, 0x103ff), + (0x1049e, 0x1049f), + (0x104aa, 0x104af), + (0x104d4, 0x104d7), + (0x104fc, 0x104ff), + (0x10528, 0x1052f), + (0x10564, 0x1056e), + (0x1057b, 0x1057b), + (0x1058b, 0x1058b), + (0x10593, 0x10593), + (0x10596, 0x10596), + (0x105a2, 0x105a2), + (0x105b2, 0x105b2), + (0x105ba, 0x105ba), + (0x105bd, 0x105ff), + (0x10737, 0x1073f), + (0x10756, 0x1075f), + (0x10768, 0x1077f), + (0x10786, 0x10786), + (0x107b1, 0x107b1), + (0x107bb, 0x107ff), + (0x10806, 0x10807), + (0x10809, 0x10809), + (0x10836, 0x10836), + (0x10839, 0x1083b), + (0x1083d, 0x1083e), + (0x10856, 0x10856), + (0x1089f, 0x108a6), + (0x108b0, 0x108df), + (0x108f3, 0x108f3), + (0x108f6, 0x108fa), + (0x1091c, 0x1091e), + (0x1093a, 0x1093e), + (0x10940, 0x1097f), + (0x109b8, 0x109bb), + (0x109d0, 0x109d1), + (0x10a04, 0x10a04), + (0x10a07, 0x10a0b), + (0x10a14, 0x10a14), + (0x10a18, 0x10a18), + (0x10a36, 0x10a37), + (0x10a3b, 0x10a3e), + (0x10a49, 0x10a4f), + (0x10a59, 0x10a5f), + (0x10aa0, 0x10abf), + (0x10ae7, 0x10aea), + (0x10af7, 0x10aff), + (0x10b36, 0x10b38), + (0x10b56, 0x10b57), + (0x10b73, 0x10b77), + (0x10b92, 0x10b98), + (0x10b9d, 0x10ba8), + (0x10bb0, 0x10bff), + (0x10c49, 0x10c7f), + (0x10cb3, 0x10cbf), + (0x10cf3, 0x10cf9), + (0x10d28, 0x10d2f), + (0x10d3a, 0x10e5f), + (0x10e7f, 0x10e7f), + (0x10eaa, 0x10eaa), + (0x10eae, 0x10eaf), + (0x10eb2, 0x10efc), + (0x10f28, 0x10f2f), + (0x10f5a, 0x10f6f), + (0x10f8a, 0x10faf), + (0x10fcc, 0x10fdf), + (0x10ff7, 0x10fff), + (0x1104e, 0x11051), + (0x11076, 0x1107e), + (0x110bd, 0x110bd), + (0x110c3, 0x110cf), + (0x110e9, 0x110ef), + (0x110fa, 0x110ff), + (0x11135, 0x11135), + (0x11148, 0x1114f), + (0x11177, 0x1117f), + (0x111e0, 0x111e0), + (0x111f5, 0x111ff), + (0x11212, 0x11212), + (0x11242, 0x1127f), + (0x11287, 0x11287), + (0x11289, 0x11289), + (0x1128e, 0x1128e), + (0x1129e, 0x1129e), + (0x112aa, 0x112af), + (0x112eb, 0x112ef), + (0x112fa, 0x112ff), + (0x11304, 0x11304), + (0x1130d, 0x1130e), + (0x11311, 0x11312), + (0x11329, 0x11329), + (0x11331, 0x11331), + (0x11334, 0x11334), + (0x1133a, 0x1133a), + (0x11345, 0x11346), + (0x11349, 0x1134a), + (0x1134e, 0x1134f), + (0x11351, 0x11356), + (0x11358, 0x1135c), + (0x11364, 0x11365), + (0x1136d, 0x1136f), + (0x11375, 0x113ff), + (0x1145c, 0x1145c), + (0x11462, 0x1147f), + (0x114c8, 0x114cf), + (0x114da, 0x1157f), + (0x115b6, 0x115b7), + (0x115de, 0x115ff), + (0x11645, 0x1164f), + (0x1165a, 0x1165f), + (0x1166d, 0x1167f), + (0x116ba, 0x116bf), + (0x116ca, 0x116ff), + (0x1171b, 0x1171c), + (0x1172c, 0x1172f), + (0x11747, 0x117ff), + (0x1183c, 0x1189f), + (0x118f3, 0x118fe), + (0x11907, 0x11908), + (0x1190a, 0x1190b), + (0x11914, 0x11914), + (0x11917, 0x11917), + (0x11936, 0x11936), + (0x11939, 0x1193a), + (0x11947, 0x1194f), + (0x1195a, 0x1199f), + (0x119a8, 0x119a9), + (0x119d8, 0x119d9), + (0x119e5, 0x119ff), + (0x11a48, 0x11a4f), + (0x11aa3, 0x11aaf), + (0x11af9, 0x11aff), + (0x11b0a, 0x11bff), + (0x11c09, 0x11c09), + (0x11c37, 0x11c37), + (0x11c46, 0x11c4f), + (0x11c6d, 0x11c6f), + (0x11c90, 0x11c91), + (0x11ca8, 0x11ca8), + (0x11cb7, 0x11cff), + (0x11d07, 0x11d07), + (0x11d0a, 0x11d0a), + (0x11d37, 0x11d39), + (0x11d3b, 0x11d3b), + (0x11d3e, 0x11d3e), + (0x11d48, 0x11d4f), + (0x11d5a, 0x11d5f), + (0x11d66, 0x11d66), + (0x11d69, 0x11d69), + (0x11d8f, 0x11d8f), + (0x11d92, 0x11d92), + (0x11d99, 0x11d9f), + (0x11daa, 0x11edf), + (0x11ef9, 0x11eff), + (0x11f11, 0x11f11), + (0x11f3b, 0x11f3d), + (0x11f5a, 0x11faf), + (0x11fb1, 0x11fbf), + (0x11ff2, 0x11ffe), + (0x1239a, 0x123ff), + (0x1246f, 0x1246f), + (0x12475, 0x1247f), + (0x12544, 0x12f8f), + (0x12ff3, 0x12fff), + (0x13430, 0x1343f), + (0x13456, 0x143ff), + (0x14647, 0x167ff), + (0x16a39, 0x16a3f), + (0x16a5f, 0x16a5f), + (0x16a6a, 0x16a6d), + (0x16abf, 0x16abf), + (0x16aca, 0x16acf), + (0x16aee, 0x16aef), + (0x16af6, 0x16aff), + (0x16b46, 0x16b4f), + (0x16b5a, 0x16b5a), + (0x16b62, 0x16b62), + (0x16b78, 0x16b7c), + (0x16b90, 0x16e3f), + (0x16e9b, 0x16eff), + (0x16f4b, 0x16f4e), + (0x16f88, 0x16f8e), + (0x16fa0, 0x16fdf), + (0x16fe5, 0x16fef), + (0x16ff2, 0x16fff), + (0x187f8, 0x187ff), + (0x18cd6, 0x18cff), + (0x18d09, 0x1afef), + (0x1aff4, 0x1aff4), + (0x1affc, 0x1affc), + (0x1afff, 0x1afff), + (0x1b123, 0x1b131), + (0x1b133, 0x1b14f), + (0x1b153, 0x1b154), + (0x1b156, 0x1b163), + (0x1b168, 0x1b16f), + (0x1b2fc, 0x1bbff), + (0x1bc6b, 0x1bc6f), + (0x1bc7d, 0x1bc7f), + (0x1bc89, 0x1bc8f), + (0x1bc9a, 0x1bc9b), + (0x1bca0, 0x1ceff), + (0x1cf2e, 0x1cf2f), + (0x1cf47, 0x1cf4f), + (0x1cfc4, 0x1cfff), + (0x1d0f6, 0x1d0ff), + (0x1d127, 0x1d128), + (0x1d173, 0x1d17a), + (0x1d1eb, 0x1d1ff), + (0x1d246, 0x1d2bf), + (0x1d2d4, 0x1d2df), + (0x1d2f4, 0x1d2ff), + (0x1d357, 0x1d35f), + (0x1d379, 0x1d3ff), + (0x1d455, 0x1d455), + (0x1d49d, 0x1d49d), + (0x1d4a0, 0x1d4a1), + (0x1d4a3, 0x1d4a4), + (0x1d4a7, 0x1d4a8), + (0x1d4ad, 0x1d4ad), + (0x1d4ba, 0x1d4ba), + (0x1d4bc, 0x1d4bc), + (0x1d4c4, 0x1d4c4), + (0x1d506, 0x1d506), + (0x1d50b, 0x1d50c), + (0x1d515, 0x1d515), + (0x1d51d, 0x1d51d), + (0x1d53a, 0x1d53a), + (0x1d53f, 0x1d53f), + (0x1d545, 0x1d545), + (0x1d547, 0x1d549), + (0x1d551, 0x1d551), + (0x1d6a6, 0x1d6a7), + (0x1d7cc, 0x1d7cd), + (0x1da8c, 0x1da9a), + (0x1daa0, 0x1daa0), + (0x1dab0, 0x1deff), + (0x1df1f, 0x1df24), + (0x1df2b, 0x1dfff), + (0x1e007, 0x1e007), + (0x1e019, 0x1e01a), + (0x1e022, 0x1e022), + (0x1e025, 0x1e025), + (0x1e02b, 0x1e02f), + (0x1e06e, 0x1e08e), + (0x1e090, 0x1e0ff), + (0x1e12d, 0x1e12f), + (0x1e13e, 0x1e13f), + (0x1e14a, 0x1e14d), + (0x1e150, 0x1e28f), + (0x1e2af, 0x1e2bf), + (0x1e2fa, 0x1e2fe), + (0x1e300, 0x1e4cf), + (0x1e4fa, 0x1e7df), + (0x1e7e7, 0x1e7e7), + (0x1e7ec, 0x1e7ec), + (0x1e7ef, 0x1e7ef), + (0x1e7ff, 0x1e7ff), + (0x1e8c5, 0x1e8c6), + (0x1e8d7, 0x1e8ff), + (0x1e94c, 0x1e94f), + (0x1e95a, 0x1e95d), + (0x1e960, 0x1ec70), + (0x1ecb5, 0x1ed00), + (0x1ed3e, 0x1edff), + (0x1ee04, 0x1ee04), + (0x1ee20, 0x1ee20), + (0x1ee23, 0x1ee23), + (0x1ee25, 0x1ee26), + (0x1ee28, 0x1ee28), + (0x1ee33, 0x1ee33), + (0x1ee38, 0x1ee38), + (0x1ee3a, 0x1ee3a), + (0x1ee3c, 0x1ee41), + (0x1ee43, 0x1ee46), + (0x1ee48, 0x1ee48), + (0x1ee4a, 0x1ee4a), + (0x1ee4c, 0x1ee4c), + (0x1ee50, 0x1ee50), + (0x1ee53, 0x1ee53), + (0x1ee55, 0x1ee56), + (0x1ee58, 0x1ee58), + (0x1ee5a, 0x1ee5a), + (0x1ee5c, 0x1ee5c), + (0x1ee5e, 0x1ee5e), + (0x1ee60, 0x1ee60), + (0x1ee63, 0x1ee63), + (0x1ee65, 0x1ee66), + (0x1ee6b, 0x1ee6b), + (0x1ee73, 0x1ee73), + (0x1ee78, 0x1ee78), + (0x1ee7d, 0x1ee7d), + (0x1ee7f, 0x1ee7f), + (0x1ee8a, 0x1ee8a), + (0x1ee9c, 0x1eea0), + (0x1eea4, 0x1eea4), + (0x1eeaa, 0x1eeaa), + (0x1eebc, 0x1eeef), + (0x1eef2, 0x1efff), + (0x1f02c, 0x1f02f), + (0x1f094, 0x1f09f), + (0x1f0af, 0x1f0b0), + (0x1f0c0, 0x1f0c0), + (0x1f0d0, 0x1f0d0), + (0x1f0f6, 0x1f0ff), + (0x1f1ae, 0x1f1e5), + (0x1f203, 0x1f20f), + (0x1f23c, 0x1f23f), + (0x1f249, 0x1f24f), + (0x1f252, 0x1f25f), + (0x1f266, 0x1f2ff), + (0x1f6d8, 0x1f6db), + (0x1f6ed, 0x1f6ef), + (0x1f6fd, 0x1f6ff), + (0x1f777, 0x1f77a), + (0x1f7da, 0x1f7df), + (0x1f7ec, 0x1f7ef), + (0x1f7f1, 0x1f7ff), + (0x1f80c, 0x1f80f), + (0x1f848, 0x1f84f), + (0x1f85a, 0x1f85f), + (0x1f888, 0x1f88f), + (0x1f8ae, 0x1f8af), + (0x1f8b2, 0x1f8ff), + (0x1fa54, 0x1fa5f), + (0x1fa6e, 0x1fa6f), + (0x1fa7d, 0x1fa7f), + (0x1fa89, 0x1fa8f), + (0x1fabe, 0x1fabe), + (0x1fac6, 0x1facd), + (0x1fadc, 0x1fadf), + (0x1fae9, 0x1faef), + (0x1faf9, 0x1faff), + (0x1fb93, 0x1fb93), + (0x1fbcb, 0x1fbef), + (0x1fbfa, 0x1ffff), + (0x2a6e0, 0x2a6ff), + (0x2b73a, 0x2b73f), + (0x2b81e, 0x2b81f), + (0x2cea2, 0x2ceaf), + (0x2ebe1, 0x2f7ff), + (0x2fa1e, 0x2ffff), + (0x3134b, 0x3134f), + (0x323b0, 0xe00ff), + (0xe01f0, 0x10ffff), +]; diff --git a/src/models/cua_s1/native/src/pyjson.rs b/src/models/cua_s1/native/src/pyjson.rs new file mode 100644 index 00000000..4cda98e4 --- /dev/null +++ b/src/models/cua_s1/native/src/pyjson.rs @@ -0,0 +1,901 @@ +//! JSON the way the Python worker sees it. +//! +//! - [`parse`] follows CPython 3.12's `json.loads` together with the checks that +//! `contract.parse_body` adds: duplicate keys, `NaN`/`Infinity`, numbers that are +//! out of range for a float, lone surrogates and very deep nesting are all rejected, +//! and errors surface in the same order as in Python. +//! - [`dumps`] follows `json.dumps(value, ensure_ascii=False)` (and the compact form +//! Starlette uses for responses). +//! - [`repr`] follows Python's `repr`, which the error messages quote. + +use std::collections::HashSet; +use std::fmt::Write as _; + +use crate::printable::NON_PRINTABLE; + +/// Nesting limits of the Python worker, measured on CPython 3.12.13 under uvicorn. +/// Both come from interpreter recursion limits, so they depend on the call stack and +/// are not documented constants. +/// +/// Parsing fails when a container would open more than this many levels deep (the +/// `json` C scanner's recursion check). +pub const MAX_PARSE_DEPTH: usize = 9990; + +/// After parsing, `parse_body` walks the value with a recursive Python function +/// (`_check_unicode`), one call per value, containers and scalars alike. A walk that +/// needs more nested calls than this fails. +pub const MAX_CHECK_CALLS: usize = 969; + +/// `sys.int_info.default_max_str_digits`: longer integers fail to convert in Python. +const MAX_INT_DIGITS: usize = 4300; + +/// A string as Python holds it: a sequence of code points that may include lone +/// surrogates, stored as generalized UTF-8. +#[derive(Clone, Debug, Default, PartialEq, Eq, Hash)] +pub struct PyStr(Vec); + +impl PyStr { + pub fn new(s: &str) -> Self { + PyStr(s.as_bytes().to_vec()) + } + + fn push(&mut self, cp: u32) { + let b = &mut self.0; + match cp { + 0..=0x7f => b.push(cp as u8), + 0x80..=0x7ff => { + b.push(0xc0 | (cp >> 6) as u8); + b.push(0x80 | (cp & 0x3f) as u8); + } + 0x800..=0xffff => { + b.push(0xe0 | (cp >> 12) as u8); + b.push(0x80 | ((cp >> 6) & 0x3f) as u8); + b.push(0x80 | (cp & 0x3f) as u8); + } + _ => { + b.push(0xf0 | (cp >> 18) as u8); + b.push(0x80 | ((cp >> 12) & 0x3f) as u8); + b.push(0x80 | ((cp >> 6) & 0x3f) as u8); + b.push(0x80 | (cp & 0x3f) as u8); + } + } + } + + /// The text, or None if it holds a lone surrogate. + pub fn as_str(&self) -> Option<&str> { + std::str::from_utf8(&self.0).ok() + } + + pub fn code_points(&self) -> impl Iterator + '_ { + let b = &self.0; + let mut i = 0; + std::iter::from_fn(move || { + let lead = *b.get(i)? as u32; + let (len, init) = match lead { + 0..=0x7f => (1, lead), + 0xc0..=0xdf => (2, lead & 0x1f), + 0xe0..=0xef => (3, lead & 0x0f), + _ => (4, lead & 0x07), + }; + let cp = b[i + 1..i + len] + .iter() + .fold(init, |acc, &c| (acc << 6) | (c as u32 & 0x3f)); + i += len; + Some(cp) + }) + } +} + +impl PartialEq for PyStr { + fn eq(&self, other: &str) -> bool { + self.0 == other.as_bytes() + } +} + +#[derive(Clone, Debug, PartialEq)] +pub enum Value { + Null, + Bool(bool), + /// An integer as Python would print it (JSON integers never have leading zeros, + /// so this is the source text, with `-0` read as `0`). + Int(String), + Float(f64), + Str(PyStr), + Array(Vec), + /// Key order is kept; keys are unique once parsing succeeds. + Object(Vec<(PyStr, Value)>), +} + +impl Value { + pub fn get(&self, key: &str) -> Option<&Value> { + match self { + Value::Object(pairs) => pairs.iter().find(|(k, _)| k == key).map(|(_, v)| v), + _ => None, + } + } +} + +/// Values can nest thousands of levels deep before the depth check rejects them, so +/// they are dropped without recursion. +impl Drop for Value { + fn drop(&mut self) { + let mut stack: Vec = Vec::new(); + take_children(self, &mut stack); + while let Some(mut v) = stack.pop() { + take_children(&mut v, &mut stack); + } + } +} + +fn take_children(value: &mut Value, stack: &mut Vec) { + match value { + Value::Array(items) => stack.append(items), + Value::Object(pairs) => stack.extend(pairs.drain(..).map(|(_, v)| v)), + _ => {} + } +} + +/// Why a body was rejected as JSON; every case is a 400. +#[derive(Debug, PartialEq)] +pub enum JsonError { + Utf8, + Syntax, + Depth, + Duplicate(PyStr), + Constant(&'static str), + Range(String), + NotObject, +} + +impl JsonError { + /// The `detail` the Python worker returns for this error. + pub fn message(&self) -> String { + match self { + JsonError::Utf8 => "request body must be valid UTF-8 text".into(), + JsonError::Syntax => "request body must be valid JSON".into(), + JsonError::Depth => "request body is nested too deeply".into(), + JsonError::Duplicate(key) => { + format!("duplicate key {} in a JSON object", repr_str(key)) + } + JsonError::Constant(name) => format!("{name} is not valid JSON"), + JsonError::Range(text) => format!("number {text} is out of range"), + JsonError::NotObject => "request body must be a JSON object".into(), + } + } +} + +/// Decode a request body into its top-level object. +pub fn parse(raw: &[u8]) -> Result, JsonError> { + let text = std::str::from_utf8(raw).map_err(|_| JsonError::Utf8)?; + let mut p = Parser { + s: text.as_bytes(), + i: 0, + }; + p.ws(); + let mut value = p.value()?; + p.ws(); + if p.i != p.s.len() { + return Err(JsonError::Syntax); + } + check(&value, 1)?; + match &mut value { + Value::Object(pairs) => Ok(std::mem::take(pairs)), + _ => Err(JsonError::NotObject), + } +} + +/// `_check_unicode` from `contract.py`: visit values in order (an object's key before +/// its value) and fail at the first lone surrogate or the first call past +/// [`MAX_CHECK_CALLS`], whichever comes first. +fn check(value: &Value, calls: usize) -> Result<(), JsonError> { + if calls > MAX_CHECK_CALLS { + return Err(JsonError::Depth); + } + match value { + Value::Str(s) if s.as_str().is_none() => Err(JsonError::Utf8), + Value::Array(items) => items.iter().try_for_each(|v| check(v, calls + 1)), + Value::Object(pairs) => pairs.iter().try_for_each(|(k, v)| { + if k.as_str().is_none() { + return Err(JsonError::Utf8); + } + check(v, calls + 1) + }), + _ => Ok(()), + } +} + +/// An array or object that has been opened but not closed yet. +enum Open { + Array(Vec), + /// the members so far, and the key whose value is being parsed + Object(Vec<(PyStr, Value)>, PyStr), +} + +struct Parser<'a> { + s: &'a [u8], + i: usize, +} + +impl Parser<'_> { + fn peek(&self) -> Option { + self.s.get(self.i).copied() + } + + fn rest_starts_with(&self, lit: &[u8]) -> bool { + self.s[self.i..].starts_with(lit) + } + + fn ws(&mut self) { + while let Some(b' ' | b'\t' | b'\n' | b'\r') = self.peek() { + self.i += 1; + } + } + + /// One JSON value. Containers are parsed with an explicit stack, so deep nesting + /// cannot overflow the thread's stack before [`MAX_PARSE_DEPTH`] stops it. + fn value(&mut self) -> Result { + let mut open: Vec = Vec::new(); + loop { + // the start of a value: open a container or read a scalar + let mut done = match self.peek() { + Some(b'{') => { + if open.len() + 1 > MAX_PARSE_DEPTH { + return Err(JsonError::Depth); + } + self.i += 1; + self.ws(); + if self.peek() == Some(b'}') { + self.i += 1; + Value::Object(Vec::new()) + } else { + let key = self.key()?; + open.push(Open::Object(Vec::new(), key)); + continue; + } + } + Some(b'[') => { + if open.len() + 1 > MAX_PARSE_DEPTH { + return Err(JsonError::Depth); + } + self.i += 1; + self.ws(); + if self.peek() == Some(b']') { + self.i += 1; + Value::Array(Vec::new()) + } else { + open.push(Open::Array(Vec::new())); + continue; + } + } + _ => self.scalar()?, + }; + // hand the finished value to its container, closing containers that end + loop { + match open.last_mut() { + None => return Ok(done), + Some(Open::Array(items)) => { + items.push(done); + self.ws(); + match self.peek() { + Some(b',') => { + self.i += 1; + self.ws(); + break; + } + Some(b']') => { + self.i += 1; + let Some(Open::Array(items)) = open.pop() else { + unreachable!() + }; + done = Value::Array(items); + } + _ => return Err(JsonError::Syntax), + } + } + Some(Open::Object(pairs, key)) => { + pairs.push((std::mem::take(key), done)); + self.ws(); + match self.peek() { + Some(b',') => { + self.i += 1; + self.ws(); + *key = self.key()?; + break; + } + Some(b'}') => { + self.i += 1; + let Some(Open::Object(pairs, _)) = open.pop() else { + unreachable!() + }; + check_duplicates(&pairs)?; + done = Value::Object(pairs); + } + _ => return Err(JsonError::Syntax), + } + } + } + } + } + } + + /// A member name and its colon. + fn key(&mut self) -> Result { + if self.peek() != Some(b'"') { + return Err(JsonError::Syntax); + } + self.i += 1; + let key = self.string()?; + self.ws(); + if self.peek() != Some(b':') { + return Err(JsonError::Syntax); + } + self.i += 1; + self.ws(); + Ok(key) + } + + fn scalar(&mut self) -> Result { + match self.peek() { + Some(b'"') => { + self.i += 1; + Ok(Value::Str(self.string()?)) + } + Some(b'n') if self.rest_starts_with(b"null") => { + self.i += 4; + Ok(Value::Null) + } + Some(b't') if self.rest_starts_with(b"true") => { + self.i += 4; + Ok(Value::Bool(true)) + } + Some(b'f') if self.rest_starts_with(b"false") => { + self.i += 5; + Ok(Value::Bool(false)) + } + Some(b'N') if self.rest_starts_with(b"NaN") => Err(JsonError::Constant("NaN")), + Some(b'I') if self.rest_starts_with(b"Infinity") => { + Err(JsonError::Constant("Infinity")) + } + Some(b'-') if self.rest_starts_with(b"-Infinity") => { + Err(JsonError::Constant("-Infinity")) + } + _ => self.number(), + } + } + + fn digits(&mut self) { + while let Some(b'0'..=b'9') = self.peek() { + self.i += 1; + } + } + + fn number(&mut self) -> Result { + let start = self.i; + if self.peek() == Some(b'-') { + self.i += 1; + } + match self.peek() { + Some(b'0') => self.i += 1, + Some(b'1'..=b'9') => self.digits(), + _ => return Err(JsonError::Syntax), + } + let mut is_float = false; + if self.peek() == Some(b'.') && matches!(self.s.get(self.i + 1), Some(b'0'..=b'9')) { + self.i += 1; + self.digits(); + is_float = true; + } + if let Some(b'e' | b'E') = self.peek() { + let e_start = self.i; + self.i += 1; + if let Some(b'+' | b'-') = self.peek() { + self.i += 1; + } + let digits_start = self.i; + self.digits(); + if self.i > digits_start { + is_float = true; + } else { + // not an exponent after all; what follows is left for the caller + self.i = e_start; + } + } + let text = std::str::from_utf8(&self.s[start..self.i]).expect("ASCII"); + if is_float { + let v: f64 = text.parse().map_err(|_| JsonError::Syntax)?; + if !v.is_finite() { + return Err(JsonError::Range(text.to_string())); + } + Ok(Value::Float(v)) + } else { + let digits = text.strip_prefix('-').unwrap_or(text); + if digits.len() > MAX_INT_DIGITS { + return Err(JsonError::Syntax); + } + Ok(Value::Int(if digits == "0" { + "0".into() + } else { + text.into() + })) + } + } + + fn hex4(&self, at: usize) -> Option { + let h = self.s.get(at..at + 4)?; + let mut v = 0; + for &c in h { + v = v * 16 + (c as char).to_digit(16)?; + } + Some(v) + } + + /// The rest of a string whose opening quote has been read. + fn string(&mut self) -> Result { + let mut out = PyStr::default(); + loop { + match self.peek().ok_or(JsonError::Syntax)? { + b'"' => { + self.i += 1; + return Ok(out); + } + b'\\' => { + let esc = *self.s.get(self.i + 1).ok_or(JsonError::Syntax)?; + self.i += 2; + let cp = match esc { + b'"' => '"' as u32, + b'\\' => '\\' as u32, + b'/' => '/' as u32, + b'b' => 0x08, + b'f' => 0x0c, + b'n' => '\n' as u32, + b'r' => '\r' as u32, + b't' => '\t' as u32, + b'u' => { + let mut c = self.hex4(self.i).ok_or(JsonError::Syntax)?; + self.i += 4; + // a high surrogate joins a directly following low one + if (0xd800..=0xdbff).contains(&c) + && self.rest_starts_with(b"\\u") + && let Some(c2) = self.hex4(self.i + 2) + && (0xdc00..=0xdfff).contains(&c2) + { + c = 0x10000 + ((c - 0xd800) << 10) + (c2 - 0xdc00); + self.i += 6; + } + c + } + _ => return Err(JsonError::Syntax), + }; + out.push(cp); + } + 0x00..=0x1f => return Err(JsonError::Syntax), + _ => { + let start = self.i; + while let Some(c) = self.peek() { + if c == b'"' || c == b'\\' || c < 0x20 { + break; + } + self.i += 1; + } + out.0.extend_from_slice(&self.s[start..self.i]); + } + } + } + } +} + +/// Python's object_pairs_hook runs when an object closes and reports the first key +/// that repeats an earlier one. +fn check_duplicates(pairs: &[(PyStr, Value)]) -> Result<(), JsonError> { + let mut seen = HashSet::with_capacity(pairs.len()); + for (key, _) in pairs { + if !seen.insert(key) { + return Err(JsonError::Duplicate(key.clone())); + } + } + Ok(()) +} + +/// Python's `repr(float)`: the shortest digits that round-trip, in fixed notation +/// for exponents from -5 to 15 and scientific notation otherwise. +pub fn float_repr(x: f64) -> String { + if x == 0.0 { + return if x.is_sign_negative() { "-0.0" } else { "0.0" }.into(); + } + let sci = format!("{x:e}"); + let (mantissa, exp) = sci.split_once('e').expect("{:e} has an exponent"); + let exp: i32 = exp.parse().expect("integer exponent"); + let (neg, mantissa) = match mantissa.strip_prefix('-') { + Some(m) => (true, m), + None => (false, mantissa), + }; + let digits: String = mantissa.chars().filter(|c| *c != '.').collect(); + let (digits, exp) = break_tie_to_even(x, digits, exp); + let mut out = String::new(); + if neg { + out.push('-'); + } + let decpt = exp + 1; + if decpt <= -4 || decpt > 16 { + out.push_str(&digits[..1]); + if digits.len() > 1 { + out.push('.'); + out.push_str(&digits[1..]); + } + let sign = if exp < 0 { '-' } else { '+' }; + write!(out, "e{sign}{:02}", exp.unsigned_abs()).unwrap(); + } else if decpt <= 0 { + out.push_str("0."); + out.extend(std::iter::repeat_n('0', (-decpt) as usize)); + out.push_str(&digits); + } else if decpt as usize >= digits.len() { + out.push_str(&digits); + out.extend(std::iter::repeat_n('0', decpt as usize - digits.len())); + out.push_str(".0"); + } else { + out.push_str(&digits[..decpt as usize]); + out.push('.'); + out.push_str(&digits[decpt as usize..]); + } + out +} + +/// When `x` lies exactly halfway between two shortest round-tripping decimals, Python +/// (David Gay's dtoa) takes the one with an even last digit, while Rust's shortest +/// formatting may take the other. Ties need an exact decimal expansion only one +/// digit longer than the shortest, which takes 16 or more significant digits. +fn break_tie_to_even(x: f64, digits: String, exp: i32) -> (String, i32) { + let n = digits.len(); + if n < 15 { + return (digits, exp); + } + // every finite double has at most 767 significant decimal digits + let exact = format!("{:.800e}", x.abs()); + let (mantissa, e) = exact.split_once('e').expect("{:e} has an exponent"); + let e: i32 = e.parse().expect("integer exponent"); + let full: String = mantissa.chars().filter(|c| *c != '.').collect(); + let full = full.trim_end_matches('0'); + if full.len() != n + 1 || !full.ends_with('5') { + return (digits, exp); + } + let floor = &full[..n]; + let last = floor.as_bytes()[n - 1] - b'0'; + let even = if last.is_multiple_of(2) { + floor.to_string() + } else if last == 9 { + // the even choice would carry, which exact ties (values in [2^50, 2^51) + // ending in .25 or .75) never need + return (digits, exp); + } else { + // the even choice is floor + 1 in the last digit + let mut up = floor.as_bytes().to_vec(); + up[n - 1] += 1; + String::from_utf8(up).expect("ASCII digits") + }; + // At a power of two the double below is half as far away, so one of the two + // candidates may read back as a different double; then it is not a tie. + let back: f64 = format!("{}.{}e{}", &even[..1], &even[1..], e) + .parse() + .expect("decimal digits"); + if back.to_bits() != x.abs().to_bits() { + return (digits, exp); + } + (even, e) +} + +/// A JSON string literal as `json.dumps(s, ensure_ascii=False)` writes it. +pub fn write_json_str(s: &str, out: &mut String) { + out.push('"'); + for c in s.chars() { + match c { + '"' => out.push_str("\\\""), + '\\' => out.push_str("\\\\"), + '\n' => out.push_str("\\n"), + '\r' => out.push_str("\\r"), + '\t' => out.push_str("\\t"), + '\u{8}' => out.push_str("\\b"), + '\u{c}' => out.push_str("\\f"), + c if (c as u32) < 0x20 => write!(out, "\\u{:04x}", c as u32).unwrap(), + c => out.push(c), + } + } + out.push('"'); +} + +/// `json.dumps(value, ensure_ascii=False)`, with the default separators when +/// `compact` is false and `(",", ":")` when it is true. The value must not hold +/// lone surrogates (`parse` rejects them). +pub fn dumps(value: &Value, compact: bool) -> String { + let mut out = String::new(); + write_value(value, compact, &mut out); + out +} + +fn write_value(value: &Value, compact: bool, out: &mut String) { + let (item_sep, key_sep) = if compact { (",", ":") } else { (", ", ": ") }; + match value { + Value::Null => out.push_str("null"), + Value::Bool(b) => out.push_str(if *b { "true" } else { "false" }), + Value::Int(text) => out.push_str(text), + Value::Float(x) => out.push_str(&float_repr(*x)), + Value::Str(s) => write_json_str(s.as_str().expect("checked UTF-8"), out), + Value::Array(items) => { + out.push('['); + for (i, item) in items.iter().enumerate() { + if i > 0 { + out.push_str(item_sep); + } + write_value(item, compact, out); + } + out.push(']'); + } + Value::Object(pairs) => { + out.push('{'); + for (i, (key, item)) in pairs.iter().enumerate() { + if i > 0 { + out.push_str(item_sep); + } + write_json_str(key.as_str().expect("checked UTF-8"), out); + out.push_str(key_sep); + write_value(item, compact, out); + } + out.push('}'); + } + } +} + +fn is_printable(cp: u32) -> bool { + NON_PRINTABLE + .binary_search_by(|&(lo, hi)| { + if hi < cp { + std::cmp::Ordering::Less + } else if lo > cp { + std::cmp::Ordering::Greater + } else { + std::cmp::Ordering::Equal + } + }) + .is_err() +} + +/// Python's `repr(str)`. +pub fn repr_str(s: &PyStr) -> String { + let cps: Vec = s.code_points().collect(); + let squote = cps.contains(&('\'' as u32)); + let dquote = cps.contains(&('"' as u32)); + let quote = if squote && !dquote { '"' } else { '\'' }; + let mut out = String::with_capacity(cps.len() + 2); + out.push(quote); + for cp in cps { + match cp { + _ if cp == quote as u32 || cp == '\\' as u32 => { + out.push('\\'); + out.push(char::from_u32(cp).unwrap()); + } + 0x09 => out.push_str("\\t"), + 0x0a => out.push_str("\\n"), + 0x0d => out.push_str("\\r"), + 0..=0x1f | 0x7f => write!(out, "\\x{cp:02x}").unwrap(), + 0x20..=0x7e => out.push(char::from_u32(cp).unwrap()), + _ if is_printable(cp) => out.push(char::from_u32(cp).unwrap()), + 0x80..=0xff => write!(out, "\\x{cp:02x}").unwrap(), + 0x100..=0xffff => write!(out, "\\u{cp:04x}").unwrap(), + _ => write!(out, "\\U{cp:08x}").unwrap(), + } + } + out.push(quote); + out +} + +/// Python's `repr` of a decoded JSON value. +pub fn repr(value: &Value) -> String { + let mut out = String::new(); + write_repr(value, &mut out); + out +} + +fn write_repr(value: &Value, out: &mut String) { + match value { + Value::Null => out.push_str("None"), + Value::Bool(b) => out.push_str(if *b { "True" } else { "False" }), + Value::Int(text) => out.push_str(text), + Value::Float(x) => out.push_str(&float_repr(*x)), + Value::Str(s) => out.push_str(&repr_str(s)), + Value::Array(items) => { + out.push('['); + for (i, item) in items.iter().enumerate() { + if i > 0 { + out.push_str(", "); + } + write_repr(item, out); + } + out.push(']'); + } + Value::Object(pairs) => { + out.push('{'); + for (i, (key, item)) in pairs.iter().enumerate() { + if i > 0 { + out.push_str(", "); + } + out.push_str(&repr_str(key)); + out.push_str(": "); + write_repr(item, out); + } + out.push('}'); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn err(body: &str) -> String { + parse(body.as_bytes()).unwrap_err().message() + } + + #[test] + fn float_repr_matches_python_examples() { + let cases = [ + (1.0, "1.0"), + (1e16, "1e+16"), + (1e15, "1000000000000000.0"), + (1e-5, "1e-05"), + (1e-4, "0.0001"), + (-0.0, "-0.0"), + (1e22, "1e+22"), + (3.14e-07, "3.14e-07"), + (0.1, "0.1"), + (5e-324, "5e-324"), + (1.7976931348623157e308, "1.7976931348623157e+308"), + // 2^-24: a digit string halfway to the next decimal, but the lower one + // reads back as the double below + (5.960464477539063e-08, "5.960464477539063e-08"), + (123456789.0, "123456789.0"), + (0.0024726232513785362, "0.0024726232513785362"), + // exactly halfway between two 17-digit candidates: Python takes the even one + (f64::from_bits(0xc31d5a973d792fa1), "-2065594985630696.2"), + ]; + for (x, want) in cases { + assert_eq!(float_repr(x), want, "{x:e}"); + } + } + + #[test] + #[ignore = "needs CUA_S1_FLOAT_VECTORS"] + fn float_repr_matches_python_vectors() { + // tests/make_float_vectors.py writes " " lines + let path = std::env::var("CUA_S1_FLOAT_VECTORS") + .expect("CUA_S1_FLOAT_VECTORS must name a file from tests/make_float_vectors.py"); + let text = std::fs::read_to_string(path).unwrap(); + let mut n = 0; + for line in text.lines() { + let (bits, want) = line.split_once(' ').unwrap(); + let x = f64::from_bits(u64::from_str_radix(bits, 16).unwrap()); + assert_eq!(float_repr(x), want, "bits {bits}"); + n += 1; + } + assert!(n > 1000); + } + + #[test] + fn dumps_matches_python() { + let v = parse(br#"{"a": [1.0, 1e16, 1e-5, 0.0001, -0.0, 1e22, 123456789012345678, 3.14e-07, -0, true, null]}"#).unwrap(); + let v = Value::Object(v); + assert_eq!( + dumps(&v, false), + r#"{"a": [1.0, 1e+16, 1e-05, 0.0001, -0.0, 1e+22, 123456789012345678, 3.14e-07, 0, true, null]}"# + ); + let s = Value::Str(PyStr::new("\u{0}\u{1f}\u{7f}\u{2028}\"\\/\t\u{8}\u{c}é😀")); + assert_eq!( + dumps(&s, false), + "\"\\u0000\\u001f\u{7f}\u{2028}\\\"\\\\/\\t\\b\\fé😀\"" + ); + } + + #[test] + fn repr_matches_python() { + let s = PyStr::new("a\u{7f}\u{a0}\u{2028}😀é"); + assert_eq!(repr_str(&s), "'a\\x7f\\xa0\\u2028😀é'"); + assert_eq!(repr_str(&PyStr::new("it's")), "\"it's\""); + assert_eq!(repr_str(&PyStr::new("both'\"")), "'both\\'\"'"); + let v = Value::Object(parse(br#"{"k": [1, 2.5, null, true, "x"], "e": {}}"#).unwrap()); + assert_eq!(repr(&v), "{'k': [1, 2.5, None, True, 'x'], 'e': {}}"); + } + + #[test] + fn surrogates() { + let v = parse(b"{\"a\": \"\\ud83d\\ude00\"}").unwrap(); + assert_eq!(v[0].1, Value::Str(PyStr::new("😀"))); + assert_eq!( + err(r#"{"a": "\ud800x"}"#), + "request body must be valid UTF-8 text" + ); + assert_eq!( + err(r#"{"a": "\ude00\ud83d"}"#), + "request body must be valid UTF-8 text" + ); + // duplicate-key errors come first, and quote the surrogate + assert_eq!( + err(r#"{"\ud800": 1, "\ud800": 2}"#), + "duplicate key '\\ud800' in a JSON object" + ); + } + + #[test] + fn errors() { + assert_eq!(err("[]"), "request body must be a JSON object"); + assert_eq!(err("\u{feff}{}"), "request body must be valid JSON"); + assert_eq!(err(r#"{"a": NaN}"#), "NaN is not valid JSON"); + assert_eq!(err(r#"{"a": -Infinity}"#), "-Infinity is not valid JSON"); + assert_eq!(err(r#"{"a": 1e400}"#), "number 1e400 is out of range"); + assert_eq!(err(r#"{"a": 1e400x"#), "number 1e400 is out of range"); + assert_eq!( + err(r#"{"a": 1, "b": 2, "a": 3}"#), + "duplicate key 'a' in a JSON object" + ); + assert_eq!(err(r#"{"a": [1,]}"#), "request body must be valid JSON"); + assert_eq!(err("{\"a\": \"\u{1}\"}"), "request body must be valid JSON"); + assert_eq!(err(r#"{"a": 01}"#), "request body must be valid JSON"); + assert_eq!(err(r#"{"a": 1.}"#), "request body must be valid JSON"); + assert_eq!(err(r#"{"a": "\x"}"#), "request body must be valid JSON"); + assert_eq!(err(r#"{} x"#), "request body must be valid JSON"); + assert!(parse(br#"{"a": 1e-400}"#).is_ok()); + let long = format!("{{\"a\": {}}}", "1".repeat(4300)); + assert!(parse(long.as_bytes()).is_ok()); + let long = format!("{{\"a\": -{}}}", "1".repeat(4301)); + assert_eq!(err(&long), "request body must be valid JSON"); + assert_eq!(parse(b"{\"a\": \"\xff\"}").unwrap_err(), JsonError::Utf8); + } + + /// The Python worker's behaviour, measured over HTTP: `state` nested `d` lists deep + /// inside the top-level object. + #[test] + fn depth() { + let body = |d: usize, leaf: &str, tail: &str| { + format!( + "{{\"model\": \"m\", \"state\": {}{leaf}{}{tail}}}", + "[".repeat(d), + "]".repeat(d) + ) + }; + let deep = "request body is nested too deeply"; + // a scalar leaf takes one more call of the check than an empty container + assert!(parse(body(967, "\"x\"", "").as_bytes()).is_ok()); + assert_eq!(err(&body(968, "\"x\"", "")), deep); + assert!(parse(body(968, "", "").as_bytes()).is_ok()); + assert_eq!(err(&body(969, "", "")), deep); + assert!(parse(body(967, "{}", "").as_bytes()).is_ok()); + assert_eq!(err(&body(968, "{}", "")), deep); + // up to the parse limit, later parse errors win over depth + assert_eq!( + err(&(body(9989, "1", "") + "x")), + "request body must be valid JSON" + ); + assert_eq!(err(&(body(9990, "1", "") + "x")), deep); + assert_eq!( + err(&body(5000, "1", ", \"state\": 1")), + "duplicate key 'state' in a JSON object" + ); + // after parsing, the first problem in walk order wins + let deep_state = format!("{}1{}", "[".repeat(1500), "]".repeat(1500)); + assert_eq!( + err(&format!( + "{{\"model\": \"\\ud800\", \"state\": {deep_state}}}" + )), + "request body must be valid UTF-8 text" + ); + assert_eq!( + err(&format!("{{\"state\": {deep_state}, \"z\": \"\\ud800\"}}")), + deep + ); + assert_eq!( + err(&format!("{{\"a\": {deep_state}, \"\\ud800\": 1}}")), + deep + ); + // far past the limit: fails cleanly, without overflowing the stack + assert_eq!(err(&"[".repeat(1_000_000)), deep); + let wide = format!("{{\"a\": {}1{}}}", "[".repeat(9989), "]".repeat(9989)); + assert_eq!(err(&wide), deep); + } +} diff --git a/src/models/cua_s1/native/src/server.rs b/src/models/cua_s1/native/src/server.rs new file mode 100644 index 00000000..c087ac23 --- /dev/null +++ b/src/models/cua_s1/native/src/server.rs @@ -0,0 +1,234 @@ +//! HTTP routes, matching the Python worker: `GET /health` and `POST /v1/systemone`, +//! one decision at a time, the same status codes and the same response bytes. + +use std::sync::Arc; + +use axum::Router; +use axum::body::Body; +use axum::extract::State; +use axum::http::{HeaderMap, StatusCode, header}; +use axum::response::{IntoResponse, Response}; +use axum::routing::{get, post}; +use http_body_util::BodyExt; + +use crate::contract::{self, Request, RequestError, detail_json, map_request, parse_body}; +use crate::engine::{Engine, Prompter}; +use crate::pyjson::{PyStr, repr_str, write_json_str}; + +pub struct Limits { + pub max_body_bytes: usize, + pub max_questions: usize, + /// per question; 0 disables the check + pub max_prompt_tokens: usize, +} + +pub struct App { + pub engine: Engine, + pub limits: Limits, + /// `Bearer ` as raw bytes, when a key is set + expected_auth: Option>, + pub identity: String, + /// held for the whole decision, so forward passes run one at a time + turn: tokio::sync::Mutex<()>, +} + +impl App { + /// `api_key` is the raw value of `CUA_S1_API_KEY`; empty means no key. + pub fn new(engine: Engine, limits: Limits, api_key: Option>, revision: &str) -> Self { + let expected_auth = api_key + .filter(|k| !k.is_empty()) + .map(|k| [b"Bearer ".as_slice(), &k].concat()); + Self { + engine, + limits, + expected_auth, + identity: contract::model_identity(revision), + turn: tokio::sync::Mutex::new(()), + } + } +} + +fn json_response(status: StatusCode, body: String) -> Response { + (status, [(header::CONTENT_TYPE, "application/json")], body).into_response() +} + +fn error(status: u16, message: &str) -> Response { + json_response( + StatusCode::from_u16(status).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), + detail_json(message), + ) +} + +pub enum DecideError { + Request(RequestError), + Internal(anyhow::Error), +} + +/// Token ids for every question, checking the per-question prompt limit before any +/// forward pass runs. +pub fn encode_all( + prompter: &Prompter, + request: &Request, + max_prompt_tokens: usize, +) -> Result>, DecideError> { + let mut encoded = Vec::with_capacity(request.questions.len()); + for question in &request.questions { + let ids = prompter + .encode(&request.state, question) + .map_err(DecideError::Internal)?; + if max_prompt_tokens > 0 && ids.len() > max_prompt_tokens { + return Err(DecideError::Request(RequestError::new( + 413, + format!( + "question {}: prompt is {} tokens, over the {max_prompt_tokens}-token limit", + repr_str(&PyStr::new(&question.name)), + ids.len() + ), + ))); + } + encoded.push(ids); + } + Ok(encoded) +} + +/// Score each question and build the response body. +pub async fn decide(app: &App, request: &Request) -> Result { + let encoded = encode_all(&app.engine.prompter, request, app.limits.max_prompt_tokens)?; + let mut answers = String::from("{"); + let mut prompt_tokens = 0; + for (i, (question, ids)) in request.questions.iter().zip(encoded).enumerate() { + prompt_tokens += ids.len(); + let probs = app + .engine + .score(ids, question.keys.len()) + .await + .map_err(DecideError::Internal)?; + let answer = contract::answer_json(question, &probs) + .map_err(|e| DecideError::Internal(anyhow::anyhow!(e)))?; + if i > 0 { + answers.push(','); + } + write_json_str(&question.name, &mut answers); + answers.push(':'); + answers.push_str(&answer); + } + answers.push('}'); + let mut out = String::from("{\"model\":"); + write_json_str(&app.identity, &mut out); + out.push_str(",\"answers\":"); + out.push_str(&answers); + out.push_str(&format!( + ",\"usage\":{{\"input_tokens\":{prompt_tokens},\"output_tokens\":0}}}}" + )); + Ok(out) +} + +/// The Python worker reads the header as Latin-1 text (Starlette) and encodes it back +/// as UTF-8 before `hmac.compare_digest`; the same bytes are compared here, so both +/// workers accept and reject the same headers. +fn authorized(app: &App, headers: &HeaderMap) -> bool { + let Some(expected) = &app.expected_auth else { + return true; + }; + let raw = headers + .get(header::AUTHORIZATION) + .map(|v| v.as_bytes()) + .unwrap_or(b""); + let mut supplied = Vec::with_capacity(raw.len()); + for &b in raw { + if b < 0x80 { + supplied.push(b); + } else { + supplied.extend_from_slice(&[0xc0 | (b >> 6), 0x80 | (b & 0x3f)]); + } + } + supplied.len() == expected.len() + && supplied + .iter() + .zip(expected) + .fold(0u8, |acc, (a, b)| acc | (a ^ b)) + == 0 +} + +async fn health(State(app): State>) -> Response { + let mut out = String::from("{\"status\":\"ready\",\"modality\":\"text\",\"model\":"); + write_json_str(&app.identity, &mut out); + out.push_str(",\"device\":"); + write_json_str(&app.engine.device, &mut out); + out.push_str(",\"dtype\":\"bfloat16\",\"mode\":\"native\"}"); + json_response(StatusCode::OK, out) +} + +async fn systemone(State(app): State>, headers: HeaderMap, body: Body) -> Response { + if !authorized(&app, &headers) { + return error(401, "invalid or missing bearer token"); + } + let max = app.limits.max_body_bytes; + if let Some(len) = headers + .get(header::CONTENT_LENGTH) + .and_then(|v| v.to_str().ok()) + && !len.is_empty() + && len.bytes().all(|b| b.is_ascii_digit()) + && len.parse::().map_or(true, |n| n > max as u128) + { + return error(413, "request body too large"); + } + let mut raw = Vec::new(); + let mut body = body; + while let Some(frame) = body.frame().await { + let Ok(frame) = frame else { + return error(400, "request body could not be read"); + }; + if let Some(chunk) = frame.data_ref() { + raw.extend_from_slice(chunk); + if raw.len() > max { + return error(413, "request body too large"); + } + } + } + let request = match parse_body(&raw).and_then(|b| map_request(&b, app.limits.max_questions)) { + Ok(r) => r, + Err(e) => return error(e.status, &e.message), + }; + let _turn = app.turn.lock().await; + match decide(&app, &request).await { + Ok(body) => json_response(StatusCode::OK, body), + Err(DecideError::Request(e)) => error(e.status, &e.message), + Err(DecideError::Internal(e)) => { + eprintln!("inference failed: {e:#}"); + error(500, "inference failed") + } + } +} + +async fn not_found() -> Response { + error(404, "Not Found") +} + +async fn method_not_allowed() -> Response { + error(405, "Method Not Allowed") +} + +pub fn router(app: Arc) -> Router { + Router::new() + .route("/health", get(health)) + .route("/v1/systemone", post(systemone)) + .fallback(not_found) + .method_not_allowed_fallback(method_not_allowed) + .with_state(app) +} + +/// One decision through the whole request path, before the server listens. +pub async fn warmup(app: &App) -> anyhow::Result<()> { + let request = map_request( + &parse_body(contract::WARMUP_BODY.as_bytes()).map_err(|e| anyhow::anyhow!(e.message))?, + 64, + ) + .map_err(|e| anyhow::anyhow!(e.message))?; + let _turn = app.turn.lock().await; + match decide(app, &request).await { + Ok(_) => Ok(()), + Err(DecideError::Request(e)) => anyhow::bail!(e.message), + Err(DecideError::Internal(e)) => Err(e), + } +} diff --git a/src/models/cua_s1/native/tests/kernels.rs b/src/models/cua_s1/native/tests/kernels.rs new file mode 100644 index 00000000..733e17a7 --- /dev/null +++ b/src/models/cua_s1/native/tests/kernels.rs @@ -0,0 +1,261 @@ +//! GPU checks of the attention and Gated DeltaNet kernels on random inputs. They need +//! a GPU and CUA_S1_CUDA_LIB pointing at libqwen3_5_cuda.so, so they only run when +//! asked for: +//! +//! CUA_S1_CUDA_LIB=$PWD/target/release/libqwen3_5_cuda.so \ +//! cargo test --release -p omni-cua-s1-native --test kernels -- --ignored + +use std::path::PathBuf; + +use half::bf16; +use omni_cua_s1_native::cuda::{self, DeviceBuffer, Stream, api, check}; + +fn setup() -> Stream { + let lib = std::env::var_os("CUA_S1_CUDA_LIB") + .map(PathBuf::from) + .expect("CUA_S1_CUDA_LIB must point at libqwen3_5_cuda.so"); + cuda::load(&lib).unwrap(); + cuda::set_device(0).unwrap(); + cuda::new_stream().unwrap() +} + +/// Uniform values in [-amp, amp), rounded to bfloat16, from a fixed seed. +fn random(n: usize, seed: u64, amp: f32) -> Vec { + let mut x = seed.wrapping_mul(0x9e37_79b9_7f4a_7c15) | 1; + (0..n) + .map(|_| { + x ^= x << 13; + x ^= x >> 7; + x ^= x << 17; + bf16::from_f32(((x >> 40) as f32 / (1u64 << 24) as f32 * 2.0 - 1.0) * amp) + }) + .collect() +} + +fn to_device(v: &[bf16], st: Stream) -> DeviceBuffer { + let bytes: Vec = v.iter().flat_map(|x| x.to_le_bytes()).collect(); + let buf = DeviceBuffer::new(bytes.len()).unwrap(); + // SAFETY: the buffer was allocated for these bytes. + unsafe { cuda::upload(buf.at(0), &bytes, st).unwrap() }; + buf +} + +fn f32_to_device(v: &[f32], st: Stream) -> DeviceBuffer { + let bytes: Vec = v.iter().flat_map(|x| x.to_le_bytes()).collect(); + let buf = DeviceBuffer::new(bytes.len()).unwrap(); + // SAFETY: the buffer was allocated for these bytes. + unsafe { cuda::upload(buf.at(0), &bytes, st).unwrap() }; + buf +} + +fn from_device(buf: &DeviceBuffer, n: usize, st: Stream) -> Vec { + let mut bytes = vec![0u8; n * 2]; + // SAFETY: the buffer holds n bfloat16 values. + unsafe { cuda::download(&mut bytes, buf.at(0), st).unwrap() }; + bytes + .chunks_exact(2) + .map(|b| bf16::from_le_bytes([b[0], b[1]]).to_f32()) + .collect() +} + +#[test] +#[ignore = "needs a GPU and CUA_S1_CUDA_LIB"] +fn flash_attention_matches_float32_kernel() { + let st = setup(); + let (hq, hk, dh) = (16usize, 4usize, 256usize); + for (t, amp) in [ + (1, 8.0), + (63, 8.0), + (65, 0.5), + (139, 2.0), + (700, 8.0), + (2048, 0.5), + ] { + let q = to_device(&random(t * hq * dh, 1, amp), st); + let k = to_device(&random(t * hk * dh, 2, amp), st); + // v is read in place from the q|k|v projection output, rows of 10240 as in the model + let (ldv, v_at) = (10240usize, (hq * 2 + hk) * dh); + let qkv = to_device(&random(t * ldv, 3, 1.0), st); + let v = qkv.at(v_at * 2); + let flash = DeviceBuffer::new(t * hq * dh * 2).unwrap(); + let simple = DeviceBuffer::new(t * hq * dh * 2).unwrap(); + let (ti, hqi, hki, dhi, ldv) = (t as i32, hq as i32, hk as i32, dh as i32, ldv as i32); + // SAFETY: every buffer holds t rows of the given widths. + unsafe { + check( + (api().cs1_attention)( + q.at(0), + k.at(0), + v, + ldv, + flash.at(0), + ti, + hqi, + hki, + dhi, + 0.0625, + st, + ), + "flash", + ) + .unwrap(); + check( + (api().cs1_attention_simple)( + q.at(0), + k.at(0), + v, + ldv, + simple.at(0), + ti, + hqi, + hki, + dhi, + 0.0625, + st, + ), + "simple", + ) + .unwrap(); + } + let a = from_device(&flash, t * hq * dh, st); + let b = from_device(&simple, t * hq * dh, st); + // per (token, head): the largest difference over the largest magnitude + let mut worst = 0f32; + for (ra, rb) in a.chunks_exact(dh).zip(b.chunks_exact(dh)) { + let d = ra + .iter() + .zip(rb) + .map(|(x, y)| (x - y).abs()) + .fold(0f32, f32::max); + let m = rb.iter().map(|y| y.abs()).fold(1e-3f32, f32::max); + assert!( + ra.iter().all(|x| x.is_finite()), + "non-finite output at t = {t}" + ); + worst = worst.max(d / m); + } + eprintln!("attention t = {t}, amplitude {amp}: largest relative difference {worst:.2e}"); + assert!(worst < 1.6e-2, "t = {t}: {worst}"); + } +} + +/// Transformers' torch_recurrent_gated_delta_rule in float64, one token at a time, +/// with the L2 norms of q and k and q scaled by K^-1/2. +#[allow(clippy::too_many_arguments)] +fn gated_delta_reference( + q: &[bf16], + k: &[bf16], + v: &[bf16], + g: &[f32], + beta: &[bf16], + t: usize, + h: usize, + hk: usize, + d: usize, +) -> Vec { + let mut out = vec![0f64; t * h * d]; + for head in 0..h { + let kh = head / (h / hk); + let mut s = vec![0f64; d * d]; // [K][V] + for tok in 0..t { + let norm = |x: &[bf16]| { + let x: Vec = x.iter().map(|v| v.to_f64()).collect(); + let inv = 1.0 / (x.iter().map(|v| v * v).sum::() + 1e-6).sqrt(); + x.into_iter().map(|v| v * inv).collect::>() + }; + let qv: Vec = norm(&q[(tok * hk + kh) * d..][..d]) + .into_iter() + .map(|x| x / (d as f64).sqrt()) + .collect(); + let kv = norm(&k[(tok * hk + kh) * d..][..d]); + let vv: Vec = v[(tok * h + head) * d..][..d] + .iter() + .map(|x| x.to_f64()) + .collect(); + let decay = (g[tok * h + head] as f64).exp(); + let b = beta[tok * h + head].to_f64(); + s.iter_mut().for_each(|x| *x *= decay); + for j in 0..d { + let mem: f64 = (0..d).map(|i| kv[i] * s[i * d + j]).sum(); + let delta = (vv[j] - mem) * b; + for i in 0..d { + s[i * d + j] += kv[i] * delta; + } + } + for j in 0..d { + out[(tok * h + head) * d + j] = (0..d).map(|i| qv[i] * s[i * d + j]).sum(); + } + } + } + out +} + +#[test] +#[ignore = "needs a GPU and CUA_S1_CUDA_LIB"] +fn gated_delta_rule_matches_recurrent_reference() { + let st = setup(); + let (h, hk, d) = (4usize, 2usize, 128usize); + for t in [1usize, 64, 150] { + // q close to k, so that q.k and the outputs are of order one as in the model + let k = random(t * hk * d, 12, 1.0); + let q: Vec = k + .iter() + .zip(random(t * hk * d, 11, 1.0)) + .map(|(k, n)| bf16::from_f32(0.8 * k.to_f32() + 0.2 * n.to_f32())) + .collect(); + let v = random(t * h * d, 13, 1.0); + // log decays in (-2, 0) and learning rates in (0, 1), as sigmoid and -exp * softplus give + let g: Vec = random(t * h, 14, 1.0) + .iter() + .map(|x| x.to_f32() - 1.0) + .collect(); + let beta: Vec = random(t * h, 15, 0.5) + .iter() + .map(|x| bf16::from_f32(x.to_f32() + 0.5)) + .collect(); + let want = gated_delta_reference(&q, &k, &v, &g, &beta, t, h, hk, d); + let (qd, kd, vd, gd, bd) = ( + to_device(&q, st), + to_device(&k, st), + to_device(&v, st), + f32_to_device(&g, st), + to_device(&beta, st), + ); + let o = DeviceBuffer::new(t * h * d * 2).unwrap(); + // SAFETY: pure function of its arguments. + let floats = unsafe { (api().cs1_gdn_workspace_floats)(t as i32, h as i32) }; + let ws = DeviceBuffer::new(floats * 4).unwrap(); + // SAFETY: every buffer holds t rows of the given widths, the workspace its size. + unsafe { + check( + (api().cs1_gdn_prefill)( + qd.at(0), + kd.at(0), + vd.at(0), + gd.at(0).cast::(), + bd.at(0), + o.at(0), + ws.at(0).cast::(), + t as i32, + h as i32, + hk as i32, + (d as f32).powf(-0.5), + st, + ), + "gdn prefill", + ) + .unwrap(); + } + let got = from_device(&o, t * h * d, st); + let scale = want.iter().fold(0f64, |m, x| m.max(x.abs())); + let worst = got + .iter() + .zip(&want) + .map(|(a, b)| (*a as f64 - b).abs()) + .fold(0f64, f64::max); + eprintln!( + "gated delta t = {t}: largest difference {worst:.2e}, largest |reference| {scale:.2}" + ); + assert!(worst <= 2e-2 * scale, "t = {t}: {worst} vs scale {scale}"); + } +} diff --git a/src/models/cua_s1/native/tests/make_float_vectors.py b/src/models/cua_s1/native/tests/make_float_vectors.py new file mode 100644 index 00000000..655b489b --- /dev/null +++ b/src/models/cua_s1/native/tests/make_float_vectors.py @@ -0,0 +1,45 @@ +"""Write float formatting cases for the pyjson tests: " " lines. + + python3 src/models/cua_s1/native/tests/make_float_vectors.py floats.txt + CUA_S1_FLOAT_VECTORS=floats.txt cargo test -p omni-cua-s1-native float_repr -- --ignored + +Special values, powers of two and ten, values next to the points where repr switches +between plain and exponent notation, random bit patterns, and random short decimals. +""" + +import math +import random +import struct +import sys + + +def bits(x: float) -> str: + return struct.pack(">d", x).hex() + + +def main() -> None: + rng = random.Random(20260928) + values = [0.0, -0.0, 1.0, -1.0, 0.5, 0.1, 0.2, 0.3, 1 / 3, 2 / 3, math.pi, math.e, + 5e-324, 2.2250738585072014e-308, 1.7976931348623157e308, 2.0**53, 2.0**53 + 2] + values += [2.0**e for e in range(-1074, 1024, 7)] + values += [10.0**e for e in range(-323, 309)] + for e in (-5, -4, 15, 16, 17): + base = 10.0**e + values += [math.nextafter(base, 0.0), base, math.nextafter(base, math.inf)] + while len(values) < 30000: + pick = rng.random() + if pick < 0.5: + x = struct.unpack(">d", rng.getrandbits(64).to_bytes(8, "big"))[0] + elif pick < 0.8: + x = float(f"{rng.randint(1, 10**rng.randint(1, 17))}e{rng.randint(-30, 30)}") + else: + x = rng.uniform(-1e6, 1e6) + if math.isfinite(x): + values.append(x) + with open(sys.argv[1], "w") as f: + for x in values: + f.write(f"{bits(x)} {x!r}\n") + + +if __name__ == "__main__": + main() diff --git a/src/models/cua_s1/native/tests/make_printable.py b/src/models/cua_s1/native/tests/make_printable.py new file mode 100644 index 00000000..63fb6fd4 --- /dev/null +++ b/src/models/cua_s1/native/tests/make_printable.py @@ -0,0 +1,27 @@ +"""Write src/printable.rs: the code points >= 0x80 that Python's repr escapes. + + python3.12 src/models/cua_s1/native/tests/make_printable.py > src/models/cua_s1/native/src/printable.rs + +Use the Python version the reference worker runs on; the table follows its Unicode +database. +""" + +import sys +import unicodedata + +ranges = [] +for cp in range(0x80, sys.maxunicode + 1): + if not chr(cp).isprintable(): + if ranges and ranges[-1][1] == cp - 1: + ranges[-1][1] = cp + else: + ranges.append([cp, cp]) +version = ".".join(map(str, sys.version_info[:3])) +print(f"//! Generated from Python {version} (Unicode {unicodedata.unidata_version}): code points >= 0x80 for which") +print("//! `str.isprintable()` is false, as inclusive ranges. Python's `repr` escapes these.") +print("//! Regenerate with tests/make_printable.py.") +print() +print("pub const NON_PRINTABLE: &[(u32, u32)] = &[") +for lo, hi in ranges: + print(f" ({lo:#x}, {hi:#x}),") +print("];") From 637ced148bee2b96b5e9f2e83126e8c973092dd7 Mon Sep 17 00:00:00 2001 From: Tianyao Wu Date: Mon, 28 Sep 2026 22:44:01 +0800 Subject: [PATCH 05/10] cua_s1: tie GEMM plans to their cuBLASLt version Cs1GemmPlan now records the cuBLASLt version a plan was tuned with, and import refuses a plan from another version, so the check no longer depends on the caller (ABI version 2; the plans file format is unchanged). A shape that was not tuned borrows a tuned algorithm only from an M at most twice its own, or from the largest tuned M when nothing larger was tuned, and otherwise takes cuBLASLt's first choice. An exhaustive tune fails when cuBLASLt cannot list its algorithms instead of timing only the shortlist. Signed-off-by: Tianyao Wu --- src/backends/cuda/qwen3_5/gemm.cu | 41 ++++++++++++------- src/backends/cuda/qwen3_5/ops.h | 12 +++--- src/models/cua_s1/native/src/cuda.rs | 3 +- src/models/cua_s1/native/src/model.rs | 10 +++-- src/models/cua_s1/native/tests/kernels.rs | 48 ++++++++++++++++++++++- 5 files changed, 88 insertions(+), 26 deletions(-) diff --git a/src/backends/cuda/qwen3_5/gemm.cu b/src/backends/cuda/qwen3_5/gemm.cu index 22f0add6..3818be61 100644 --- a/src/backends/cuda/qwen3_5/gemm.cu +++ b/src/backends/cuda/qwen3_5/gemm.cu @@ -13,10 +13,10 @@ // different ones with about the same speed. // It replaces the heuristic's first choice only when it is more than 3% faster, so // near ties rarely change between runs, and -// cs1_gemm_export / cs1_gemm_import let a caller keep the choices across runs. A shape -// that was not tuned uses the algorithm tuned for the nearest M with the same N, K -// and ldy (the smallest tuned M above it, else the largest below), or the heuristic's -// first choice. Split-K reductions that accumulate into the output in place are +// cs1_gemm_export / cs1_gemm_import let a caller keep the choices across runs, for the +// cuBLASLt version they were tuned with. A shape that was not tuned borrows the +// algorithm tuned for a nearby M with the same N, K and ldy (see plan_for), or takes +// the heuristic's first choice. Split-K reductions that accumulate into the output in place are // excluded, since their order, and so the rounding, is not fixed. #include @@ -142,12 +142,11 @@ T cap(const cublasLtMatmulAlgo_t& algo, cublasLtMatmulAlgoCapAttributes_t attr) // algorithm id with each tile, stage count, custom option and swizzle it supports, and // split-K factors from `splits` (reduced in the compute or the output type, not in // place). Other attributes stay at their defaults. -std::vector every_config(Gemm& g, const Plan& p) { +int every_config(Gemm& g, const Plan& p, std::vector& out) { int ids[256], nids = 0; - std::vector out; - if (cublasLtMatmulAlgoGetIds(g.handle, CUBLAS_COMPUTE_32F, CUDA_R_32F, CUDA_R_16BF, CUDA_R_16BF, CUDA_R_16BF, - CUDA_R_16BF, 256, ids, &nids) != CUBLAS_STATUS_SUCCESS) - return out; + const cublasStatus_t s = cublasLtMatmulAlgoGetIds(g.handle, CUBLAS_COMPUTE_32F, CUDA_R_32F, CUDA_R_16BF, + CUDA_R_16BF, CUDA_R_16BF, CUDA_R_16BF, 256, ids, &nids); + if (s != CUBLAS_STATUS_SUCCESS) return status(s); const int splits[] = {1, 2, 3, 4, 5, 6, 8, 12, 16}; const uint32_t schemes[] = {CUBLASLT_REDUCTION_SCHEME_NONE, CUBLASLT_REDUCTION_SCHEME_COMPUTE_TYPE, CUBLASLT_REDUCTION_SCHEME_OUTPUT_TYPE}; @@ -190,7 +189,7 @@ std::vector every_config(Gemm& g, const Plan& p) { } } } - return out; + return 0; } // The plan for a shape, created on first use. @@ -207,7 +206,9 @@ int plan_for(Gemm& g, int M, int N, int K, int ldy, Plan*& out) { destroy(p); return rc; } - // The algorithm tuned for the nearest M: the smallest above, else the largest below. + // Borrow a tuned algorithm: the one for the smallest tuned M above, if that M is at + // most twice this one; the one for the largest tuned M below, if no larger M was + // tuned; else take the heuristic's first choice. const Plan* above = nullptr; const Plan* below = nullptr; int above_m = 0, below_m = 0; @@ -217,9 +218,9 @@ int plan_for(Gemm& g, int M, int N, int K, int ldy, Plan*& out) { if (m > M && (!above || m < above_m)) above = &kv.second, above_m = m; if (m < M && (!below || m > below_m)) below = &kv.second, below_m = m; } - if (above && usable(g, p, above->algo)) { + if (above && above_m <= 2 * M && usable(g, p, above->algo)) { p.algo = above->algo; - } else if (below && usable(g, p, below->algo)) { + } else if (!above && below && usable(g, p, below->algo)) { p.algo = below->algo; } else { std::vector cands; @@ -304,7 +305,13 @@ extern "C" int cs1_gemm_tune(void* gemm, const void* x, const void* w, void* y, std::vector cands; for (auto& r : shortlist) cands.push_back(r.algo); if (exhaustive) { - for (auto& a : every_config(*g, p)) + std::vector all; + rc = every_config(*g, p, all); + if (rc != 0) { + destroy(p); + return rc; + } + for (auto& a : all) if (std::none_of(cands.begin(), cands.end(), [&](const cublasLtMatmulAlgo_t& c) { return std::memcmp(&c, &a, sizeof a) == 0; })) cands.push_back(a); @@ -388,7 +395,7 @@ extern "C" size_t cs1_gemm_export(void* gemm, Cs1GemmPlan* out, size_t cap) { if (!kv.second.tuned) continue; if (n < cap) { const auto [m, nn, k, l] = kv.first; - out[n] = Cs1GemmPlan{m, nn, k, l, {}}; + out[n] = Cs1GemmPlan{m, nn, k, l, (uint64_t)cublasLtGetVersion(), {}}; static_assert(sizeof(cublasLtMatmulAlgo_t) == sizeof(out[n].algo), "algo layout"); std::memcpy(out[n].algo, &kv.second.algo, sizeof(out[n].algo)); } @@ -409,6 +416,10 @@ extern "C" int cs1_gemm_import(void* gemm, const Cs1GemmPlan* plans, size_t n) { rc = cudaErrorInvalidValue; break; } + if (r.cublaslt_version != (uint64_t)cublasLtGetVersion()) { + rc = status(CUBLAS_STATUS_NOT_SUPPORTED); + break; + } Plan p; rc = describe(r.m, r.n, r.k, r.ldy, p); if (rc == 0) { diff --git a/src/backends/cuda/qwen3_5/ops.h b/src/backends/cuda/qwen3_5/ops.h index 70ce39ca..e625686c 100644 --- a/src/backends/cuda/qwen3_5/ops.h +++ b/src/backends/cuda/qwen3_5/ops.h @@ -13,7 +13,7 @@ #include // Bumped whenever a signature below changes. -#define CS1_ABI_VERSION 1 +#define CS1_ABI_VERSION 2 #ifdef __cplusplus extern "C" { @@ -97,9 +97,11 @@ int cs1_silu_mul(const void* gate_up, int ld, void* out, int T, int I, void* str // cs1_gemm_tune picks the algorithm for one shape by timing, among the heuristic's // shortlist or (exhaustive) a wider enumeration (see gemm.cu); it must not // run during stream capture. cs1_gemm_tune_done frees the buffers tuning used. -// A tuned algorithm for one shape; `algo` holds a cublasLtMatmulAlgo_t. +// A tuned algorithm for one shape: `algo` holds a cublasLtMatmulAlgo_t, valid for the +// cuBLASLt version (cublasLtGetVersion) it was tuned with. typedef struct { int32_t m, n, k, ldy; + uint64_t cublaslt_version; uint64_t algo[8]; } Cs1GemmPlan; @@ -111,9 +113,9 @@ void cs1_gemm_tune_done(void* gemm); // Copy up to `cap` tuned plans to `out`; returns how many there are. size_t cs1_gemm_export(void* gemm, Cs1GemmPlan* out, size_t cap); // Use these plans (from cs1_gemm_export, possibly of an earlier run): all of them, or -// none if one fails cuBLASLt's check on this device or reduces split-K in place. -// The check does not tell whether a plan was tuned on this GPU and cuBLASLt version; -// the caller keeps that with the plans. +// none if one was tuned with another cuBLASLt version, fails cuBLASLt's check on this +// device, or reduces split-K in place. Whether a plan was tuned on this GPU model is +// not checked; the caller keeps that with the plans. int cs1_gemm_import(void* gemm, const Cs1GemmPlan* plans, size_t n); // cublasLtGetVersion(). size_t cs1_gemm_version(void); diff --git a/src/models/cua_s1/native/src/cuda.rs b/src/models/cua_s1/native/src/cuda.rs index d6e8c47e..769f3208 100644 --- a/src/models/cua_s1/native/src/cuda.rs +++ b/src/models/cua_s1/native/src/cuda.rs @@ -9,7 +9,7 @@ use std::sync::OnceLock; use anyhow::{Context, Result, bail, ensure}; /// `CS1_ABI_VERSION` in ops.h. -const ABI_VERSION: u32 = 1; +const ABI_VERSION: u32 = 2; pub const LIBRARY: &str = "libqwen3_5_cuda.so"; /// A `cudaStream_t`. @@ -29,6 +29,7 @@ pub struct GemmPlan { pub n: i32, pub k: i32, pub ldy: i32, + pub cublaslt_version: u64, pub algo: [u64; 8], } diff --git a/src/models/cua_s1/native/src/model.rs b/src/models/cua_s1/native/src/model.rs index 625810ff..d4c77d06 100644 --- a/src/models/cua_s1/native/src/model.rs +++ b/src/models/cua_s1/native/src/model.rs @@ -9,7 +9,7 @@ //! eager pass of that length, so both give bitwise identical results. //! //! GEMM algorithms are tuned at startup for the lengths in TUNE_ROWS; other lengths -//! use the nearest tuned one (see gemm.cu). The choices can be saved to a file and +//! borrow a nearby tuned one (see gemm.cu). The choices can be saved to a file and //! reused, so that restarts do not change them; the file records the GPU, the //! cuBLASLt version and the tuned lengths, and one that does not match is refused. //! @@ -49,8 +49,10 @@ fn plan_setup(graph_max_tokens: usize) -> Result { })) } -/// Prompt lengths the GEMM algorithms are tuned for. Past these, cuBLASLt's first -/// choice for long prompts is an older, half-rate tensor-core kernel on sm_89. +/// Prompt lengths the GEMM algorithms are tuned for, at most twice apart, so that a +/// length up to the last one borrows a tuned algorithm for at most twice its length +/// and longer ones borrow the last one's. Past these, cuBLASLt's first choice for long +/// prompts is an older, half-rate tensor-core kernel on sm_89. pub const TUNE_ROWS: &[usize] = &[ 64, 96, 128, 160, 192, 224, 256, 320, 384, 448, 512, 640, 768, 1024, 1536, 2048, 4096, 8192, 16384, @@ -796,6 +798,8 @@ impl Model { n: int(&p["n"])?, k: int(&p["k"])?, ldy: int(&p["ldy"])?, + // the file's setup, checked above, records the version + cublaslt_version: setup["cublaslt"].as_u64().context("no cuBLASLt version")?, algo, }) }) diff --git a/src/models/cua_s1/native/tests/kernels.rs b/src/models/cua_s1/native/tests/kernels.rs index 733e17a7..83664dab 100644 --- a/src/models/cua_s1/native/tests/kernels.rs +++ b/src/models/cua_s1/native/tests/kernels.rs @@ -1,4 +1,5 @@ -//! GPU checks of the attention and Gated DeltaNet kernels on random inputs. They need +//! GPU checks of the attention and Gated DeltaNet kernels on random inputs, and of the +//! GEMM plan import. They need //! a GPU and CUA_S1_CUDA_LIB pointing at libqwen3_5_cuda.so, so they only run when //! asked for: //! @@ -8,7 +9,7 @@ use std::path::PathBuf; use half::bf16; -use omni_cua_s1_native::cuda::{self, DeviceBuffer, Stream, api, check}; +use omni_cua_s1_native::cuda::{self, DeviceBuffer, GemmPlan, Stream, api, check}; fn setup() -> Stream { let lib = std::env::var_os("CUA_S1_CUDA_LIB") @@ -259,3 +260,46 @@ fn gated_delta_rule_matches_recurrent_reference() { assert!(worst <= 2e-2 * scale, "t = {t}: {worst} vs scale {scale}"); } } + +/// A tuned GEMM plan moves to another cuBLASLt handle through export and import, and a +/// plan that says it was tuned with another cuBLASLt version is refused without +/// changing anything. +#[test] +#[ignore = "needs a GPU and CUA_S1_CUDA_LIB"] +fn gemm_plans_import_only_for_their_cublaslt_version() { + let st = setup(); + let (m, n, k) = (64i32, 256i32, 512i32); + let x = to_device(&random((m * k) as usize, 21, 1.0), st); + let w = to_device(&random((n * k) as usize, 22, 1.0), st); + let y = DeviceBuffer::new((m * n * 2) as usize).unwrap(); + // SAFETY: the buffers hold m x k, n x k and m x n values; the handles are destroyed + // at the end and not used after. + unsafe { + let api = api(); + let (a, b) = ( + (api.cs1_gemm_create)(32 << 20), + (api.cs1_gemm_create)(32 << 20), + ); + assert!(!a.is_null() && !b.is_null()); + check( + (api.cs1_gemm_tune)(a, x.at(0), w.at(0), y.at(0), m, n, k, n, 0, st), + "tune", + ) + .unwrap(); + (api.cs1_gemm_tune_done)(a); + let count = (api.cs1_gemm_export)(a, std::ptr::null_mut(), 0); + assert_eq!(count, 1); + let mut plans = vec![GemmPlan::default(); count]; + (api.cs1_gemm_export)(a, plans.as_mut_ptr(), count); + assert_eq!(plans[0].cublaslt_version, (api.cs1_gemm_version)() as u64); + + let mut other = plans.clone(); + other[0].cublaslt_version += 1; + assert_ne!((api.cs1_gemm_import)(b, other.as_ptr(), 1), 0); + assert_eq!((api.cs1_gemm_export)(b, std::ptr::null_mut(), 0), 0); + check((api.cs1_gemm_import)(b, plans.as_ptr(), 1), "import").unwrap(); + assert_eq!((api.cs1_gemm_export)(b, std::ptr::null_mut(), 0), 1); + (api.cs1_gemm_destroy)(a); + (api.cs1_gemm_destroy)(b); + } +} From 6104d17cc0d879c726002be64f6a8ac3e91e00e6 Mon Sep 17 00:00:00 2001 From: Tianyao Wu Date: Tue, 29 Sep 2026 13:47:05 +0800 Subject: [PATCH 06/10] Keep the Cua-S1 native files clean under Rust 1.98 clippy and ruff Rust 1.98 adds clippy's chunks_exact_to_as_chunks, which fails the three bfloat16 decodes under -D warnings; they now use as_chunks. The recipe and test scripts pass ruff check (E4, E7, E9, F, I) and ruff format, as the multimodal workflow runs them over recipe/cua_s1. The reformatted scripts produce the same corpus, float vectors and printable table as before. Signed-off-by: Tianyao Wu --- recipe/cua_s1/check_native.py | 29 +- recipe/cua_s1/diff_corpus.py | 279 ++++++++++++++---- recipe/cua_s1/diff_workers.py | 77 ++++- src/models/cua_s1/native/src/engine.rs | 6 +- src/models/cua_s1/native/src/model.rs | 7 +- src/models/cua_s1/native/tests/kernels.rs | 7 +- .../cua_s1/native/tests/make_float_vectors.py | 25 +- .../cua_s1/native/tests/make_printable.py | 8 +- 8 files changed, 341 insertions(+), 97 deletions(-) diff --git a/recipe/cua_s1/check_native.py b/recipe/cua_s1/check_native.py index e6732054..2c239773 100644 --- a/recipe/cua_s1/check_native.py +++ b/recipe/cua_s1/check_native.py @@ -12,24 +12,31 @@ prompt fits) must be bitwise identical to the eager one. With a second file (another process start), both runs must be bitwise identical too. """ + import json import sys from pathlib import Path def load_parity(path): - rows = [json.loads(l) for l in path.read_text().splitlines() if l] + rows = [json.loads(line) for line in path.read_text().splitlines() if line] return {(r["case"], r["question"]): r["worker"] for r in rows} def load(path): - return [json.loads(l) for l in Path(path).read_text().splitlines() if l] + return [json.loads(line) for line in Path(path).read_text().splitlines() if line] + + +def diff(a, b): + return max(abs(a[k] - b[k]) for k in b) runs = load(sys.argv[1]) parity = Path(sys.argv[2]) -fp32, bf16 = load_parity(parity / "parity_float32.jsonl"), load_parity(parity / "parity_bfloat16.jsonl") -diff = lambda a, b: max(abs(a[k] - b[k]) for k in b) +fp32, bf16 = ( + load_parity(parity / "parity_float32.jsonl"), + load_parity(parity / "parity_bfloat16.jsonl"), +) allowance = 2 * max(diff(bf16[k], fp32[k]) for k in fp32) + 0.01 eager = {(r["case"], r["question"]): r for r in runs if r["mode"] == "eager"} served = {(r["case"], r["question"]): r for r in runs if r["mode"] == "served"} @@ -45,10 +52,18 @@ def load(path): if margin >= 0.05 and max(ref, key=ref.get) != max(got, key=got.get): flips.append(f"{key[0]}/{key[1]}") graph_keys = [k for k, r in served.items() if r["graph"]] -mismatch = [f"{k[0]}/{k[1]}" for k in served if served[k]["probabilities"] != eager[k]["probabilities"]] +mismatch = [ + f"{k[0]}/{k[1]}" + for k in served + if served[k]["probabilities"] != eager[k]["probabilities"] +] print(f"{len(eager)} questions; allowance {allowance:.4f}") -print(f"largest |eager - fp32| {worst:.4f} ({worst_at}); top-option changes: {flips or 'none'}; missing: {missing or 'none'}") -print(f"served from a graph: {len(graph_keys)}; served differs from eager: {mismatch or 'none'}") +print( + f"largest |eager - fp32| {worst:.4f} ({worst_at}); top-option changes: {flips or 'none'}; missing: {missing or 'none'}" +) +print( + f"served from a graph: {len(graph_keys)}; served differs from eager: {mismatch or 'none'}" +) ok = worst <= allowance and not flips and not missing and not mismatch if len(sys.argv) > 3: other = load(sys.argv[3]) diff --git a/recipe/cua_s1/diff_corpus.py b/recipe/cua_s1/diff_corpus.py index e9343a94..635ccef7 100644 --- a/recipe/cua_s1/diff_corpus.py +++ b/recipe/cua_s1/diff_corpus.py @@ -26,7 +26,9 @@ def req(state="Button: OK", questions=None, model=M, **extra): def choice(instructions="Press OK.", criteria=None, type_="choice"): q = {"type": type_, "instructions": instructions} - q["criteria"] = criteria if criteria is not None else {"ok": "OK", "cancel": "Cancel"} + q["criteria"] = ( + criteria if criteria is not None else {"ok": "OK", "cancel": "Cancel"} + ) return q @@ -43,44 +45,77 @@ def add(name, body): fixtures = json.load(open(inputs_path, encoding="utf-8")) for name, body in fixtures.items(): add(f"fixture/{name}", body) - add(f"fixture_ascii_compact/{name}", json.dumps(body, ensure_ascii=True, separators=(",", ":"))) + add( + f"fixture_ascii_compact/{name}", + json.dumps(body, ensure_ascii=True, separators=(",", ":")), + ) # --- JSON syntax and decoding --- ok = json.dumps(req()) raw = { - "empty": "", "space": " ", "open": "{", "close": "}", "array": "[]", "null": "null", - "number": "1", "string": '"x"', "true": "true", "extra": ok + " x", "bom": "" + ok, - "ws": " \t\n" + ok + "\r\n", "trailing_comma_obj": ok[:-1] + ",}", - "single_quotes": ok.replace('"', "'"), "comment": "// c\n" + ok, - "nan": ok.replace('"Button: OK"', "NaN"), "inf": ok.replace('"Button: OK"', "Infinity"), - "neg_inf": ok.replace('"Button: OK"', "-Infinity"), "nan_nested": ok.replace('"OK"', "[1, NaN]"), - "float_range": ok.replace('"OK"', "[1e400]"), "float_range_neg": ok.replace('"OK"', "[-1E309]"), + "empty": "", + "space": " ", + "open": "{", + "close": "}", + "array": "[]", + "null": "null", + "number": "1", + "string": '"x"', + "true": "true", + "extra": ok + " x", + "bom": "" + ok, + "ws": " \t\n" + ok + "\r\n", + "trailing_comma_obj": ok[:-1] + ",}", + "single_quotes": ok.replace('"', "'"), + "comment": "// c\n" + ok, + "nan": ok.replace('"Button: OK"', "NaN"), + "inf": ok.replace('"Button: OK"', "Infinity"), + "neg_inf": ok.replace('"Button: OK"', "-Infinity"), + "nan_nested": ok.replace('"OK"', "[1, NaN]"), + "float_range": ok.replace('"OK"', "[1e400]"), + "float_range_neg": ok.replace('"OK"', "[-1E309]"), "float_range_then_garbage": ok.replace('"OK"', "[1e400x]"), "float_underflow": ok.replace('"OK"', "[1e-400, 1.0e+308, -0.0, 5e-324]"), "int_4300": ok.replace('"OK"', "[" + "9" * 4300 + "]"), "int_4301": ok.replace('"OK"', "[" + "9" * 4301 + "]"), "int_neg_4301": ok.replace('"OK"', "[-" + "9" * 4301 + "]"), - "leading_zero": ok.replace('"OK"', "[01]"), "dot_no_digit": ok.replace('"OK"', "[1.]"), - "exp_no_digit": ok.replace('"OK"', "[1e]"), "exp_sign_no_digit": ok.replace('"OK"', "[1e+]"), - "minus_alone": ok.replace('"OK"', "[-]"), "plus": ok.replace('"OK"', "[+1]"), + "leading_zero": ok.replace('"OK"', "[01]"), + "dot_no_digit": ok.replace('"OK"', "[1.]"), + "exp_no_digit": ok.replace('"OK"', "[1e]"), + "exp_sign_no_digit": ok.replace('"OK"', "[1e+]"), + "minus_alone": ok.replace('"OK"', "[-]"), + "plus": ok.replace('"OK"', "[+1]"), "hex": ok.replace('"OK"', "[0x10]"), - "raw_control": ok.replace("Press OK.", "Press\u0001OK."), "raw_tab": ok.replace("Press OK.", "Press\tOK."), + "raw_control": ok.replace("Press OK.", "Press\u0001OK."), + "raw_tab": ok.replace("Press OK.", "Press\tOK."), "raw_del": ok.replace("Press OK.", "Press\u007fOK."), - "bad_escape": ok.replace("Press OK.", "Press \\x OK."), "short_u": ok.replace("Press OK.", "\\u12"), + "bad_escape": ok.replace("Press OK.", "Press \\x OK."), + "short_u": ok.replace("Press OK.", "\\u12"), "bad_u": ok.replace("Press OK.", "\\uZZZZ"), - "escapes": ok.replace("Press OK.", "\\/\\b\\f\\n\\r\\t\\\"\\\\ \\u00e9\\u0000"), - "missing_colon": ok.replace('"model":', '"model"'), "unterminated": ok[:-3], - "dup_top": '{"model": "' + M + '", "model": "' + M + '", "state": "s", "questions": {}}', + "escapes": ok.replace("Press OK.", '\\/\\b\\f\\n\\r\\t\\"\\\\ \\u00e9\\u0000'), + "missing_colon": ok.replace('"model":', '"model"'), + "unterminated": ok[:-3], + "dup_top": '{"model": "' + + M + + '", "model": "' + + M + + '", "state": "s", "questions": {}}', "dup_criteria": ok.replace('"cancel": "Cancel"', '"ok": "Again"'), "dup_then_syntax": ok.replace('"cancel": "Cancel"', '"ok": "Again"') + "x", - "syntax_inside_dup_object": ok.replace('"cancel": "Cancel"', '"ok": "Again", x'), + "syntax_inside_dup_object": ok.replace( + '"cancel": "Cancel"', '"ok": "Again", x' + ), "dup_state_nested": ok.replace('"Button: OK"', '{"a": {"b": 1, "b": 2}}'), "dup_after_nan": ok.replace('"cancel": "Cancel"', '"ok": NaN'), - "lone_high": ok.replace("Press OK.", "\\ud800"), "lone_low": ok.replace("Press OK.", "\\udc00x"), - "pair": ok.replace("Press OK.", "\\ud83d\\ude00"), "reversed_pair": ok.replace("Press OK.", "\\ude00\\ud83d"), + "lone_high": ok.replace("Press OK.", "\\ud800"), + "lone_low": ok.replace("Press OK.", "\\udc00x"), + "pair": ok.replace("Press OK.", "\\ud83d\\ude00"), + "reversed_pair": ok.replace("Press OK.", "\\ude00\\ud83d"), "high_then_bmp": ok.replace("Press OK.", "\\ud800\\u0041"), "lone_in_key": ok.replace('"ok": "OK"', '"\\udfff": "OK"'), - "lone_dup_key": ok.replace('"ok": "OK", "cancel": "Cancel"', '"\\ud800": 1, "\\ud800": 2'), + "lone_dup_key": ok.replace( + '"ok": "OK", "cancel": "Cancel"', '"\\ud800": 1, "\\ud800": 2' + ), "lone_then_syntax": ok.replace("Press OK.", "\\ud800") + ",", } for name, text in raw.items(): @@ -99,7 +134,10 @@ def nested(d, leaf='"x"', open_="[", close="]"): for d in (964, 965, 966): add(f"depth/criteria_{d}", ok.replace('"OK"', nested(d))) add(f"depth/type_{d}", ok.replace('"choice"', nested(d))) - add(f"depth/instructions_obj_{d}", ok.replace('"Press OK."', nested(d, '"x"', '{"a":', "}"))) + add( + f"depth/instructions_obj_{d}", + ok.replace('"Press OK."', nested(d, '"x"', '{"a":', "}")), + ) head = '{"model": "cua-s1-4b-0.2", "state": ' tail = ', "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "A"}}}}' @@ -111,40 +149,80 @@ def nested(d, leaf='"x"', open_="[", close="]"): add(f"depth2/parse_limit_garbage_{d}", head + nested(d, "1") + tail + "x") add("depth2/deep_then_dup", head + nested(5000, "1") + ', "state": 1' + tail) add("depth2/deep_then_nan", head + nested(3000, "1") + ', "z": NaN' + tail) - add("depth2/surrogate_before_deep", '{"model": "\\ud800", "state": ' + nested(1500, "1") + tail) - add("depth2/deep_before_surrogate", head + nested(1500, "1") + ', "z": "\\ud800"' + tail) - add("depth2/deep_before_surrogate_key", '{"a": ' + nested(1500, "1") + ', "\\ud800": 1}') + add( + "depth2/surrogate_before_deep", + '{"model": "\\ud800", "state": ' + nested(1500, "1") + tail, + ) + add( + "depth2/deep_before_surrogate", + head + nested(1500, "1") + ', "z": "\\ud800"' + tail, + ) + add( + "depth2/deep_before_surrogate_key", + '{"a": ' + nested(1500, "1") + ', "\\ud800": 1}', + ) add("depth2/hundred_thousand", "[" * 100000) add("depth2/deep_top_level_array", nested(5000, "1")) - sa = {"_sample": choice(criteria={"_save": "Save", "_sa": None, "b": "B"}), "q": choice()} + sa = { + "_sample": choice(criteria={"_save": "Save", "_sa": None, "b": "B"}), + "q": choice(), + } add("keys/_sa_names_and_options", req(questions=sa)) add("keys/_sa_state_keys", req(state={"_sa": 1, "_sample": [2]})) # --- mapping errors, in Python's check order --- sem = { "no_model": {"state": "s", "questions": {"q": choice()}}, - "model_null": req(model=None), "model_case": req(model="Cua-S1-4B-0.2"), - "model_space": req(model=M + " "), "model_number": req(model=1), + "model_null": req(model=None), + "model_case": req(model="Cua-S1-4B-0.2"), + "model_space": req(model=M + " "), + "model_number": req(model=1), "model_wrong_and_no_state": {"model": "x"}, "no_state": {"model": M, "questions": {"q": choice()}}, - "state_null": req(state=None), "state_true": req(state=True), "state_zero": req(state=0), - "state_float": req(state=1.5), "state_empty": req(state=""), "state_empty_obj": req(state={}), - "state_empty_list": req(state=[]), "state_spaces": req(state=" "), - "state_obj": req(state={"a": None, "b": [1, 2.5, -0.0, 1e22, 12345678901234567890, True], "é": "
"}), + "state_null": req(state=None), + "state_true": req(state=True), + "state_zero": req(state=0), + "state_float": req(state=1.5), + "state_empty": req(state=""), + "state_empty_obj": req(state={}), + "state_empty_list": req(state=[]), + "state_spaces": req(state=" "), + "state_obj": req( + state={ + "a": None, + "b": [1, 2.5, -0.0, 1e22, 12345678901234567890, True], + "é": "
", + } + ), "state_list": req(state=["Button: OK", {"k": "v"}, 3.14e-07]), - "no_questions": {"model": M, "state": "s"}, "questions_null": req(questions=None) | {"questions": None}, - "questions_list": req(questions=[]), "questions_empty": req(questions={}), + "no_questions": {"model": M, "state": "s"}, + "questions_null": req(questions=None) | {"questions": None}, + "questions_list": req(questions=[]), + "questions_empty": req(questions={}), "questions_str": req(questions="q"), "questions_65": req(questions={f"q{i}": choice() for i in range(65)}), - "questions_65_bad_types": req(questions={f"q{i}": {"type": "score"} for i in range(65)}), - "questions_12": req(questions={f"q{i}": choice(f"Pick {i}.") for i in range(12)}), - "question_list": req(questions={"q": []}), "question_null": req(questions={"q": None}), + "questions_65_bad_types": req( + questions={f"q{i}": {"type": "score"} for i in range(65)} + ), + "questions_12": req( + questions={f"q{i}": choice(f"Pick {i}.") for i in range(12)} + ), + "question_list": req(questions={"q": []}), + "question_null": req(questions={"q": None}), "question_str": req(questions={"q": "choice"}), - "score_second": req(questions={"a": {"type": "choice"}, "b": {"type": "score"}}), + "score_second": req( + questions={"a": {"type": "choice"}, "b": {"type": "score"}} + ), "noul": req(questions={"a": {"type": "noul"}}), - "type_missing": req(questions={"a": {"instructions": "x", "criteria": {"a": "A"}}}), - "missing_instructions_then_score": req(questions={"a": {"type": "choice"}, "b": {"type": "noul"}}), - "no_instructions": req(questions={"a": {"type": "choice", "criteria": {"x": "X"}}}), + "type_missing": req( + questions={"a": {"instructions": "x", "criteria": {"a": "A"}}} + ), + "missing_instructions_then_score": req( + questions={"a": {"type": "choice"}, "b": {"type": "noul"}} + ), + "no_instructions": req( + questions={"a": {"type": "choice", "criteria": {"x": "X"}}} + ), "instructions_null": req(questions={"a": choice(None)}), "instructions_empty": req(questions={"a": choice("")}), "instructions_zero": req(questions={"a": choice(0)}), @@ -153,14 +231,40 @@ def nested(d, leaf='"x"', open_="[", close="]"): "instructions_empty_obj": req(questions={"a": choice({})}), "instructions_list": req(questions={"a": choice([])}), "no_criteria": req(questions={"a": {"type": "choice", "instructions": "x"}}), - "criteria_null": req(questions={"a": choice(criteria=None) | {"criteria": None}}), + "criteria_null": req( + questions={"a": choice(criteria=None) | {"criteria": None}} + ), "criteria_list": req(questions={"a": choice(criteria=[])}), "criteria_empty": req(questions={"a": choice(criteria={})}), - "criteria_26": req(questions={"a": choice(criteria={f"o{i}": f"Option {i}" for i in range(26)})}), - "criteria_27": req(questions={"a": choice(criteria={f"o{i}": f"Option {i}" for i in range(27)})}), - "criteria_values": req(questions={"a": choice(criteria={ - "n": None, "o": {"k": [1, 2]}, "l": [], "e": {}, "q": 'quote"s', "b": "back\\slash", - "nl": "new\nline", "t": "tab\t", "z": "\u0000", "ls": "
", "é": "ünï 中文 😀"})}), + "criteria_26": req( + questions={ + "a": choice(criteria={f"o{i}": f"Option {i}" for i in range(26)}) + } + ), + "criteria_27": req( + questions={ + "a": choice(criteria={f"o{i}": f"Option {i}" for i in range(27)}) + } + ), + "criteria_values": req( + questions={ + "a": choice( + criteria={ + "n": None, + "o": {"k": [1, 2]}, + "l": [], + "e": {}, + "q": 'quote"s', + "b": "back\\slash", + "nl": "new\nline", + "t": "tab\t", + "z": "\u0000", + "ls": "
", + "é": "ünï 中文 😀", + } + ) + } + ), "criteria_bool": req(questions={"a": choice(criteria={"x": "X", "y": True})}), "criteria_number": req(questions={"a": choice(criteria={"x": 1})}), "second_question_bad": req(questions={"a": choice(), "b": choice(criteria={})}), @@ -168,27 +272,71 @@ def nested(d, leaf='"x"', open_="[", close="]"): for name, body in sem.items(): add(f"map/{name}", body) - weird_names = ["it's", 'say "hi"', "both'\"", "back\\slash", "new\nline", "nul\u0000", "del\u007f", - "nbsp ", "ls
", "zw​", "tag\U000e0001", "pua", "unassigned͸", - "emoji😀", "é", "combining é", "rtl א", "soft­", "space ", "", "\t"] + weird_names = [ + "it's", + 'say "hi"', + "both'\"", + "back\\slash", + "new\nline", + "nul\u0000", + "del\u007f", + "nbsp ", + "ls
", + "zw​", + "tag\U000e0001", + "pua", + "unassigned͸", + "emoji😀", + "é", + "combining é", + "rtl א", + "soft­", + "space ", + "", + "\t", + ] for i, name in enumerate(weird_names): add(f"names/type_{i}", req(questions={name: {"type": name}})) add(f"names/option_{i}", req(questions={name: choice(criteria={name: 1})})) - add(f"names/ok_{i}", req(questions={name: choice(criteria={name: None, "other": "Other"})})) - for i, kind in enumerate(["", "Choice", 1, 1.5, 1e16, -0.0, 1e-5, True, None, [], {}, {"a": [1, {"b": None}]}, - 12345678901234567890, 2.5e-310]): + add( + f"names/ok_{i}", + req(questions={name: choice(criteria={name: None, "other": "Other"})}), + ) + for i, kind in enumerate( + [ + "", + "Choice", + 1, + 1.5, + 1e16, + -0.0, + 1e-5, + True, + None, + [], + {}, + {"a": [1, {"b": None}]}, + 12345678901234567890, + 2.5e-310, + ] + ): add(f"types/{i}", req(questions={"q": {"type": kind}})) # --- prompt length --- long_state = "Row: item " * 3000 # a little over 16384 tokens add("limit/prompt_too_long", req(state=long_state)) - add("limit/second_prompt_too_long", req(questions={"a": choice(), "b": choice("x " * 17000)})) + add( + "limit/second_prompt_too_long", + req(questions={"a": choice(), "b": choice("x " * 17000)}), + ) add("limit/prompt_just_under", req(state="Row: item " * 2600)) # --- seeded random bodies --- rnd = random.Random(0) - pool = ("abcXYZ019 _-.:/\\\"'\n\t\r{}[]<>|é中文😀
​ ́א\u0000\u001f\u007f" - "<|im_start|><|im_end|>") + pool = ( + "abcXYZ019 _-.:/\\\"'\n\t\r{}[]<>|é中文😀
​ ́א\u0000\u001f\u007f" + "<|im_start|><|im_end|>" + ) def rstr(n=12): return "".join(rnd.choice(pool) for _ in range(rnd.randint(0, n))) @@ -196,7 +344,7 @@ def rstr(n=12): def rnum(): k = rnd.random() if k < 0.3: - return rnd.randint(-10**rnd.randint(1, 30), 10**rnd.randint(1, 30)) + return rnd.randint(-(10 ** rnd.randint(1, 30)), 10 ** rnd.randint(1, 30)) if k < 0.9: return rnd.uniform(-1, 1) * 10 ** rnd.randint(-320, 300) return rnd.choice([0.0, -0.0, 5e-324, 1e16, 1e-5, 0.1, 1 / 3]) @@ -216,9 +364,15 @@ def rval(depth=0): for i in range(n_fuzz): qs = {} for _ in range(rnd.randint(1, 3)): - crit = {rstr(8) or "k": rnd.choice([rstr(), None, rval(), rval()]) for _ in range(rnd.randint(1, 6))} - qs[rstr(8)] = {"type": "choice" if rnd.random() < 0.9 else rval(), - "instructions": rnd.choice([rstr(40), None, rval(), ""]), "criteria": crit} + crit = { + rstr(8) or "k": rnd.choice([rstr(), None, rval(), rval()]) + for _ in range(rnd.randint(1, 6)) + } + qs[rstr(8)] = { + "type": "choice" if rnd.random() < 0.9 else rval(), + "instructions": rnd.choice([rstr(40), None, rval(), ""]), + "criteria": crit, + } body = req(state=rnd.choice([rstr(200), rval(), rval()]), questions=qs) text = json.dumps(body, ensure_ascii=rnd.random() < 0.3) add(f"fuzz/{i}", text) @@ -241,7 +395,10 @@ def main() -> None: cases = build(sys.argv[1], n_fuzz) with open(sys.argv[2], "w") as f: for name, body in cases: - f.write(json.dumps({"name": name, "body": base64.b64encode(body).decode()}) + "\n") + f.write( + json.dumps({"name": name, "body": base64.b64encode(body).decode()}) + + "\n" + ) print(f"{len(cases)} bodies") diff --git a/recipe/cua_s1/diff_workers.py b/recipe/cua_s1/diff_workers.py index 19ce7706..976e6ae3 100644 --- a/recipe/cua_s1/diff_workers.py +++ b/recipe/cua_s1/diff_workers.py @@ -63,14 +63,20 @@ def compare(name, py, rs): rs_, rc, rb, _ = rs if ps != 200 or rs_ != 200: same = (ps, pc, pb) == (rs_, rc, rb) - return same, "" if same else f"python {ps} {pb[:200]!r} | native {rs_} {rb[:200]!r}", 0.0 + return ( + same, + "" if same else f"python {ps} {pb[:200]!r} | native {rs_} {rb[:200]!r}", + 0.0, + ) p, r = ordered(pb), ordered(rb) pd, rd = dict(p), dict(r) notes, worst = [], 0.0 if [k for k, _ in p] != [k for k, _ in r]: notes.append("top-level keys differ") if pd["model"] != rd["model"] or pd["usage"] != rd["usage"]: - notes.append(f"model/usage differ: {pd['model']} {pd['usage']} vs {rd['model']} {rd['usage']}") + notes.append( + f"model/usage differ: {pd['model']} {pd['usage']} vs {rd['model']} {rd['usage']}" + ) pa, ra = pd["answers"], rd["answers"] if [k for k, _ in pa] != [k for k, _ in ra]: notes.append("question order differs") @@ -85,29 +91,58 @@ def compare(name, py, rs): top = sorted(pv, reverse=True) margin = top[0] - (top[1] if len(top) > 1 else 0.0) if pans["choice"] != rans["choice"] and margin >= 0.05: - notes.append(f"{qn}: choice {pans['choice']!r} vs {rans['choice']!r} (margin {margin:.3f})") + notes.append( + f"{qn}: choice {pans['choice']!r} vs {rans['choice']!r} (margin {margin:.3f})" + ) return not notes, "; ".join(notes), worst +def median(values): + return sorted(values)[len(values) // 2] + + def main() -> None: - corpus, pport, rport, out_path = sys.argv[1], int(sys.argv[2]), int(sys.argv[3]), sys.argv[4] - cases = [json.loads(l) for l in open(corpus)] + corpus, pport, rport, out_path = ( + sys.argv[1], + int(sys.argv[2]), + int(sys.argv[3]), + sys.argv[4], + ) + cases = [json.loads(line) for line in open(corpus)] extra = [] max_body = 4 << 20 big = b'{"model": "cua-s1-4b-0.2", "state": "' + b"x" * (max_body + 1) + b'"}' extra.append(("http/too_large_content_length", "POST", "/v1/systemone", big, False)) extra.append(("http/too_large_chunked", "POST", "/v1/systemone", big, True)) - fill = max_body - len(b'{"model": "cua-s1-4b-0.2", "state": "", "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "A"}}}}') - exact = b'{"model": "cua-s1-4b-0.2", "state": "' + b"ab " * (fill // 3) + b"a" * (fill % 3) + b'", "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "A"}}}}' + fill = max_body - len( + b'{"model": "cua-s1-4b-0.2", "state": "", "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "A"}}}}' + ) + exact = ( + b'{"model": "cua-s1-4b-0.2", "state": "' + + b"ab " * (fill // 3) + + b"a" * (fill % 3) + + b'", "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "A"}}}}' + ) assert len(exact) == max_body, len(exact) extra.append(("http/exactly_max_body", "POST", "/v1/systemone", exact, False)) - extra.append(("http/chunked_small", "POST", "/v1/systemone", base64.b64decode(cases[0]["body"]), True)) + extra.append( + ( + "http/chunked_small", + "POST", + "/v1/systemone", + base64.b64decode(cases[0]["body"]), + True, + ) + ) extra.append(("http/get_systemone", "GET", "/v1/systemone", None, False)) extra.append(("http/post_health", "POST", "/health", b"{}", False)) extra.append(("http/not_found", "GET", "/nope", None, False)) results, failures, worst, times = [], [], 0.0, {"py": [], "rs": []} - items = [(c["name"], "POST", "/v1/systemone", base64.b64decode(c["body"]), False) for c in cases] + extra + items = [ + (c["name"], "POST", "/v1/systemone", base64.b64decode(c["body"]), False) + for c in cases + ] + extra for i, (name, method, path, body, chunked) in enumerate(items): py = request(pport, method, path, body, chunked=chunked) rs = request(rport, method, path, body, chunked=chunked) @@ -116,8 +151,18 @@ def main() -> None: if py[0] == 200 and rs[0] == 200: times["py"].append(py[3]) times["rs"].append(rs[3]) - results.append({"name": name, "ok": ok, "note": note, "python_status": py[0], - "native_status": rs[0], "max_prob_diff": diff, "python_ms": py[3], "native_ms": rs[3]}) + results.append( + { + "name": name, + "ok": ok, + "note": note, + "python_status": py[0], + "native_status": rs[0], + "max_prob_diff": diff, + "python_ms": py[3], + "native_ms": rs[3], + } + ) if not ok: failures.append((name, note)) if (i + 1) % 100 == 0: @@ -126,7 +171,9 @@ def main() -> None: hp = request(pport, "GET", "/health") hr = request(rport, "GET", "/health") hpj, hrj = json.loads(hp[2]), json.loads(hr[2]) - health_note = {k: (hpj.get(k), hrj.get(k)) for k in {**hpj, **hrj} if hpj.get(k) != hrj.get(k)} + health_note = { + k: (hpj.get(k), hrj.get(k)) for k in {**hpj, **hrj} if hpj.get(k) != hrj.get(k) + } with open(out_path, "w") as f: for r in results: @@ -140,8 +187,10 @@ def main() -> None: print(f"- {name}: {note}") print(f"largest probability difference on answered requests: {worst:.4f}") if times["py"]: - s = lambda v: sorted(v)[len(v) // 2] - print(f"answered requests: python median {s(times['py']):.1f} ms, native median {s(times['rs']):.1f} ms") + py_ms, rs_ms = median(times["py"]), median(times["rs"]) + print( + f"answered requests: python median {py_ms:.1f} ms, native median {rs_ms:.1f} ms" + ) print(f"/health differences (python, native): {health_note}") diff --git a/src/models/cua_s1/native/src/engine.rs b/src/models/cua_s1/native/src/engine.rs index 600b0d4b..cccf1112 100644 --- a/src/models/cua_s1/native/src/engine.rs +++ b/src/models/cua_s1/native/src/engine.rs @@ -168,10 +168,8 @@ fn letter_rows(dir: &Path, letter_ids: &[u32]) -> Result<(Vec, usize)> { let id = id as usize; ensure!(id < vocab, "letter id {id} outside the vocabulary"); let row = &data[id * hidden * 2..(id + 1) * hidden * 2]; - rows.extend( - row.chunks_exact(2) - .map(|b| half::bf16::from_le_bytes([b[0], b[1]]).to_f32()), - ); + let (pairs, _) = row.as_chunks::<2>(); + rows.extend(pairs.iter().map(|&b| half::bf16::from_le_bytes(b).to_f32())); } Ok((rows, hidden)) } diff --git a/src/models/cua_s1/native/src/model.rs b/src/models/cua_s1/native/src/model.rs index d4c77d06..e8a583f9 100644 --- a/src/models/cua_s1/native/src/model.rs +++ b/src/models/cua_s1/native/src/model.rs @@ -921,9 +921,10 @@ impl Model { let mut last = vec![0u8; h * BF16]; // SAFETY: x holds at least t rows of the hidden size. unsafe { cuda::download(&mut last, s.at(s.x + (t - 1) * h * BF16), self.stream)? }; - Ok(last - .chunks_exact(2) - .map(|b| half::bf16::from_le_bytes([b[0], b[1]]).to_f32()) + let (pairs, _) = last.as_chunks::<2>(); + Ok(pairs + .iter() + .map(|&b| half::bf16::from_le_bytes(b).to_f32()) .collect()) } diff --git a/src/models/cua_s1/native/tests/kernels.rs b/src/models/cua_s1/native/tests/kernels.rs index 83664dab..6b373232 100644 --- a/src/models/cua_s1/native/tests/kernels.rs +++ b/src/models/cua_s1/native/tests/kernels.rs @@ -53,9 +53,10 @@ fn from_device(buf: &DeviceBuffer, n: usize, st: Stream) -> Vec { let mut bytes = vec![0u8; n * 2]; // SAFETY: the buffer holds n bfloat16 values. unsafe { cuda::download(&mut bytes, buf.at(0), st).unwrap() }; - bytes - .chunks_exact(2) - .map(|b| bf16::from_le_bytes([b[0], b[1]]).to_f32()) + let (pairs, _) = bytes.as_chunks::<2>(); + pairs + .iter() + .map(|&b| bf16::from_le_bytes(b).to_f32()) .collect() } diff --git a/src/models/cua_s1/native/tests/make_float_vectors.py b/src/models/cua_s1/native/tests/make_float_vectors.py index 655b489b..329f13fd 100644 --- a/src/models/cua_s1/native/tests/make_float_vectors.py +++ b/src/models/cua_s1/native/tests/make_float_vectors.py @@ -19,8 +19,25 @@ def bits(x: float) -> str: def main() -> None: rng = random.Random(20260928) - values = [0.0, -0.0, 1.0, -1.0, 0.5, 0.1, 0.2, 0.3, 1 / 3, 2 / 3, math.pi, math.e, - 5e-324, 2.2250738585072014e-308, 1.7976931348623157e308, 2.0**53, 2.0**53 + 2] + values = [ + 0.0, + -0.0, + 1.0, + -1.0, + 0.5, + 0.1, + 0.2, + 0.3, + 1 / 3, + 2 / 3, + math.pi, + math.e, + 5e-324, + 2.2250738585072014e-308, + 1.7976931348623157e308, + 2.0**53, + 2.0**53 + 2, + ] values += [2.0**e for e in range(-1074, 1024, 7)] values += [10.0**e for e in range(-323, 309)] for e in (-5, -4, 15, 16, 17): @@ -31,7 +48,9 @@ def main() -> None: if pick < 0.5: x = struct.unpack(">d", rng.getrandbits(64).to_bytes(8, "big"))[0] elif pick < 0.8: - x = float(f"{rng.randint(1, 10**rng.randint(1, 17))}e{rng.randint(-30, 30)}") + x = float( + f"{rng.randint(1, 10 ** rng.randint(1, 17))}e{rng.randint(-30, 30)}" + ) else: x = rng.uniform(-1e6, 1e6) if math.isfinite(x): diff --git a/src/models/cua_s1/native/tests/make_printable.py b/src/models/cua_s1/native/tests/make_printable.py index 63fb6fd4..93c285fb 100644 --- a/src/models/cua_s1/native/tests/make_printable.py +++ b/src/models/cua_s1/native/tests/make_printable.py @@ -17,8 +17,12 @@ else: ranges.append([cp, cp]) version = ".".join(map(str, sys.version_info[:3])) -print(f"//! Generated from Python {version} (Unicode {unicodedata.unidata_version}): code points >= 0x80 for which") -print("//! `str.isprintable()` is false, as inclusive ranges. Python's `repr` escapes these.") +print( + f"//! Generated from Python {version} (Unicode {unicodedata.unidata_version}): code points >= 0x80 for which" +) +print( + "//! `str.isprintable()` is false, as inclusive ranges. Python's `repr` escapes these." +) print("//! Regenerate with tests/make_printable.py.") print() print("pub const NON_PRINTABLE: &[(u32, u32)] = &[") From baa80d6c6328eb413613ef0b33a9f8425e9c8a55 Mon Sep 17 00:00:00 2001 From: Tianyao Wu Date: Wed, 30 Sep 2026 09:55:25 +0800 Subject: [PATCH 07/10] cua_s1: keep the text worker to the model, its HTTP worker and tests Move the HTTP worker to src/frontend/cua_s1_text.py, so src/models/cua_s1/text/ holds only the model: contract.py (request mapping, prompts, answers) and model.py (adapter checks, loading, readout; formerly engine.py and adapter.py). The upstream MIT notice moves into contract.py. The upstream comparison and latency scripts are no longer part of the change; the recipe keeps setup and launch only. Signed-off-by: Tianyao Wu --- .gitignore | 9 +- README.md | 2 +- recipe/README.md | 2 +- recipe/cua_s1/bench_text.py | 122 ----------- recipe/cua_s1/compare_text_with_upstream.py | 204 ------------------ recipe/cua_s1/requirements-text.txt | 3 +- recipe/cua_s1/text.md | 23 +- .../server.py => frontend/cua_s1_text.py} | 30 +-- src/models/cua_s1/README.md | 10 +- src/models/cua_s1/text/THIRD_PARTY_NOTICES.md | 23 -- src/models/cua_s1/text/adapter.py | 51 ----- src/models/cua_s1/text/contract.py | 31 ++- src/models/cua_s1/text/engine.py | 84 -------- src/models/cua_s1/text/model.py | 129 +++++++++++ ...est_text_adapter.py => test_text_model.py} | 2 +- tests/cua_s1/test_text_server.py | 30 +-- 16 files changed, 198 insertions(+), 557 deletions(-) delete mode 100644 recipe/cua_s1/bench_text.py delete mode 100644 recipe/cua_s1/compare_text_with_upstream.py rename src/{models/cua_s1/text/server.py => frontend/cua_s1_text.py} (91%) delete mode 100644 src/models/cua_s1/text/THIRD_PARTY_NOTICES.md delete mode 100644 src/models/cua_s1/text/adapter.py delete mode 100644 src/models/cua_s1/text/engine.py create mode 100644 src/models/cua_s1/text/model.py rename tests/cua_s1/{test_text_adapter.py => test_text_model.py} (94%) diff --git a/.gitignore b/.gitignore index 07a95a78..99ba095f 100644 --- a/.gitignore +++ b/.gitignore @@ -20,12 +20,9 @@ target # option (not recommended) you can uncomment the following to ignore the entire idea folder. #.idea/ -# Python workers and local model artifacts -__pycache__/ -.pytest_cache/ -.ruff_cache/ +# Local Python worker environment .venv/ +__pycache__/ +# Model weights downloaded by the recipes weights/ - -# macOS metadata .DS_Store diff --git a/README.md b/README.md index 76944aa3..05d9d08b 100644 --- a/README.md +++ b/README.md @@ -40,7 +40,7 @@ Implementation code lives under `src/`; recipes and documentation stay at the re | Directory | Responsibility | | --- | --- | -| [`src/frontend/`](src/frontend/) | Rust serving code and the small engine interface. | +| [`src/frontend/`](src/frontend/) | Rust serving code, Python worker adapters, and the small engine interface. | | [`src/models/`](src/models/) | Model implementations, one directory per model: preprocessing, batching, state, execution, and output processing. | | [`src/backends/cuda/`](src/backends/cuda/) | NVIDIA GPU operations and kernel integration. | | [`src/backends/metal/`](src/backends/metal/) | Apple GPU operations and kernel integration. | diff --git a/recipe/README.md b/recipe/README.md index 875486cb..48d22d5a 100644 --- a/recipe/README.md +++ b/recipe/README.md @@ -3,7 +3,7 @@ - [Laya text worker](laya/README.md): start the external Python worker, connect the Rust frontend and compare direct and proxied responses. - [Cua-S1 4B 0.2 text worker](cua_s1/text.md): download the pinned weights, start - the worker, connect the Rust frontend and check the worker against upstream. + the worker and connect the Rust frontend. Recipes contain setup, launch commands and examples. Reusable implementation code belongs under `src/`. diff --git a/recipe/cua_s1/bench_text.py b/recipe/cua_s1/bench_text.py deleted file mode 100644 index f5ee9c88..00000000 --- a/recipe/cua_s1/bench_text.py +++ /dev/null @@ -1,122 +0,0 @@ -"""Send the fixed input set to a running worker, directly and through the frontend. - -For each case it checks that the frontend returns the same status, content -type and body bytes as the worker, then measures warm end-to-end latency on -both paths. Warmup requests are sent first and reported separately. Requests -are sequential (concurrency 1). - - python recipe/cua_s1/bench_text.py --direct http://127.0.0.1:8000 \ - --frontend http://127.0.0.1:8080 --warmup 3 --repeat 20 --out results.json -""" - -from __future__ import annotations - -import argparse -import json -import statistics -import time -import urllib.error -import urllib.request -from pathlib import Path - -ROOT = Path(__file__).resolve().parents[2] -# Ignore http_proxy and friends: the worker and the frontend are local. -OPENER = urllib.request.build_opener(urllib.request.ProxyHandler({})) - - -def post(url: str, body: bytes, token: str | None) -> tuple[int, str, bytes, float]: - headers = {"Content-Type": "application/json"} - if token: - headers["Authorization"] = f"Bearer {token}" - request = urllib.request.Request(url + "/v1/systemone", data=body, headers=headers) - started = time.perf_counter() - try: - with OPENER.open(request, timeout=120) as response: - data = response.read() - status, ctype = response.status, response.headers.get("content-type", "") - except urllib.error.HTTPError as error: - data, status, ctype = ( - error.read(), - error.code, - error.headers.get("content-type", ""), - ) - return status, ctype, data, (time.perf_counter() - started) * 1000 - - -def pct(values: list[float], q: float) -> float: - ordered = sorted(values) - return ordered[min(len(ordered) - 1, round(q * (len(ordered) - 1)))] - - -def main() -> int: - parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) - parser.add_argument("--direct", required=True, help="worker base URL") - parser.add_argument( - "--frontend", help="frontend base URL; omit to measure the worker only" - ) - parser.add_argument( - "--inputs", default=str(ROOT / "tests/cua_s1/data/text_inputs.json") - ) - parser.add_argument("--warmup", type=int, default=3) - parser.add_argument("--repeat", type=int, default=20) - parser.add_argument("--token", help="bearer token, if the worker requires one") - parser.add_argument("--out") - args = parser.parse_args() - - cases = json.loads(Path(args.inputs).read_text(encoding="utf-8")) - paths = {"direct": args.direct} - if args.frontend: - paths["frontend"] = args.frontend - results, mismatches = {}, 0 - for name, body in cases.items(): - raw = json.dumps(body, ensure_ascii=False).encode() - status, ctype, direct_body, _ = post(args.direct, raw, args.token) - row = {"status": status, "content_type": ctype} - if status == 200: - reply = json.loads(direct_body) - row["answers"], row["input_tokens"] = ( - reply["answers"], - reply["usage"]["input_tokens"], - ) - else: - row["body"] = direct_body.decode("utf-8", "replace")[:500] - if args.frontend: - f_status, f_ctype, f_body, _ = post(args.frontend, raw, args.token) - row["frontend_identical"] = (f_status, f_ctype, f_body) == ( - status, - ctype, - direct_body, - ) - mismatches += not row["frontend_identical"] - for label, url in paths.items(): - warm = [post(url, raw, args.token)[3] for _ in range(args.warmup)] - times = [post(url, raw, args.token)[3] for _ in range(args.repeat)] - row[label] = { - "warmup_ms": [round(t, 2) for t in warm], - "p50_ms": round(statistics.median(times), 2), - "p95_ms": round(pct(times, 0.95), 2), - "min_ms": round(min(times), 2), - "raw_ms": [round(t, 2) for t in times], - } - results[name] = row - line = f"{name}: status {status}, tokens {row.get('input_tokens')}" - for label in paths: - line += ( - f", {label} p50 {row[label]['p50_ms']} ms p95 {row[label]['p95_ms']} ms" - ) - if args.frontend: - line += f", identical {row['frontend_identical']}" - print(line, flush=True) - if args.out: - Path(args.out).write_text( - json.dumps(results, ensure_ascii=False, indent=1) + "\n" - ) - if args.frontend: - print( - f"{len(cases) - mismatches}/{len(cases)} cases byte-identical through the frontend" - ) - return 1 if mismatches else 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/recipe/cua_s1/compare_text_with_upstream.py b/recipe/cua_s1/compare_text_with_upstream.py deleted file mode 100644 index 8a50a546..00000000 --- a/recipe/cua_s1/compare_text_with_upstream.py +++ /dev/null @@ -1,204 +0,0 @@ -"""Compare the Cua-S1 worker with upstream `FourBModel` on the fixed input set. - -Needs a checkout of trycua/cua at the pinned commit (for `cua_s1.four_b` and -the jev-use chooser) and the pinned weights. The worker's model scores every -question first and is freed; then `FourBModel` is loaded with the same device -and dtype and scores the same questions. For every question it checks: - -- prompt token ids: worker vs upstream `build_prompt` plus the chat template; -- probabilities: worker vs `FourBModel.forward`, exact fp32 equality; -- for the two upstream fixtures, also worker vs the chooser's own path - (`S1DecisionModel.score`), matched by option key. - -Upstream `build_prompt` is given the worker's mapped labels, state and goal, -so the id check covers the prompt layout, chat template and tokenizer. The -request mapping itself (escaping, structured values, `null` labels) is -covered by the unit tests and, independently, by the two fixtures. - -Only one model is resident at a time. Both compute logits for every prompt -position, so the longest input (15,446 tokens) needs about 8 GB for logits in -bfloat16 and 15 GB in float32, on top of the weights. - -Run from the repository root: - - python recipe/cua_s1/compare_text_with_upstream.py --upstream ../cua \\ - --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text --device cuda -""" - -from __future__ import annotations - -import argparse -import gc -import json -import sys -import time -from pathlib import Path - -ROOT = Path(__file__).resolve().parents[2] -sys.path.insert(0, str(ROOT / "src")) - -FIXTURES = { - "fixture_positive": "jev-choice-request-v1.json", - "fixture_negative": "jev-choice-negative-v1.json", -} - - -def free(device: str) -> None: - import torch - - gc.collect() - if device.startswith("cuda"): - torch.cuda.empty_cache() - - -def peak_gib(device: str) -> float | None: - import torch - - if not device.startswith("cuda"): - return None - peak = torch.cuda.max_memory_allocated() / 2**30 - torch.cuda.reset_peak_memory_stats() - return round(peak, 2) - - -def main() -> int: - parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) - parser.add_argument( - "--upstream", required=True, help="trycua/cua checkout at the pinned commit" - ) - parser.add_argument("--base", required=True) - parser.add_argument("--adapter", required=True) - parser.add_argument("--device", default="cuda") - parser.add_argument("--dtype", default="bfloat16") - parser.add_argument( - "--inputs", default=str(ROOT / "tests/cua_s1/data/text_inputs.json") - ) - parser.add_argument("--out", help="write one JSON line per question here") - parser.add_argument( - "--no-tf32", - action="store_true", - help="disable TF32 in cuBLAS and cuDNN (use for fp32 reference runs)", - ) - args = parser.parse_args() - - upstream = Path(args.upstream) - sys.path.insert(0, str(upstream / "libs/cua-s1/python/src")) - sys.path.insert(0, str(upstream / "libs/cua-driver/examples/jev-use/python")) - import torch - - if args.no_tf32: - torch.backends.cuda.matmul.allow_tf32 = False - torch.backends.cudnn.allow_tf32 = False - from cua_s1.four_b import FourBModel, Option, assign_letters, build_prompt - from decision_models import DecisionRequest, S1DecisionModel - - from models.cua_s1.text.contract import ( - ACTION, - APP, - ROLE, - TASK_FAMILY, - map_request, - parse_body, - ) - from models.cua_s1.text.engine import TextEngine - - cases = json.loads(Path(args.inputs).read_text(encoding="utf-8")) - questions = [] - for name, body in cases.items(): - request = map_request(parse_body(json.dumps(body).encode())) - questions += [(name, request, question) for question in request.questions] - - # Pass 1: the worker. - engine = TextEngine(args.base, args.adapter, args.device, args.dtype) - print(f"worker loaded in {engine.load_seconds:.1f} s", flush=True) - worker = {} - for name, request, question in questions: - worker[name, question.name] = ( - engine.prompt_ids(request.state, question), - engine.score(request.state, question).probabilities, - ) - worker_peak = peak_gib(args.device) - del engine - free(args.device) - - # Pass 2: upstream FourBModel, and the chooser for the two fixtures. - started = time.perf_counter() - reference = FourBModel( - base_model=args.base, - lora_adapter_path=args.adapter, - device=args.device, - dtype=args.dtype, - modality="text", - ) - reference.load() - print(f"upstream loaded in {time.perf_counter() - started:.1f} s", flush=True) - fixture_dir = upstream / "libs/cua-driver/examples/jev-use/fixtures" - rows, failures = [], 0 - for name, request, question in questions: - options = [ - Option(element_id=k, role=ROLE, label=label, action=ACTION) - for k, label in zip(question.keys, question.labels, strict=True) - ] - kwargs = dict( - app=APP, - task_family=TASK_FAMILY, - ax_tree=request.state, - modality="text", - goal=question.goal or None, - ) - messages = build_prompt(assign_letters(options), **kwargs) - chat = reference._tokenizer.apply_chat_template( - messages, tokenize=False, add_generation_prompt=True - ) - upstream_ids = reference._tokenizer(chat)["input_ids"] - upstream_p = [r.probability for r in reference.forward(options, **kwargs)] - worker_ids, worker_p = worker[name, question.name] - row = { - "case": name, - "question": question.name, - "options": len(options), - "prompt_tokens": len(worker_ids), - "ids_equal": worker_ids == upstream_ids, - "probs_equal": worker_p == upstream_p, - "max_abs_diff": max( - abs(a - b) for a, b in zip(worker_p, upstream_p, strict=True) - ), - "worker": dict(zip(question.keys, worker_p, strict=True)), - "upstream": dict(zip(question.keys, upstream_p, strict=True)), - } - if name in FIXTURES: - raw = json.loads((fixture_dir / FIXTURES[name]).read_text(encoding="utf-8")) - chooser = ( - S1DecisionModel(reference, modality="text") - .score(DecisionRequest.from_validated(raw)) - .probabilities - ) - row["chooser_equal"] = all( - chooser.get(k) == p for k, p in row["worker"].items() - ) - ok = row["ids_equal"] and row["probs_equal"] and row.get("chooser_equal", True) - failures += not ok - rows.append(row) - print( - f"{'ok ' if ok else 'FAIL'} {name}/{question.name}: {len(options)} options, " - f"{len(worker_ids)} tokens, max |diff| {row['max_abs_diff']:.3g}", - flush=True, - ) - upstream_peak = peak_gib(args.device) - - if args.out: - with open(args.out, "w", encoding="utf-8") as f: - for row in rows: - f.write(json.dumps(row, ensure_ascii=False) + "\n") - if worker_peak is not None: - print(f"peak allocated: worker {worker_peak} GiB, upstream {upstream_peak} GiB") - print( - f"{len(rows) - failures}/{len(rows)} questions identical " - f"(device {args.device}, dtype {args.dtype}, torch {torch.__version__}, " - f"tf32 {'off' if args.no_tf32 else 'default'})" - ) - return 1 if failures else 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/recipe/cua_s1/requirements-text.txt b/recipe/cua_s1/requirements-text.txt index cde47c55..136ef2b3 100644 --- a/recipe/cua_s1/requirements-text.txt +++ b/recipe/cua_s1/requirements-text.txt @@ -1,5 +1,6 @@ # Versions match upstream's `four-b` lock (trycua/cua libs/cua-s1/python/uv.lock -# at 0e75660ce4c2edda519e0c795fa3ad98abf4e76f), which the parity checks rely on. +# at 0e75660ce4c2edda519e0c795fa3ad98abf4e76f), the reference environment in +# src/models/cua_s1/README.md. torch==2.14.0 transformers==5.17.0 tokenizers==0.23.2 diff --git a/recipe/cua_s1/text.md b/recipe/cua_s1/text.md index d7cebdbd..53be6947 100644 --- a/recipe/cua_s1/text.md +++ b/recipe/cua_s1/text.md @@ -1,6 +1,6 @@ # Cua-S1 4B 0.2 text worker -This recipe runs the Cua-S1 4B 0.2 `text` adapter behind the Rust frontend. The worker lives in [`src/models/cua_s1/text/`](../../src/models/cua_s1/text/), and [`src/models/cua_s1/README.md`](../../src/models/cua_s1/README.md) documents the inference contract and the request mapping. Only `choice` questions are supported. +This recipe runs the Cua-S1 4B 0.2 `text` adapter behind the Rust frontend. The model lives in [`src/models/cua_s1/text/`](../../src/models/cua_s1/text/) and the HTTP worker in [`src/frontend/cua_s1_text.py`](../../src/frontend/cua_s1_text.py); [`src/models/cua_s1/README.md`](../../src/models/cua_s1/README.md) documents the inference contract and the request mapping. Only `choice` questions are supported. Run all commands from the repository root, on Linux with an NVIDIA GPU. @@ -33,7 +33,7 @@ To verify every file against upstream's lock, clone [trycua/cua](https://github. ## Start the worker ```sh -PYTHONPATH=src .venv/bin/python -m models.cua_s1.text.server \ +PYTHONPATH=src .venv/bin/python -m frontend.cua_s1_text \ --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text \ --device cuda --dtype bfloat16 --host 127.0.0.1 --port 8000 ``` @@ -68,25 +68,6 @@ The answer has the Jev choice shape. On an RTX 6000 Ada in bfloat16, the respons {"model":"cua-ai/cua-s1-4b-0.2@16818868b0cc7813808aae4e87b417657046ab79:text","answers":{"pick":{"type":"choice","choice":"cancel","probabilities":{"delete":0.0024726232513785362,"cancel":0.9975274205207825},"confidence":0.9750249565060322}},"usage":{"input_tokens":153,"output_tokens":0}} ``` -## Check against upstream - -`compare_text_with_upstream.py` scores the fixed input set (`tests/cua_s1/data/text_inputs.json`) with the worker's model and then with upstream `FourBModel`, one model at a time, and compares the results. It needs the trycua/cua checkout from above and two extra packages for upstream's processor: - -```sh -.venv/bin/python -m pip install torchvision==0.29.0 pillow==11.3.0 -.venv/bin/python recipe/cua_s1/compare_text_with_upstream.py --upstream ../cua \ - --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text --device cuda -``` - -Every question must have identical prompt token ids and identical fp32 probabilities. - -With the worker and the frontend running, `bench_text.py` checks that the frontend returns the same bytes as the worker for every input, then measures warm latency on both paths: - -```sh -.venv/bin/python recipe/cua_s1/bench_text.py --direct http://127.0.0.1:8000 \ - --frontend http://127.0.0.1:8080 --warmup 3 --repeat 20 -``` - ## Tests The contract and HTTP tests need neither weights nor a GPU. The tokenizer tests also run when `CUA_S1_BASE` points to the downloaded base model; they read only its tokenizer files: diff --git a/src/models/cua_s1/text/server.py b/src/frontend/cua_s1_text.py similarity index 91% rename from src/models/cua_s1/text/server.py rename to src/frontend/cua_s1_text.py index cb7da3bc..eb2295a1 100644 --- a/src/models/cua_s1/text/server.py +++ b/src/frontend/cua_s1_text.py @@ -1,9 +1,10 @@ """HTTP worker for Cua-S1 4B 0.2 (`text` adapter) behind the Rust frontend. Routes: `GET /health` and `POST /v1/systemone`. The model is loaded before the -server starts listening, and one forward pass runs at a time. +server starts listening, and one forward pass runs at a time. The model itself +is in `src/models/cua_s1/text/`. - PYTHONPATH=src python -m models.cua_s1.text.server --base --adapter + PYTHONPATH=src python -m frontend.cua_s1_text --base --adapter """ from __future__ import annotations @@ -24,7 +25,7 @@ from fastapi import FastAPI, Request from fastapi.responses import JSONResponse -from .contract import ( +from models.cua_s1.text.contract import ( ADAPTER_REVISION, MODEL_NAME, RequestError, @@ -48,7 +49,7 @@ def build_app( - engine: Any, + model: Any, *, api_key: str | None, max_body_bytes: int, @@ -80,8 +81,8 @@ def health(): "status": "ready", "modality": "text", "model": identity, - "device": engine.device, - "dtype": engine.dtype, + "device": model.device, + "dtype": model.dtype, } def decide(mapped): @@ -89,7 +90,7 @@ def decide(mapped): # before any forward pass runs. encoded = [] for question in mapped.questions: - inputs = engine.encode(mapped.state, question) + inputs = model.encode(mapped.state, question) n = int(inputs["input_ids"].shape[1]) if max_prompt_tokens and n > max_prompt_tokens: raise RequestError( @@ -100,7 +101,7 @@ def decide(mapped): encoded.append((question, inputs)) answers, prompt_tokens = {}, 0 for question, inputs in encoded: - scored = engine.score_encoded(inputs, len(question.keys)) + scored = model.score_encoded(inputs, len(question.keys)) answers[question.name] = answer(question, scored.probabilities) prompt_tokens += scored.prompt_tokens return { @@ -189,8 +190,11 @@ def main(argv: list[str] | None = None) -> None: import uvicorn - from .adapter import downloaded_revision, text_adapter_dir - from .engine import TextEngine + from models.cua_s1.text.model import ( + TextModel, + downloaded_revision, + text_adapter_dir, + ) # Fail before loading weights if this is not the text adapter. text_adapter_dir(args.adapter) @@ -212,13 +216,13 @@ def main(argv: list[str] | None = None) -> None: flush=True, ) - engine = TextEngine(args.base, args.adapter, args.device, args.dtype) + model = TextModel(args.base, args.adapter, args.device, args.dtype) print( - f"loaded in {engine.load_seconds:.1f} s on {args.device} ({args.dtype})", + f"loaded in {model.load_seconds:.1f} s on {args.device} ({args.dtype})", flush=True, ) app = build_app( - engine, + model, api_key=env("CUA_S1_API_KEY") or None, max_body_bytes=args.max_body_bytes, max_questions=args.max_questions, diff --git a/src/models/cua_s1/README.md b/src/models/cua_s1/README.md index f0712d26..8711b17e 100644 --- a/src/models/cua_s1/README.md +++ b/src/models/cua_s1/README.md @@ -2,15 +2,7 @@ This directory owns Cua-S1 4B 0.2 ([#10](https://github.com/ThinkFlowLab/system1-omni/issues/10)): request mapping, prompt construction, adapter selection, execution, and the answer-letter readout. This page records the pinned upstream revisions, the inference contract an implementation must match, and how its outputs will be compared with the upstream reference. -Status: a reference worker for the `text` adapter is in [`text/`](text/). It loads the model directly through Hugging Face Transformers and PEFT and serves `/v1/systemone`; setup and checks are in [`recipe/cua_s1/text.md`](../../../recipe/cua_s1/text.md). A worker for the `multimodal` adapter is proposed in [#12](https://github.com/ThinkFlowLab/system1-omni/pull/12). - -| Path | Contents | -| --- | --- | -| `text/contract.py` | Request validation, the `/v1/systemone` mapping, prompt construction and answers. No torch imports. | -| `text/engine.py` | Model and adapter loading and the answer-letter readout. | -| `text/server.py` | The HTTP worker (`GET /health`, `POST /v1/systemone`). | -| `text/adapter.py` | Finds and checks the local `text` adapter and its downloaded revision. | -| `tests/cua_s1/test_text_*.py` (repository root) | Tests that need neither weights nor a GPU, and tokenizer checks. The fixed input set is `tests/cua_s1/data/text_inputs.json`. | +Status: a reference worker for the `text` adapter loads the model directly through Hugging Face Transformers and PEFT and serves `/v1/systemone`. The model is in [`text/`](text/) (`contract.py` for request mapping, prompts and answers; `model.py` for loading and the answer-letter readout), the HTTP worker is [`src/frontend/cua_s1_text.py`](../../frontend/cua_s1_text.py), and setup is in [`recipe/cua_s1/text.md`](../../../recipe/cua_s1/text.md). A worker for the `multimodal` adapter is proposed in [#12](https://github.com/ThinkFlowLab/system1-omni/pull/12). ## Pinned revisions diff --git a/src/models/cua_s1/text/THIRD_PARTY_NOTICES.md b/src/models/cua_s1/text/THIRD_PARTY_NOTICES.md deleted file mode 100644 index f054d433..00000000 --- a/src/models/cua_s1/text/THIRD_PARTY_NOTICES.md +++ /dev/null @@ -1,23 +0,0 @@ -The system message, prompt layout and fixed values in `contract.py`, and the two upstream fixtures converted in `tests/cua_s1/data/text_inputs.json`, come from [trycua/cua](https://github.com/trycua/cua) at `0e75660ce4c2edda519e0c795fa3ad98abf4e76f` under the following license. - -MIT License - -Copyright (c) 2025 Cua AI, Inc. - -Permission is hereby granted, free of charge, to any person obtaining a copy -of this software and associated documentation files (the "Software"), to deal -in the Software without restriction, including without limitation the rights -to use, copy, modify, merge, publish, distribute, sublicense, and/or sell -copies of the Software, and to permit persons to whom the Software is -furnished to do so, subject to the following conditions: - -The above copyright notice and this permission notice shall be included in all -copies or substantial portions of the Software. - -THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR -IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, -FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE -AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER -LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, -OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE -SOFTWARE. diff --git a/src/models/cua_s1/text/adapter.py b/src/models/cua_s1/text/adapter.py deleted file mode 100644 index 472e533a..00000000 --- a/src/models/cua_s1/text/adapter.py +++ /dev/null @@ -1,51 +0,0 @@ -"""Locate and check the local Cua-S1 `text` adapter. No torch imports.""" - -from __future__ import annotations - -import json -import re -from pathlib import Path - - -def text_adapter_dir(adapter_root: str | Path) -> Path: - """Return the `text` adapter directory under the adapter root. - - Accepts the repository root (`/text`) or the `text/` directory - itself, and refuses the `multimodal/` adapter: PEFT only warns about keys - it cannot place, so loading the wrong adapter would otherwise go unnoticed. - """ - root = Path(adapter_root) - path = root / "text" if (root / "text" / "adapter_config.json").exists() else root - config_file = path / "adapter_config.json" - if not config_file.exists(): - raise RuntimeError(f"no adapter_config.json under {root}") - config = json.loads(config_file.read_text()) - if config.get("base_model_name_or_path") != "Qwen/Qwen3.5-4B": - raise RuntimeError(f"{config_file}: base model is not Qwen/Qwen3.5-4B") - if {"linear_fc1", "linear_fc2"} & set(config.get("target_modules") or []): - raise RuntimeError( - f"{config_file}: this is the multimodal adapter, not the text adapter" - ) - return path - - -def downloaded_revision(adapter_root: str | Path) -> str | None: - """The commit that `hf download --local-dir` recorded for the text adapter, if any. - - `hf download` keeps its metadata under the repository root, so this also - looks one level up when `adapter_root` is the `text/` directory itself. - """ - root = Path(adapter_root) - places = [(root, "text/"), (root, "")] - if root.name == "text": - places.insert(0, (root.parent, "text/")) - for base, prefix in places: - cache = base / ".cache" / "huggingface" / "download" - try: - meta = (cache / f"{prefix}adapter_model.safetensors.metadata").read_text() - first = meta.splitlines()[0].strip() - except (OSError, IndexError): - continue - if re.fullmatch(r"[0-9a-f]{40}", first): - return first - return None diff --git a/src/models/cua_s1/text/contract.py b/src/models/cua_s1/text/contract.py index a0a2d1d3..5fcfcb3e 100644 --- a/src/models/cua_s1/text/contract.py +++ b/src/models/cua_s1/text/contract.py @@ -16,7 +16,6 @@ ADAPTER_REPO = "cua-ai/cua-s1-4b-0.2" ADAPTER_REVISION = "16818868b0cc7813808aae4e87b417657046ab79" BASE_REPO = "Qwen/Qwen3.5-4B" -BASE_REVISION = "851bf6e806efd8d0a36b00ddf55e13ccb7b8cd0a" LETTERS = string.ascii_uppercase MAX_OPTIONS = len(LETTERS) @@ -25,8 +24,30 @@ # copied from trycua/cua at 0e75660ce4c2edda519e0c795fa3ad98abf4e76f: # `libs/cua-s1/python/src/cua_s1/four_b.py` (SYSTEM_PROMPT, build_prompt, # _describe_option) and `libs/cua-driver/examples/jev-use/python/ -# decision_models.py` (S1DecisionModel.score). MIT License, Copyright (c) 2025 -# Cua AI, Inc.; see THIRD_PARTY_NOTICES.md. +# decision_models.py` (S1DecisionModel.score), as are the two upstream fixtures +# converted in `tests/cua_s1/data/text_inputs.json`. +# +# MIT License +# +# Copyright (c) 2025 Cua AI, Inc. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. SYSTEM_PROMPT = ( "You are a one-pass computer-use decision model. You are shown the " "current state of a screen and a fixed, closed list of candidate " @@ -269,5 +290,5 @@ def answer(question: Question, probabilities: list[float]) -> dict[str, Any]: } -def model_identity(revision: str = ADAPTER_REVISION, modality: str = "text") -> str: - return f"{ADAPTER_REPO}@{revision}:{modality}" +def model_identity(revision: str = ADAPTER_REVISION) -> str: + return f"{ADAPTER_REPO}@{revision}:text" diff --git a/src/models/cua_s1/text/engine.py b/src/models/cua_s1/text/engine.py deleted file mode 100644 index cdcc351d..00000000 --- a/src/models/cua_s1/text/engine.py +++ /dev/null @@ -1,84 +0,0 @@ -"""Load Qwen3.5-4B with the Cua-S1 `text` adapter and score one prompt. - -The calls mirror upstream `cua_s1.four_b.FourBModel` (text modality): the same -model class, an unmerged PEFT adapter, the chat template with its default -generation prompt, full logits, and a fp32 softmax over the letter logits at -the last position. Keeping them the same is what makes the worker's -probabilities bitwise identical to the reference in the same environment. -""" - -from __future__ import annotations - -import time -from dataclasses import dataclass - -import torch -from peft import PeftModel -from transformers import AutoModelForCausalLM, AutoTokenizer - -from .adapter import text_adapter_dir -from .contract import LETTERS, Question, build_messages - - -@dataclass -class Scored: - probabilities: list[float] - prompt_tokens: int - - -class TextEngine: - def __init__( - self, base_model: str, adapter_root: str, device: str, dtype: str - ) -> None: - self.device = device - self.dtype = dtype - started = time.perf_counter() - self.tokenizer = AutoTokenizer.from_pretrained(base_model) - model = AutoModelForCausalLM.from_pretrained( - base_model, dtype=getattr(torch, dtype), device_map=device - ) - model = PeftModel.from_pretrained(model, str(text_adapter_dir(adapter_root))) - model.eval() - self.model = model - self.load_seconds = time.perf_counter() - started - self.letter_ids = self._letter_ids() - - def _letter_ids(self) -> list[int]: - ids = [] - for letter in LETTERS: - tokens = self.tokenizer.encode(letter, add_special_tokens=False) - if len(tokens) != 1: - raise RuntimeError(f"letter {letter!r} is not a single token: {tokens}") - ids.append(tokens[0]) - return ids - - def encode(self, state: str, question: Question): - """Tokenized prompt for one question, on CPU. - - The Qwen3.5 tokenizer adds no special tokens here (contract point 4); - the chat template already contains them. - """ - chat_text = self.tokenizer.apply_chat_template( - build_messages(state, question), tokenize=False, add_generation_prompt=True - ) - return self.tokenizer(chat_text, return_tensors="pt") - - def prompt_ids(self, state: str, question: Question) -> list[int]: - return self.encode(state, question)["input_ids"][0].tolist() - - @torch.no_grad() - def score_encoded(self, inputs, n_options: int) -> Scored: - inputs = inputs.to(self.model.device) - out = self.model(**inputs) - final_logits = out.logits[0, -1, :] - letter_ids = self.letter_ids[:n_options] - option_logits = final_logits[ - torch.tensor(letter_ids, device=final_logits.device) - ] - probabilities = torch.softmax(option_logits.float(), dim=-1).tolist() - return Scored( - probabilities=probabilities, prompt_tokens=int(inputs["input_ids"].shape[1]) - ) - - def score(self, state: str, question: Question) -> Scored: - return self.score_encoded(self.encode(state, question), len(question.keys)) diff --git a/src/models/cua_s1/text/model.py b/src/models/cua_s1/text/model.py new file mode 100644 index 00000000..9724c426 --- /dev/null +++ b/src/models/cua_s1/text/model.py @@ -0,0 +1,129 @@ +"""Load Qwen3.5-4B with the Cua-S1 `text` adapter and score one prompt. + +The calls mirror upstream `cua_s1.four_b.FourBModel` (text modality): the same +model class, an unmerged PEFT adapter, the chat template with its default +generation prompt, full logits, and a fp32 softmax over the letter logits at +the last position. Keeping them the same is what makes the worker's +probabilities bitwise identical to the reference in the same environment. + +Torch, Transformers and PEFT are imported only when a model is loaded, so the +adapter checks can run without them. +""" + +from __future__ import annotations + +import json +import re +import time +from dataclasses import dataclass +from pathlib import Path + +from .contract import BASE_REPO, LETTERS, Question, build_messages + + +def text_adapter_dir(adapter_root: str | Path) -> Path: + """Return the `text` adapter directory under the adapter root. + + Accepts the repository root (`/text`) or the `text/` directory + itself, and refuses the `multimodal/` adapter: PEFT only warns about keys + it cannot place, so loading the wrong adapter would otherwise go unnoticed. + """ + root = Path(adapter_root) + path = root / "text" if (root / "text" / "adapter_config.json").exists() else root + config_file = path / "adapter_config.json" + if not config_file.exists(): + raise RuntimeError(f"no adapter_config.json under {root}") + config = json.loads(config_file.read_text()) + if config.get("base_model_name_or_path") != BASE_REPO: + raise RuntimeError(f"{config_file}: base model is not {BASE_REPO}") + if {"linear_fc1", "linear_fc2"} & set(config.get("target_modules") or []): + raise RuntimeError( + f"{config_file}: this is the multimodal adapter, not the text adapter" + ) + return path + + +def downloaded_revision(adapter_root: str | Path) -> str | None: + """The commit that `hf download --local-dir` recorded for the text adapter, if any. + + `hf download` keeps its metadata under the repository root, so this also + looks one level up when `adapter_root` is the `text/` directory itself. + """ + root = Path(adapter_root) + places = [(root, "text/"), (root, "")] + if root.name == "text": + places.insert(0, (root.parent, "text/")) + for base, prefix in places: + cache = base / ".cache" / "huggingface" / "download" + try: + meta = (cache / f"{prefix}adapter_model.safetensors.metadata").read_text() + first = meta.splitlines()[0].strip() + except (OSError, IndexError): + continue + if re.fullmatch(r"[0-9a-f]{40}", first): + return first + return None + + +@dataclass +class Scored: + probabilities: list[float] + prompt_tokens: int + + +class TextModel: + def __init__( + self, base_model: str, adapter_root: str, device: str, dtype: str + ) -> None: + import torch + from peft import PeftModel + from transformers import AutoModelForCausalLM, AutoTokenizer + + self.device = device + self.dtype = dtype + started = time.perf_counter() + self.tokenizer = AutoTokenizer.from_pretrained(base_model) + model = AutoModelForCausalLM.from_pretrained( + base_model, dtype=getattr(torch, dtype), device_map=device + ) + model = PeftModel.from_pretrained(model, str(text_adapter_dir(adapter_root))) + model.eval() + self.model = model + self.load_seconds = time.perf_counter() - started + self.letter_ids = self._letter_ids() + + def _letter_ids(self) -> list[int]: + ids = [] + for letter in LETTERS: + tokens = self.tokenizer.encode(letter, add_special_tokens=False) + if len(tokens) != 1: + raise RuntimeError(f"letter {letter!r} is not a single token: {tokens}") + ids.append(tokens[0]) + return ids + + def encode(self, state: str, question: Question): + """Tokenized prompt for one question, on CPU. + + The Qwen3.5 tokenizer adds no special tokens here (contract point 4); + the chat template already contains them. + """ + chat_text = self.tokenizer.apply_chat_template( + build_messages(state, question), tokenize=False, add_generation_prompt=True + ) + return self.tokenizer(chat_text, return_tensors="pt") + + def score_encoded(self, inputs, n_options: int) -> Scored: + import torch + + with torch.no_grad(): + inputs = inputs.to(self.model.device) + out = self.model(**inputs) + final_logits = out.logits[0, -1, :] + letter_ids = self.letter_ids[:n_options] + option_logits = final_logits[ + torch.tensor(letter_ids, device=final_logits.device) + ] + probabilities = torch.softmax(option_logits.float(), dim=-1).tolist() + return Scored( + probabilities=probabilities, prompt_tokens=int(inputs["input_ids"].shape[1]) + ) diff --git a/tests/cua_s1/test_text_adapter.py b/tests/cua_s1/test_text_model.py similarity index 94% rename from tests/cua_s1/test_text_adapter.py rename to tests/cua_s1/test_text_model.py index e3255440..8dd0fd83 100644 --- a/tests/cua_s1/test_text_adapter.py +++ b/tests/cua_s1/test_text_model.py @@ -4,7 +4,7 @@ import pytest -from models.cua_s1.text.adapter import downloaded_revision, text_adapter_dir +from models.cua_s1.text.model import downloaded_revision, text_adapter_dir REV = "16818868b0cc7813808aae4e87b417657046ab79" diff --git a/tests/cua_s1/test_text_server.py b/tests/cua_s1/test_text_server.py index cd4e51dd..ed5cb7fd 100644 --- a/tests/cua_s1/test_text_server.py +++ b/tests/cua_s1/test_text_server.py @@ -1,4 +1,4 @@ -"""HTTP tests for the worker with a fake engine: no weights, no torch.""" +"""HTTP tests for the worker with a fake model: no weights, no torch.""" import json from dataclasses import dataclass @@ -11,7 +11,7 @@ pytest.importorskip("httpx") from fastapi.testclient import TestClient # noqa: E402 -from models.cua_s1.text.server import build_app # noqa: E402 +from frontend.cua_s1_text import build_app # noqa: E402 @dataclass @@ -19,7 +19,7 @@ class _Ids: shape: tuple -class FakeEngine: +class FakeModel: device = "cpu" dtype = "float32" @@ -49,9 +49,9 @@ class Scored: return Scored(probabilities, inputs["input_ids"].shape[1]) -def client(engine=None, api_key=None, max_body_bytes=4 << 20, max_prompt_tokens=32768): +def client(model=None, api_key=None, max_body_bytes=4 << 20, max_prompt_tokens=32768): app = build_app( - engine or FakeEngine(), + model or FakeModel(), api_key=api_key, max_body_bytes=max_body_bytes, max_questions=64, @@ -136,26 +136,26 @@ def test_limits(): headers={"content-type": "application/json"}, ) assert streamed.status_code == 413 - engine = FakeEngine(tokens=40000) + model = FakeModel(tokens=40000) body = json.loads(json.dumps(BODY)) body["questions"]["r"] = body["questions"]["q"] - response = client(engine).post("/v1/systemone", json=body) + response = client(model).post("/v1/systemone", json=body) assert response.status_code == 413 assert "token limit" in response.json()["detail"] - assert engine.forward_calls == 0 + assert model.forward_calls == 0 -@pytest.mark.parametrize("engine", [FakeEngine(fail=True), FakeEngine(nan=True)]) -def test_engine_failure_is_json_500(engine): - response = client(engine).post("/v1/systemone", json=BODY) +@pytest.mark.parametrize("model", [FakeModel(fail=True), FakeModel(nan=True)]) +def test_model_failure_is_json_500(model): + response = client(model).post("/v1/systemone", json=BODY) assert response.status_code == 500 assert response.json() == {"detail": "inference failed"} def test_warmup_runs_the_request_path(): - engine = FakeEngine() + model = FakeModel() app = build_app( - engine, + model, api_key=None, max_body_bytes=1 << 20, max_questions=64, @@ -163,10 +163,10 @@ def test_warmup_runs_the_request_path(): revision="r", ) app.state.warmup() - assert engine.forward_calls == 1 + assert model.forward_calls == 1 with pytest.raises(ValueError): build_app( - FakeEngine(nan=True), + FakeModel(nan=True), api_key=None, max_body_bytes=1 << 20, max_questions=64, From a086a316babe796d41ba96cc4b4be0cc5123f15d Mon Sep 17 00:00:00 2001 From: Tianyao Wu Date: Wed, 30 Sep 2026 10:00:01 +0800 Subject: [PATCH 08/10] cua_s1: keep the native worker to what serving needs Drop the comparison and benchmark tooling (diff_corpus.py, diff_workers.py, check_native.py, and the worker's --score-all, --bench and --encode-only modes with the eager/served switch they used), the float-vector generator, and the Unicode table that made error messages escape non-ASCII names exactly as Python's repr does, along with code only they used. The export record and tokenizer are now checked once at startup, and the export script refuses weights without download metadata, which the worker would refuse later. The upstream MIT notice moves into contract.rs, and the README and recipe keep setup, launch and tests. Signed-off-by: Tianyao Wu --- recipe/cua_s1/check_native.py | 73 -- recipe/cua_s1/diff_corpus.py | 406 ---------- recipe/cua_s1/diff_workers.py | 198 ----- recipe/cua_s1/export_text_merged.py | 33 +- recipe/cua_s1/native.md | 64 +- src/models/cua_s1/native/README.md | 34 +- .../cua_s1/native/THIRD_PARTY_NOTICES.md | 23 - src/models/cua_s1/native/src/contract.rs | 31 +- src/models/cua_s1/native/src/cuda.rs | 4 - src/models/cua_s1/native/src/engine.rs | 58 +- src/models/cua_s1/native/src/lib.rs | 1 - src/models/cua_s1/native/src/main.rs | 160 +--- src/models/cua_s1/native/src/model.rs | 14 +- src/models/cua_s1/native/src/printable.rs | 717 ------------------ src/models/cua_s1/native/src/pyjson.rs | 93 +-- src/models/cua_s1/native/src/server.rs | 31 +- .../cua_s1/native/tests/make_float_vectors.py | 64 -- .../cua_s1/native/tests/make_printable.py | 31 - 18 files changed, 136 insertions(+), 1899 deletions(-) delete mode 100644 recipe/cua_s1/check_native.py delete mode 100644 recipe/cua_s1/diff_corpus.py delete mode 100644 recipe/cua_s1/diff_workers.py delete mode 100644 src/models/cua_s1/native/THIRD_PARTY_NOTICES.md delete mode 100644 src/models/cua_s1/native/src/printable.rs delete mode 100644 src/models/cua_s1/native/tests/make_float_vectors.py delete mode 100644 src/models/cua_s1/native/tests/make_printable.py diff --git a/recipe/cua_s1/check_native.py b/recipe/cua_s1/check_native.py deleted file mode 100644 index 2c239773..00000000 --- a/recipe/cua_s1/check_native.py +++ /dev/null @@ -1,73 +0,0 @@ -"""Check `omni-cua-s1-native --score-all` output against the float32 reference. - - python recipe/cua_s1/check_native.py scores.jsonl [scores from another start.jsonl] - -The parity directory holds parity_float32.jsonl and parity_bfloat16.jsonl, written by -compare_text_with_upstream.py --dtype float32 / bfloat16 --out (see native.md). - -The rule is the one src/models/cua_s1/README.md declares for a native engine: the -largest per-option difference from the float32 worker is at most 2 x (bfloat16 -worker vs float32) + 0.01, and the top option matches float32 wherever float32's -top-two margin is at least 0.05. The served result (from a graph when the -prompt fits) must be bitwise identical to the eager one. With a second file (another -process start), both runs must be bitwise identical too. -""" - -import json -import sys -from pathlib import Path - - -def load_parity(path): - rows = [json.loads(line) for line in path.read_text().splitlines() if line] - return {(r["case"], r["question"]): r["worker"] for r in rows} - - -def load(path): - return [json.loads(line) for line in Path(path).read_text().splitlines() if line] - - -def diff(a, b): - return max(abs(a[k] - b[k]) for k in b) - - -runs = load(sys.argv[1]) -parity = Path(sys.argv[2]) -fp32, bf16 = ( - load_parity(parity / "parity_float32.jsonl"), - load_parity(parity / "parity_bfloat16.jsonl"), -) -allowance = 2 * max(diff(bf16[k], fp32[k]) for k in fp32) + 0.01 -eager = {(r["case"], r["question"]): r for r in runs if r["mode"] == "eager"} -served = {(r["case"], r["question"]): r for r in runs if r["mode"] == "served"} -missing = sorted(set(fp32) - set(eager)) -worst, worst_at, flips = 0.0, None, [] -for key, r in eager.items(): - ref, got = fp32[key], r["probabilities"] - d = diff(got, ref) - if d > worst: - worst, worst_at = d, f"{key[0]}/{key[1]}" - top = sorted(ref.values(), reverse=True) - margin = top[0] - (top[1] if len(top) > 1 else 0.0) - if margin >= 0.05 and max(ref, key=ref.get) != max(got, key=got.get): - flips.append(f"{key[0]}/{key[1]}") -graph_keys = [k for k, r in served.items() if r["graph"]] -mismatch = [ - f"{k[0]}/{k[1]}" - for k in served - if served[k]["probabilities"] != eager[k]["probabilities"] -] -print(f"{len(eager)} questions; allowance {allowance:.4f}") -print( - f"largest |eager - fp32| {worst:.4f} ({worst_at}); top-option changes: {flips or 'none'}; missing: {missing or 'none'}" -) -print( - f"served from a graph: {len(graph_keys)}; served differs from eager: {mismatch or 'none'}" -) -ok = worst <= allowance and not flips and not missing and not mismatch -if len(sys.argv) > 3: - other = load(sys.argv[3]) - same = len(other) == len(runs) and all(a == b for a, b in zip(runs, other)) - print(f"identical to the other start: {same}") - ok = ok and same -print("PASS" if ok else "FAIL") diff --git a/recipe/cua_s1/diff_corpus.py b/recipe/cua_s1/diff_corpus.py deleted file mode 100644 index 635ccef7..00000000 --- a/recipe/cua_s1/diff_corpus.py +++ /dev/null @@ -1,406 +0,0 @@ -"""Request bodies for comparing the Python and native Cua-S1 text workers. - - python recipe/cua_s1/diff_corpus.py tests/cua_s1/data/text_inputs.json corpus.jsonl [n_fuzz] - -Each line is {"name": ..., "body": }. The set has -the fixed input set, hand-written edge cases for every error path of -`contract.parse_body` / `map_request` / the server, and seeded random bodies. -""" - -from __future__ import annotations - -import base64 -import json -import random -import sys - -M = "cua-s1-4b-0.2" - - -def req(state="Button: OK", questions=None, model=M, **extra): - body = {"model": model, "state": state} - body["questions"] = questions if questions is not None else {"q": choice()} - body.update(extra) - return body - - -def choice(instructions="Press OK.", criteria=None, type_="choice"): - q = {"type": type_, "instructions": instructions} - q["criteria"] = ( - criteria if criteria is not None else {"ok": "OK", "cancel": "Cancel"} - ) - return q - - -def build(inputs_path: str, n_fuzz: int) -> list[tuple[str, bytes]]: - cases: list[tuple[str, bytes]] = [] - - def add(name, body): - if isinstance(body, (dict, list)): - body = json.dumps(body, ensure_ascii=False) - if isinstance(body, str): - body = body.encode("utf-8", "surrogatepass") - cases.append((name, body)) - - fixtures = json.load(open(inputs_path, encoding="utf-8")) - for name, body in fixtures.items(): - add(f"fixture/{name}", body) - add( - f"fixture_ascii_compact/{name}", - json.dumps(body, ensure_ascii=True, separators=(",", ":")), - ) - - # --- JSON syntax and decoding --- - ok = json.dumps(req()) - raw = { - "empty": "", - "space": " ", - "open": "{", - "close": "}", - "array": "[]", - "null": "null", - "number": "1", - "string": '"x"', - "true": "true", - "extra": ok + " x", - "bom": "" + ok, - "ws": " \t\n" + ok + "\r\n", - "trailing_comma_obj": ok[:-1] + ",}", - "single_quotes": ok.replace('"', "'"), - "comment": "// c\n" + ok, - "nan": ok.replace('"Button: OK"', "NaN"), - "inf": ok.replace('"Button: OK"', "Infinity"), - "neg_inf": ok.replace('"Button: OK"', "-Infinity"), - "nan_nested": ok.replace('"OK"', "[1, NaN]"), - "float_range": ok.replace('"OK"', "[1e400]"), - "float_range_neg": ok.replace('"OK"', "[-1E309]"), - "float_range_then_garbage": ok.replace('"OK"', "[1e400x]"), - "float_underflow": ok.replace('"OK"', "[1e-400, 1.0e+308, -0.0, 5e-324]"), - "int_4300": ok.replace('"OK"', "[" + "9" * 4300 + "]"), - "int_4301": ok.replace('"OK"', "[" + "9" * 4301 + "]"), - "int_neg_4301": ok.replace('"OK"', "[-" + "9" * 4301 + "]"), - "leading_zero": ok.replace('"OK"', "[01]"), - "dot_no_digit": ok.replace('"OK"', "[1.]"), - "exp_no_digit": ok.replace('"OK"', "[1e]"), - "exp_sign_no_digit": ok.replace('"OK"', "[1e+]"), - "minus_alone": ok.replace('"OK"', "[-]"), - "plus": ok.replace('"OK"', "[+1]"), - "hex": ok.replace('"OK"', "[0x10]"), - "raw_control": ok.replace("Press OK.", "Press\u0001OK."), - "raw_tab": ok.replace("Press OK.", "Press\tOK."), - "raw_del": ok.replace("Press OK.", "Press\u007fOK."), - "bad_escape": ok.replace("Press OK.", "Press \\x OK."), - "short_u": ok.replace("Press OK.", "\\u12"), - "bad_u": ok.replace("Press OK.", "\\uZZZZ"), - "escapes": ok.replace("Press OK.", '\\/\\b\\f\\n\\r\\t\\"\\\\ \\u00e9\\u0000'), - "missing_colon": ok.replace('"model":', '"model"'), - "unterminated": ok[:-3], - "dup_top": '{"model": "' - + M - + '", "model": "' - + M - + '", "state": "s", "questions": {}}', - "dup_criteria": ok.replace('"cancel": "Cancel"', '"ok": "Again"'), - "dup_then_syntax": ok.replace('"cancel": "Cancel"', '"ok": "Again"') + "x", - "syntax_inside_dup_object": ok.replace( - '"cancel": "Cancel"', '"ok": "Again", x' - ), - "dup_state_nested": ok.replace('"Button: OK"', '{"a": {"b": 1, "b": 2}}'), - "dup_after_nan": ok.replace('"cancel": "Cancel"', '"ok": NaN'), - "lone_high": ok.replace("Press OK.", "\\ud800"), - "lone_low": ok.replace("Press OK.", "\\udc00x"), - "pair": ok.replace("Press OK.", "\\ud83d\\ude00"), - "reversed_pair": ok.replace("Press OK.", "\\ude00\\ud83d"), - "high_then_bmp": ok.replace("Press OK.", "\\ud800\\u0041"), - "lone_in_key": ok.replace('"ok": "OK"', '"\\udfff": "OK"'), - "lone_dup_key": ok.replace( - '"ok": "OK", "cancel": "Cancel"', '"\\ud800": 1, "\\ud800": 2' - ), - "lone_then_syntax": ok.replace("Press OK.", "\\ud800") + ",", - } - for name, text in raw.items(): - add(f"json/{name}", text) - add("json/invalid_utf8", ok.encode().replace(b"Press", b"Pr\xffss")) - add("json/overlong_utf8", ok.encode().replace(b"Press", b"Pr\xc0\xafss")) - add("json/utf8_surrogate", ok.encode().replace(b"Press", b"Pr\xed\xa0\x80ss")) - add("json/utf8_bom_bytes", b"\xef\xbb\xbf" + ok.encode()) - add("json/latin1", ok.encode().replace(b"Press", b"Pr\xe9ss")) - - def nested(d, leaf='"x"', open_="[", close="]"): - return open_ * d + leaf + close * d - - for d in (1, 50, 966, 967, 968, 969, 1500): - add(f"depth/state_list_{d}", ok.replace('"Button: OK"', nested(d))) - for d in (964, 965, 966): - add(f"depth/criteria_{d}", ok.replace('"OK"', nested(d))) - add(f"depth/type_{d}", ok.replace('"choice"', nested(d))) - add( - f"depth/instructions_obj_{d}", - ok.replace('"Press OK."', nested(d, '"x"', '{"a":', "}")), - ) - - head = '{"model": "cua-s1-4b-0.2", "state": ' - tail = ', "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "A"}}}}' - for d in (967, 968, 969, 970): - add(f"depth2/empty_list_leaf_{d}", head + nested(d, "") + tail) - add(f"depth2/empty_obj_leaf_{d}", head + nested(d, "{}") + tail) - add(f"depth2/number_leaf_{d}", head + nested(d, "1") + tail) - for d in (9988, 9989, 9990, 9991): - add(f"depth2/parse_limit_garbage_{d}", head + nested(d, "1") + tail + "x") - add("depth2/deep_then_dup", head + nested(5000, "1") + ', "state": 1' + tail) - add("depth2/deep_then_nan", head + nested(3000, "1") + ', "z": NaN' + tail) - add( - "depth2/surrogate_before_deep", - '{"model": "\\ud800", "state": ' + nested(1500, "1") + tail, - ) - add( - "depth2/deep_before_surrogate", - head + nested(1500, "1") + ', "z": "\\ud800"' + tail, - ) - add( - "depth2/deep_before_surrogate_key", - '{"a": ' + nested(1500, "1") + ', "\\ud800": 1}', - ) - add("depth2/hundred_thousand", "[" * 100000) - add("depth2/deep_top_level_array", nested(5000, "1")) - sa = { - "_sample": choice(criteria={"_save": "Save", "_sa": None, "b": "B"}), - "q": choice(), - } - add("keys/_sa_names_and_options", req(questions=sa)) - add("keys/_sa_state_keys", req(state={"_sa": 1, "_sample": [2]})) - - # --- mapping errors, in Python's check order --- - sem = { - "no_model": {"state": "s", "questions": {"q": choice()}}, - "model_null": req(model=None), - "model_case": req(model="Cua-S1-4B-0.2"), - "model_space": req(model=M + " "), - "model_number": req(model=1), - "model_wrong_and_no_state": {"model": "x"}, - "no_state": {"model": M, "questions": {"q": choice()}}, - "state_null": req(state=None), - "state_true": req(state=True), - "state_zero": req(state=0), - "state_float": req(state=1.5), - "state_empty": req(state=""), - "state_empty_obj": req(state={}), - "state_empty_list": req(state=[]), - "state_spaces": req(state=" "), - "state_obj": req( - state={ - "a": None, - "b": [1, 2.5, -0.0, 1e22, 12345678901234567890, True], - "é": "
", - } - ), - "state_list": req(state=["Button: OK", {"k": "v"}, 3.14e-07]), - "no_questions": {"model": M, "state": "s"}, - "questions_null": req(questions=None) | {"questions": None}, - "questions_list": req(questions=[]), - "questions_empty": req(questions={}), - "questions_str": req(questions="q"), - "questions_65": req(questions={f"q{i}": choice() for i in range(65)}), - "questions_65_bad_types": req( - questions={f"q{i}": {"type": "score"} for i in range(65)} - ), - "questions_12": req( - questions={f"q{i}": choice(f"Pick {i}.") for i in range(12)} - ), - "question_list": req(questions={"q": []}), - "question_null": req(questions={"q": None}), - "question_str": req(questions={"q": "choice"}), - "score_second": req( - questions={"a": {"type": "choice"}, "b": {"type": "score"}} - ), - "noul": req(questions={"a": {"type": "noul"}}), - "type_missing": req( - questions={"a": {"instructions": "x", "criteria": {"a": "A"}}} - ), - "missing_instructions_then_score": req( - questions={"a": {"type": "choice"}, "b": {"type": "noul"}} - ), - "no_instructions": req( - questions={"a": {"type": "choice", "criteria": {"x": "X"}}} - ), - "instructions_null": req(questions={"a": choice(None)}), - "instructions_empty": req(questions={"a": choice("")}), - "instructions_zero": req(questions={"a": choice(0)}), - "instructions_false": req(questions={"a": choice(False)}), - "instructions_obj": req(questions={"a": choice({"x": 1.5e-7, "y": [None]})}), - "instructions_empty_obj": req(questions={"a": choice({})}), - "instructions_list": req(questions={"a": choice([])}), - "no_criteria": req(questions={"a": {"type": "choice", "instructions": "x"}}), - "criteria_null": req( - questions={"a": choice(criteria=None) | {"criteria": None}} - ), - "criteria_list": req(questions={"a": choice(criteria=[])}), - "criteria_empty": req(questions={"a": choice(criteria={})}), - "criteria_26": req( - questions={ - "a": choice(criteria={f"o{i}": f"Option {i}" for i in range(26)}) - } - ), - "criteria_27": req( - questions={ - "a": choice(criteria={f"o{i}": f"Option {i}" for i in range(27)}) - } - ), - "criteria_values": req( - questions={ - "a": choice( - criteria={ - "n": None, - "o": {"k": [1, 2]}, - "l": [], - "e": {}, - "q": 'quote"s', - "b": "back\\slash", - "nl": "new\nline", - "t": "tab\t", - "z": "\u0000", - "ls": "
", - "é": "ünï 中文 😀", - } - ) - } - ), - "criteria_bool": req(questions={"a": choice(criteria={"x": "X", "y": True})}), - "criteria_number": req(questions={"a": choice(criteria={"x": 1})}), - "second_question_bad": req(questions={"a": choice(), "b": choice(criteria={})}), - } - for name, body in sem.items(): - add(f"map/{name}", body) - - weird_names = [ - "it's", - 'say "hi"', - "both'\"", - "back\\slash", - "new\nline", - "nul\u0000", - "del\u007f", - "nbsp ", - "ls
", - "zw​", - "tag\U000e0001", - "pua", - "unassigned͸", - "emoji😀", - "é", - "combining é", - "rtl א", - "soft­", - "space ", - "", - "\t", - ] - for i, name in enumerate(weird_names): - add(f"names/type_{i}", req(questions={name: {"type": name}})) - add(f"names/option_{i}", req(questions={name: choice(criteria={name: 1})})) - add( - f"names/ok_{i}", - req(questions={name: choice(criteria={name: None, "other": "Other"})}), - ) - for i, kind in enumerate( - [ - "", - "Choice", - 1, - 1.5, - 1e16, - -0.0, - 1e-5, - True, - None, - [], - {}, - {"a": [1, {"b": None}]}, - 12345678901234567890, - 2.5e-310, - ] - ): - add(f"types/{i}", req(questions={"q": {"type": kind}})) - - # --- prompt length --- - long_state = "Row: item " * 3000 # a little over 16384 tokens - add("limit/prompt_too_long", req(state=long_state)) - add( - "limit/second_prompt_too_long", - req(questions={"a": choice(), "b": choice("x " * 17000)}), - ) - add("limit/prompt_just_under", req(state="Row: item " * 2600)) - - # --- seeded random bodies --- - rnd = random.Random(0) - pool = ( - "abcXYZ019 _-.:/\\\"'\n\t\r{}[]<>|é中文😀
​ ́א\u0000\u001f\u007f" - "<|im_start|><|im_end|>" - ) - - def rstr(n=12): - return "".join(rnd.choice(pool) for _ in range(rnd.randint(0, n))) - - def rnum(): - k = rnd.random() - if k < 0.3: - return rnd.randint(-(10 ** rnd.randint(1, 30)), 10 ** rnd.randint(1, 30)) - if k < 0.9: - return rnd.uniform(-1, 1) * 10 ** rnd.randint(-320, 300) - return rnd.choice([0.0, -0.0, 5e-324, 1e16, 1e-5, 0.1, 1 / 3]) - - def rval(depth=0): - k = rnd.random() - if depth > 3 or k < 0.35: - return rstr() - if k < 0.5: - return rnum() - if k < 0.55: - return rnd.choice([None, True, False]) - if k < 0.75: - return [rval(depth + 1) for _ in range(rnd.randint(0, 4))] - return {rstr(6): rval(depth + 1) for _ in range(rnd.randint(0, 4))} - - for i in range(n_fuzz): - qs = {} - for _ in range(rnd.randint(1, 3)): - crit = { - rstr(8) or "k": rnd.choice([rstr(), None, rval(), rval()]) - for _ in range(rnd.randint(1, 6)) - } - qs[rstr(8)] = { - "type": "choice" if rnd.random() < 0.9 else rval(), - "instructions": rnd.choice([rstr(40), None, rval(), ""]), - "criteria": crit, - } - body = req(state=rnd.choice([rstr(200), rval(), rval()]), questions=qs) - text = json.dumps(body, ensure_ascii=rnd.random() < 0.3) - add(f"fuzz/{i}", text) - # byte-level damage to a copy - b = bytearray(text.encode("utf-8", "surrogatepass")) - for _ in range(rnd.randint(1, 3)): - op, pos = rnd.random(), rnd.randrange(len(b)) - if op < 0.4: - del b[pos] - elif op < 0.8: - b.insert(pos, rnd.choice(b'{}[],:"\\ 0e-.\x00\xff')) - else: - b[pos] = rnd.randrange(256) - add(f"fuzz_damaged/{i}", bytes(b)) - return cases - - -def main() -> None: - n_fuzz = int(sys.argv[3]) if len(sys.argv) > 3 else 150 - cases = build(sys.argv[1], n_fuzz) - with open(sys.argv[2], "w") as f: - for name, body in cases: - f.write( - json.dumps({"name": name, "body": base64.b64encode(body).decode()}) - + "\n" - ) - print(f"{len(cases)} bodies") - - -if __name__ == "__main__": - main() diff --git a/recipe/cua_s1/diff_workers.py b/recipe/cua_s1/diff_workers.py deleted file mode 100644 index 976e6ae3..00000000 --- a/recipe/cua_s1/diff_workers.py +++ /dev/null @@ -1,198 +0,0 @@ -"""Send every corpus body to the Python and the native worker and compare the answers. - - python recipe/cua_s1/diff_workers.py corpus.jsonl out.jsonl - -Errors must match exactly: status, content type and body bytes. For answers, the -model identity, usage, question and option order and the answer type must match; -probabilities are compared numerically, and the choice may differ only where the -Python worker's top-two margin is under 0.05. -""" - -from __future__ import annotations - -import base64 -import http.client -import json -import sys -import time - - -def request(port, method, path, body=None, headers=None, chunked=False): - conn = http.client.HTTPConnection("127.0.0.1", port, timeout=600) - headers = dict(headers or {}) - if body is not None and not chunked: - headers.setdefault("content-type", "application/json") - started = time.perf_counter() - if chunked: - conn.putrequest(method, path) - conn.putheader("transfer-encoding", "chunked") - conn.putheader("content-type", "application/json") - conn.endheaders() - step = 1 << 16 - try: - for i in range(0, len(body), step): - part = body[i : i + step] - conn.send(f"{len(part):x}\r\n".encode() + part + b"\r\n") - conn.send(b"0\r\n\r\n") - except (BrokenPipeError, ConnectionResetError): - pass - else: - try: - conn.request(method, path, body=body, headers=headers) - except (BrokenPipeError, ConnectionResetError): - pass # the server answered before reading the whole body - try: - r = conn.getresponse() - except (ConnectionResetError, http.client.RemoteDisconnected) as e: - conn.close() - return -1, None, repr(e).encode(), (time.perf_counter() - started) * 1000 - data = r.read() - ms = (time.perf_counter() - started) * 1000 - ctype = r.getheader("content-type") - conn.close() - return r.status, ctype, data, ms - - -def ordered(data: bytes): - return json.loads(data, object_pairs_hook=lambda pairs: pairs) - - -def compare(name, py, rs): - """Return (ok, note, max_prob_diff).""" - ps, pc, pb, _ = py - rs_, rc, rb, _ = rs - if ps != 200 or rs_ != 200: - same = (ps, pc, pb) == (rs_, rc, rb) - return ( - same, - "" if same else f"python {ps} {pb[:200]!r} | native {rs_} {rb[:200]!r}", - 0.0, - ) - p, r = ordered(pb), ordered(rb) - pd, rd = dict(p), dict(r) - notes, worst = [], 0.0 - if [k for k, _ in p] != [k for k, _ in r]: - notes.append("top-level keys differ") - if pd["model"] != rd["model"] or pd["usage"] != rd["usage"]: - notes.append( - f"model/usage differ: {pd['model']} {pd['usage']} vs {rd['model']} {rd['usage']}" - ) - pa, ra = pd["answers"], rd["answers"] - if [k for k, _ in pa] != [k for k, _ in ra]: - notes.append("question order differs") - for (qn, pans), (_, rans) in zip(pa, ra): - pans, rans = dict(pans), dict(rans) - pp, rp = pans["probabilities"], rans["probabilities"] - if [k for k, _ in pp] != [k for k, _ in rp] or pans["type"] != rans["type"]: - notes.append(f"{qn}: option keys or type differ") - continue - pv, rv = [v for _, v in pp], [v for _, v in rp] - worst = max(worst, max(abs(a - b) for a, b in zip(pv, rv))) - top = sorted(pv, reverse=True) - margin = top[0] - (top[1] if len(top) > 1 else 0.0) - if pans["choice"] != rans["choice"] and margin >= 0.05: - notes.append( - f"{qn}: choice {pans['choice']!r} vs {rans['choice']!r} (margin {margin:.3f})" - ) - return not notes, "; ".join(notes), worst - - -def median(values): - return sorted(values)[len(values) // 2] - - -def main() -> None: - corpus, pport, rport, out_path = ( - sys.argv[1], - int(sys.argv[2]), - int(sys.argv[3]), - sys.argv[4], - ) - cases = [json.loads(line) for line in open(corpus)] - extra = [] - max_body = 4 << 20 - big = b'{"model": "cua-s1-4b-0.2", "state": "' + b"x" * (max_body + 1) + b'"}' - extra.append(("http/too_large_content_length", "POST", "/v1/systemone", big, False)) - extra.append(("http/too_large_chunked", "POST", "/v1/systemone", big, True)) - fill = max_body - len( - b'{"model": "cua-s1-4b-0.2", "state": "", "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "A"}}}}' - ) - exact = ( - b'{"model": "cua-s1-4b-0.2", "state": "' - + b"ab " * (fill // 3) - + b"a" * (fill % 3) - + b'", "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "A"}}}}' - ) - assert len(exact) == max_body, len(exact) - extra.append(("http/exactly_max_body", "POST", "/v1/systemone", exact, False)) - extra.append( - ( - "http/chunked_small", - "POST", - "/v1/systemone", - base64.b64decode(cases[0]["body"]), - True, - ) - ) - extra.append(("http/get_systemone", "GET", "/v1/systemone", None, False)) - extra.append(("http/post_health", "POST", "/health", b"{}", False)) - extra.append(("http/not_found", "GET", "/nope", None, False)) - - results, failures, worst, times = [], [], 0.0, {"py": [], "rs": []} - items = [ - (c["name"], "POST", "/v1/systemone", base64.b64decode(c["body"]), False) - for c in cases - ] + extra - for i, (name, method, path, body, chunked) in enumerate(items): - py = request(pport, method, path, body, chunked=chunked) - rs = request(rport, method, path, body, chunked=chunked) - ok, note, diff = compare(name, py, rs) - worst = max(worst, diff) - if py[0] == 200 and rs[0] == 200: - times["py"].append(py[3]) - times["rs"].append(rs[3]) - results.append( - { - "name": name, - "ok": ok, - "note": note, - "python_status": py[0], - "native_status": rs[0], - "max_prob_diff": diff, - "python_ms": py[3], - "native_ms": rs[3], - } - ) - if not ok: - failures.append((name, note)) - if (i + 1) % 100 == 0: - print(f"{i + 1}/{len(items)} done, {len(failures)} failures", flush=True) - - hp = request(pport, "GET", "/health") - hr = request(rport, "GET", "/health") - hpj, hrj = json.loads(hp[2]), json.loads(hr[2]) - health_note = { - k: (hpj.get(k), hrj.get(k)) for k in {**hpj, **hrj} if hpj.get(k) != hrj.get(k) - } - - with open(out_path, "w") as f: - for r in results: - f.write(json.dumps(r, ensure_ascii=False) + "\n") - statuses = {} - for r in results: - statuses[r["python_status"]] = statuses.get(r["python_status"], 0) + 1 - print(f"{len(results)} requests, python statuses {statuses}") - print(f"mismatches: {len(failures)}") - for name, note in failures[:40]: - print(f"- {name}: {note}") - print(f"largest probability difference on answered requests: {worst:.4f}") - if times["py"]: - py_ms, rs_ms = median(times["py"]), median(times["rs"]) - print( - f"answered requests: python median {py_ms:.1f} ms, native median {rs_ms:.1f} ms" - ) - print(f"/health differences (python, native): {health_note}") - - -if __name__ == "__main__": - main() diff --git a/recipe/cua_s1/export_text_merged.py b/recipe/cua_s1/export_text_merged.py index da38237c..4036bcd2 100644 --- a/recipe/cua_s1/export_text_merged.py +++ b/recipe/cua_s1/export_text_merged.py @@ -19,6 +19,7 @@ import argparse import hashlib import json +import re import time from pathlib import Path @@ -26,15 +27,18 @@ import torch import transformers -from models.cua_s1.text.adapter import downloaded_revision from models.cua_s1.text.contract import ADAPTER_REPO, BASE_REPO, LETTERS -from models.cua_s1.text.engine import TextEngine +from models.cua_s1.text.model import TextModel, downloaded_revision def base_revision(base: Path) -> str | None: - """The commit Hugging Face recorded when it downloaded config.json.""" + """The commit Hugging Face recorded when it downloaded config.json, if any.""" meta = base / ".cache/huggingface/download/config.json.metadata" - return meta.read_text().splitlines()[0].strip() if meta.exists() else None + try: + first = meta.read_text().splitlines()[0].strip() + except (OSError, IndexError): + return None + return first if re.fullmatch(r"[0-9a-f]{40}", first) else None def main() -> None: @@ -44,19 +48,30 @@ def main() -> None: parser.add_argument("--out", required=True, type=Path) parser.add_argument("--device", default="cuda") args = parser.parse_args() + # The native worker refuses an export that does not record both revisions. + revisions = { + "base": base_revision(args.base), + "adapter": downloaded_revision(args.adapter), + } + missing = [name for name, revision in revisions.items() if revision is None] + if missing: + parser.error( + f"no download metadata for the {' and '.join(missing)} weights; " + "download them with `hf download --revision ... --local-dir ...`" + ) started = time.perf_counter() - engine = TextEngine(str(args.base), str(args.adapter), args.device, "bfloat16") - model = engine.model.merge_and_unload().eval() + loaded = TextModel(str(args.base), str(args.adapter), args.device, "bfloat16") + model = loaded.model.merge_and_unload().eval() model.save_pretrained(args.out, safe_serialization=True, max_shard_size="5GB") - engine.tokenizer.save_pretrained(args.out) + loaded.tokenizer.save_pretrained(args.out) tokenizer = (args.out / "tokenizer.json").read_bytes() record = { "format": "cua-s1-text-merged/1", - "base": {"repo": BASE_REPO, "revision": base_revision(args.base)}, + "base": {"repo": BASE_REPO, "revision": revisions["base"]}, "adapter": { "repo": ADAPTER_REPO, - "revision": downloaded_revision(args.adapter), + "revision": revisions["adapter"], "subfolder": "text", }, "merge": { diff --git a/recipe/cua_s1/native.md b/recipe/cua_s1/native.md index 38073854..e4f8da45 100644 --- a/recipe/cua_s1/native.md +++ b/recipe/cua_s1/native.md @@ -2,7 +2,7 @@ The native worker ([`src/models/cua_s1/native/`](../../src/models/cua_s1/native/)) serves the `text` adapter like the reference worker in [`text.md`](text.md), with the forward pass on the CUDA kernels in [`src/backends/cuda/qwen3_5/`](../../src/backends/cuda/qwen3_5/). It needs an NVIDIA GPU with compute capability 8.0 or newer and was measured on an RTX 6000 Ada (sm_89). -Run the commands from the repository root. The reference worker's setup from `text.md` is needed once, to export the merged weights and for the checks. +Run the commands from the repository root. The reference worker's setup from `text.md` is needed once, to export the merged weights. ## Build @@ -36,9 +36,9 @@ The first start tunes the GEMMs for this GPU and writes the choices to `--gemm-p The same options exist as environment variables (`CUA_S1_MODEL`, `CUA_S1_PORT`, `CUA_S1_GEMM_PLANS`, ...; see `--help`), as do the reference worker's request limits and `CUA_S1_API_KEY`. The Rust frontend and the requests are as in `text.md`. -## Check it +## Tests -Tests without a GPU, then the kernel checks (attention against a float32 kernel, the gated delta rule against a float64 token-by-token reference): +Tests without a GPU, then the kernel checks (attention against a float32 kernel, the gated delta rule against a float64 token-by-token reference, and the GEMM plans): ```sh cargo test -p omni-cua-s1-native @@ -46,64 +46,6 @@ CUA_S1_CUDA_LIB=$PWD/target/release/libqwen3_5_cuda.so \ cargo test --release -p omni-cua-s1-native --test kernels -- --ignored ``` -Accuracy against the float32 reference worker, under the tolerance in [`src/models/cua_s1/README.md`](../../src/models/cua_s1/README.md#validation). First write the reference results with the reference worker's check (`text.md`), once in each dtype: - -```sh -mkdir -p parity -.venv/bin/python recipe/cua_s1/compare_text_with_upstream.py --upstream ../cua \ - --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text --device cuda \ - --dtype float32 --no-tf32 --out parity/parity_float32.jsonl -# the same with --dtype bfloat16 --out parity/parity_bfloat16.jsonl -``` - -Then score the fixed input set eagerly and as served, in two separate processes, and compare: - -```sh -for run in 1 2; do - target/release/omni-cua-s1-native --model weights/cua-s1-4b-0.2-text-merged \ - --gemm-plans weights/gemm-plans.json --score-all tests/cua_s1/data/text_inputs.json > scores-$run.jsonl -done -python3 recipe/cua_s1/check_native.py scores-1.jsonl parity scores-2.jsonl -``` - -It passes when the largest difference and the top options are within the tolerance, the served results (CUDA graphs) equal the eager ones bit for bit, and the two processes agree bit for bit. - -The HTTP behaviour against the reference worker: start the reference worker on port 8001 and the native worker on port 8002, then - -```sh -python3 recipe/cua_s1/diff_corpus.py tests/cua_s1/data/text_inputs.json corpus.jsonl 150 -python3 recipe/cua_s1/diff_workers.py corpus.jsonl 8001 8002 diff.jsonl -``` - -sends the fixed input set, edge cases for every error the workers return, and random bodies to both. Errors must be identical (status, content type and body); answers must have the same model, usage, keys and types, and the same choice wherever the reference's top-two margin is at least 0.05. - -Latency, through the frontend and directly, as for the reference worker: - -```sh -.venv/bin/python recipe/cua_s1/bench_text.py --direct http://127.0.0.1:8000 \ - --frontend http://127.0.0.1:8080 --warmup 3 --repeat 20 -``` - -`--bench ... --bench-gap-ms 50` times the forward pass alone, with idle time between passes. - -## Results - -On one RTX 6000 Ada (48 GB, sm_89), CUDA 13.2, driver 595.91.07, with the pinned revisions: - -- Accuracy: over the 16 questions of the fixed input set, the largest difference from the float32 reference worker is 0.0126 (the allowance is 2 × 0.0145 + 0.01 = 0.039), and no top option changes. Served and eager results are bitwise identical, and so are two separate processes with the same plans file. -- HTTP: 567 requests to both workers (the fixed input set, error cases and random bodies) give identical status and body for every error and the same keys, types and choices for every answer; the largest probability difference is 0.029. -- Frontend: all 14 requests return identical status, content type and body bytes directly and through the frontend. -- Startup: 2.0 s to a ready `/health` with a plans file (about 70 s on the first start with `--gemm-search`); the first request after that took 16 ms. The card peaked at 12.6 GiB in use while serving. -- Latency, p50 in milliseconds, one request at a time. `bench_text.py` sends requests back to back, which keeps this card at its 300 W power limit; the forward-only numbers leave 50 ms idle before each pass, closer to decisions that arrive one by one. The reference worker's numbers come from the same two methods. - -| Case | Prompt tokens | Reference worker, `bench_text.py` | Native, `bench_text.py` | Reference worker, forward only | Native, forward only | -| --- | --- | --- | --- | --- | --- | -| `one_option` | 139 | 46.2 | 15.7 | 44.6 | 12.2 | -| `fixture_positive` | 218 | 49.2 | 18.8 | 47.4 | 13.8 | -| `fixture_negative` | 292 | 52.7 | 25.6 | 51.1 | 17.0 | -| `max_26_options` | 712 | 109.3 | 50.3 | 103.0 | 33.4 | -| `long_state` | 15446 | 3084.7 | 1364.5 | 3047.6 | 1308.9 | - ## Not covered - `score` and `noul` questions, and the `multimodal` adapter, as in the reference worker. diff --git a/src/models/cua_s1/native/README.md b/src/models/cua_s1/native/README.md index eb247a5d..1910b0dd 100644 --- a/src/models/cua_s1/native/README.md +++ b/src/models/cua_s1/native/README.md @@ -1,17 +1,16 @@ # Cua-S1 4B 0.2 native text worker -A `/v1/systemone` worker for the `text` adapter in Rust. It answers every request the way the reference worker in [`../text/`](../text/) does (same validation, error bodies, prompt token ids and answer format) and runs the Qwen3.5-4B forward pass on the CUDA kernels in [`src/backends/cuda/qwen3_5/`](../../../backends/cuda/qwen3_5/), with no Python and no PyTorch. Setup, launch and checks are in [`recipe/cua_s1/native.md`](../../../../recipe/cua_s1/native.md). +A `/v1/systemone` worker for the `text` adapter in Rust. It handles requests like the reference worker (same validation, status codes, prompt token ids and answer format) and runs the Qwen3.5-4B forward pass on the CUDA kernels in [`src/backends/cuda/qwen3_5/`](../../../backends/cuda/qwen3_5/), with no Python and no PyTorch. Setup and launch are in [`recipe/cua_s1/native.md`](../../../../recipe/cua_s1/native.md). | File | Contents | | --- | --- | -| `src/pyjson.rs` | JSON parsing and output that behave like CPython 3.12's `json` module and `repr`, so request errors and answers match the Python worker byte for byte. | +| `src/pyjson.rs` | JSON parsing and `json.dumps` output as the Python worker has them, since structured `state`, `instructions` and `criteria` values reach the prompt as `json.dumps` text. | | `src/contract.rs` | The request mapping, prompt text, confidence and answers of `../text/contract.py`. | -| `src/server.rs` | The HTTP worker (`GET /health`, `POST /v1/systemone`), with the reference worker's limits, status codes and error bodies. | +| `src/server.rs` | The HTTP worker (`GET /health`, `POST /v1/systemone`) with the reference worker's limits and status codes. | | `src/engine.rs` | Tokenization, the letter rows of the output projection, and scoring. | | `src/model.rs` | The Qwen3.5 text model: weights, buffers, the layer loop, CUDA graphs and GEMM tuning. | | `src/cuda.rs` | Loading `libqwen3_5_cuda.so` and the calls into it. | -| `tests/kernels.rs` | GPU checks of the attention and Gated DeltaNet kernels (ignored unless asked for; they need `CUA_S1_CUDA_LIB`). | -| `THIRD_PARTY_NOTICES.md` | The license of the prompt text and fixed values that `src/contract.rs` copies from trycua/cua. | +| `tests/kernels.rs` | GPU checks of the attention and Gated DeltaNet kernels and of the GEMM plans. | ## How a question is answered @@ -19,30 +18,21 @@ A `/v1/systemone` worker for the `text` adapter in Rust. It answers every reques 2. One forward pass runs over the prompt: bfloat16 weights, with the `text` adapter merged into them by `recipe/cua_s1/export_text_merged.py`. The operations follow the Transformers implementation and round to bfloat16 where it does, except inside attention and the Gated DeltaNet prefill, which keep some intermediate results in bfloat16 as FlashAttention and flash-linear-attention do. 3. The final-norm hidden state at the last position is multiplied by the 26 letter rows of the output projection (float32 with float64 accumulation), and a softmax over the question's letters gives the option probabilities. -The whole path from request to answer is in this crate. `src/backends/cuda/qwen3_5/` provides the operations: RMSNorm variants, the Gated DeltaNet convolution, gates and chunked prefill, rotary embedding and attention, and bfloat16 GEMMs through cuBLASLt. - ## CUDA library, graphs and GEMM plans -The kernels are built into `libqwen3_5_cuda.so` by `src/backends/cuda/qwen3_5/build.sh` and loaded when the worker starts (`--cuda-lib`, by default next to the executable), so building the crate needs no CUDA toolkit and the workspace checks run anywhere. - -Prompts up to `--graph-max-tokens` (2048) run as a CUDA graph captured for their exact length on first use; the 128 most recently used lengths keep theirs. There is no padding, and a graph queues the same kernels with the same GEMM algorithms as the eager pass, so both give bitwise identical results. Longer prompts run eagerly. +`libqwen3_5_cuda.so` is built by `src/backends/cuda/qwen3_5/build.sh` and loaded when the worker starts, so building the crate needs no CUDA toolkit. -Projections that share an input run as one GEMM (`gate_proj` and `up_proj`; `in_proj_qkv`, `in_proj_z`, `in_proj_b` and `in_proj_a`; `q_proj`, `k_proj` and `v_proj`), with their weights stored back to back. GEMM algorithms are tuned at startup for a set of prompt lengths, by timing cuBLASLt's candidates. For prompts up to `--graph-max-tokens`, `--gemm-search` enumerates far more configurations (each algorithm with its tiles, stage counts, swizzles and several split-K factors), times each once and times the 12 fastest properly. `--gemm-plans` keeps the choices in a file, so that later starts reuse them and give the same results. The file records the GPU, the cuBLASLt version, the workspace size and `--graph-max-tokens`; a file that does not match is refused, and removing it tunes again. +Prompts up to `--graph-max-tokens` (2048) run as a CUDA graph captured for their exact length on first use, and the 128 most recently used lengths keep theirs. A graph queues the same kernels with the same GEMM algorithms as an eager pass, so both give bitwise identical results. Longer prompts run eagerly. -## Known differences from the reference worker +Projections that share an input run as one GEMM. GEMM algorithms are tuned at startup for a set of prompt lengths by timing cuBLASLt's candidates; `--gemm-search` times far more configurations for prompts up to `--graph-max-tokens`. `--gemm-plans` keeps the choices in a file, so later starts reuse them and give the same results. The file records the GPU, the cuBLASLt version, the workspace size and `--graph-max-tokens`; a file that does not match is refused, and removing it tunes again. -Request bodies are handled the same way. The nesting limits copy what the Python worker does on CPython 3.12.13 (parsing fails for arrays nested more than 9,990 deep, and the check after parsing past 969 nested calls). In Python both come from recursion limits, so they move with the interpreter's call stack, and some values need more of it: close to the limit, `NaN` inside 9,985 arrays is "nested too deeply" in Python and "NaN is not valid JSON" here. The status is 400 either way. +## Differences from the reference worker -Outside the request body the HTTP stacks differ (Starlette and h11 in Python, axum and hyper here): - -- A trailing slash (`/v1/systemone/`) gets a 307 redirect from Starlette and a 404 here; a percent-encoded path is decoded by uvicorn and not here. -- FastAPI also serves `/docs`, `/redoc` and `/openapi.json`; they are not served here. -- `HEAD /health` returns 200 here and 405 from FastAPI. -- The HTTP parsers reject different malformed requests (control bytes in header values, a missing `Host`, a `Content-Length` too large to parse), and those rejections are not JSON. A body that fails partway through reading gets a JSON 400 here. +- The probabilities are not bitwise identical: the adapter is merged, the kernels differ, and the GEMM algorithms depend on the prompt length and the GPU. They are held to the tolerance in [`../README.md`](../README.md#validation). +- Error messages quote names as Python's `repr` does, except that non-printable characters outside ASCII, such as U+00A0 or U+200B, are written as they are instead of escaped. +- The HTTP stacks differ outside the request body (Starlette and uvicorn there, axum and hyper here): trailing slashes, percent-encoded paths, `HEAD /health`, FastAPI's `/docs`, and which malformed HTTP requests are refused. - `GET /health` also reports `"mode": "native"`, and its `dtype` is always `bfloat16`. -The probabilities are not bitwise identical to the reference worker's: the adapter is merged, the kernels differ, and the GEMM algorithms depend on the prompt length and the GPU. They are checked against the float32 worker under the tolerance in [`../README.md`](../README.md#validation). - ## Tests -`cargo test -p omni-cua-s1-native` runs the request-handling tests. Three more are ignored unless asked for with `-- --ignored`: a check of float formatting against Python, on a file that `tests/make_float_vectors.py` writes (`CUA_S1_FLOAT_VECTORS`), and the kernel checks in `tests/kernels.rs`, which need a GPU and `CUA_S1_CUDA_LIB` pointing to a built `libqwen3_5_cuda.so`. +`cargo test -p omni-cua-s1-native` runs the request-handling tests. The kernel checks in `tests/kernels.rs` need a GPU and run with `-- --ignored` and `CUA_S1_CUDA_LIB` pointing to a built `libqwen3_5_cuda.so`. diff --git a/src/models/cua_s1/native/THIRD_PARTY_NOTICES.md b/src/models/cua_s1/native/THIRD_PARTY_NOTICES.md deleted file mode 100644 index 32a974d6..00000000 --- a/src/models/cua_s1/native/THIRD_PARTY_NOTICES.md +++ /dev/null @@ -1,23 +0,0 @@ -The system message, prompt layout and fixed values in `src/contract.rs` come from [trycua/cua](https://github.com/trycua/cua) at `0e75660ce4c2edda519e0c795fa3ad98abf4e76f` under the following license. - -MIT License - -Copyright (c) 2025 Cua AI, Inc. - -Permission is hereby granted, free of charge, to any person obtaining a copy -of this software and associated documentation files (the "Software"), to deal -in the Software without restriction, including without limitation the rights -to use, copy, modify, merge, publish, distribute, sublicense, and/or sell -copies of the Software, and to permit persons to whom the Software is -furnished to do so, subject to the following conditions: - -The above copyright notice and this permission notice shall be included in all -copies or substantial portions of the Software. - -THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR -IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, -FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE -AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER -LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, -OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE -SOFTWARE. diff --git a/src/models/cua_s1/native/src/contract.rs b/src/models/cua_s1/native/src/contract.rs index 4336689a..a117f6e1 100644 --- a/src/models/cua_s1/native/src/contract.rs +++ b/src/models/cua_s1/native/src/contract.rs @@ -1,6 +1,6 @@ //! Request mapping, prompt construction and answers for Cua-S1 4B 0.2, ported from -//! `src/models/cua_s1/text/contract.py` so the two workers answer alike, down to the -//! error messages. +//! `src/models/cua_s1/text/contract.py`, so the two workers build the same prompts and +//! reject the same requests with the same status codes. use std::fmt::Write as _; @@ -19,8 +19,29 @@ pub const MAX_OPTIONS: usize = 26; // from trycua/cua at 0e75660ce4c2edda519e0c795fa3ad98abf4e76f: // `libs/cua-s1/python/src/cua_s1/four_b.py` (SYSTEM_PROMPT, build_prompt, // _describe_option) and `libs/cua-driver/examples/jev-use/python/decision_models.py` -// (S1DecisionModel.score). MIT License, Copyright (c) 2025 Cua AI, Inc.; the full -// notice is in THIRD_PARTY_NOTICES.md. +// (S1DecisionModel.score). +// +// MIT License +// +// Copyright (c) 2025 Cua AI, Inc. +// +// Permission is hereby granted, free of charge, to any person obtaining a copy +// of this software and associated documentation files (the "Software"), to deal +// in the Software without restriction, including without limitation the rights +// to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +// copies of the Software, and to permit persons to whom the Software is +// furnished to do so, subject to the following conditions: +// +// The above copyright notice and this permission notice shall be included in all +// copies or substantial portions of the Software. +// +// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +// IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +// FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +// AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +// LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +// OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +// SOFTWARE. pub const SYSTEM_PROMPT: &str = "You are a one-pass computer-use decision model. You are shown the \ current state of a screen and a fixed, closed list of candidate \ (element, action) options, each given a single letter. Choose exactly \ @@ -85,7 +106,7 @@ fn text(s: &PyStr) -> &str { fn as_text(value: &Value) -> String { match value { Value::Str(s) => text(s).to_string(), - other => dumps(other, false), + other => dumps(other), } } diff --git a/src/models/cua_s1/native/src/cuda.rs b/src/models/cua_s1/native/src/cuda.rs index 769f3208..c33c7af9 100644 --- a/src/models/cua_s1/native/src/cuda.rs +++ b/src/models/cua_s1/native/src/cuda.rs @@ -224,10 +224,6 @@ impl DeviceBuffer { Ok(Self { ptr, bytes }) } - pub fn bytes(&self) -> usize { - self.bytes - } - /// The device address `offset` bytes into the buffer. pub fn at(&self, offset: usize) -> *mut c_void { debug_assert!(offset <= self.bytes); diff --git a/src/models/cua_s1/native/src/engine.rs b/src/models/cua_s1/native/src/engine.rs index cccf1112..d2b10814 100644 --- a/src/models/cua_s1/native/src/engine.rs +++ b/src/models/cua_s1/native/src/engine.rs @@ -3,7 +3,7 @@ //! the output projection. use std::path::Path; -use std::sync::{Arc, Mutex, MutexGuard}; +use std::sync::{Arc, Mutex}; use std::time::Instant; use anyhow::{Context, Result, bail, ensure}; @@ -13,7 +13,7 @@ use tokenizers::Tokenizer; use crate::contract::{self, LETTERS, Question}; use crate::model::Model; -pub use crate::model::{Mode, Options}; +use crate::model::Options; /// What `cua_s1_export.json` records about a merged checkpoint. #[derive(Debug, Clone)] @@ -67,16 +67,14 @@ pub fn provenance(dir: &Path) -> Result { }) } -/// Chat text and token ids for a question; needs only `tokenizer.json`. +/// Chat text and token ids for a question. pub struct Prompter { tokenizer: Tokenizer, pub letter_ids: Vec, } impl Prompter { - /// Checks `tokenizer.json` against `cua_s1_export.json` first (see `provenance`). - pub fn load(dir: &Path) -> Result { - provenance(dir)?; + fn load(dir: &Path) -> Result { let tokenizer = Tokenizer::from_file(dir.join("tokenizer.json")) .map_err(|e| anyhow::anyhow!("tokenizer.json: {e}"))?; let mut letter_ids = Vec::with_capacity(LETTERS.len()); @@ -181,12 +179,15 @@ pub struct Engine { hidden: usize, pub load_seconds: f64, pub device: String, + pub provenance: Provenance, } impl Engine { - /// Load the CUDA library and the model, and prepare CUDA graphs (see `Options`). + /// Check the export record and `tokenizer.json` (see `provenance`), load the CUDA + /// library and the model, and prepare CUDA graphs (see `Options`). pub async fn load(dir: &Path, opts: &Options) -> Result { let started = Instant::now(); + let provenance = provenance(dir)?; let prompter = Prompter::load(dir)?; let (letters, hidden) = letter_rows(dir, &prompter.letter_ids)?; let dir = dir.to_path_buf(); @@ -204,52 +205,27 @@ impl Engine { hidden, load_seconds: started.elapsed().as_secs_f64(), device: "cuda".to_string(), + provenance, }) } - fn lock(&self) -> Result> { - self.model - .lock() - .map_err(|_| anyhow::anyhow!("model lock poisoned")) - } - /// The longest prompt that runs as a CUDA graph (0: none). pub fn graph_max_tokens(&self) -> usize { - self.lock().map(|m| m.graph_max_tokens()).unwrap_or(0) + self.model.lock().map(|m| m.graph_max_tokens()).unwrap_or(0) } - /// The final-norm hidden state at the last position. - async fn last_hidden(&self, ids: Vec, mode: Mode) -> Result> { + /// Option probabilities for one prompt: the final-norm hidden state at the last + /// position times the letter rows, in float32 with float64 accumulation, then a + /// softmax over the first `n_options` letters. + pub async fn score(&self, ids: Vec, n_options: usize) -> Result> { let model = self.model.clone(); - tokio::task::spawn_blocking(move || { + let last = tokio::task::spawn_blocking(move || { let mut model = model .lock() .map_err(|_| anyhow::anyhow!("model lock poisoned"))?; - model.forward(&ids, mode) + model.forward(&ids) }) - .await? - } - - /// One forward pass on the calling thread, for timing. - pub fn last_hidden_blocking(&self, ids: &[u32]) -> Result> { - self.lock()?.forward(ids, Mode::Auto) - } - - /// Option probabilities for one prompt: the final-norm hidden state at the last - /// position times the letter rows, in float32 with float64 accumulation, then a - /// softmax over the first `n_options` letters. - pub async fn score(&self, ids: Vec, n_options: usize) -> Result> { - self.score_mode(ids, n_options, Mode::Auto).await - } - - /// As `score`, run eagerly or from a graph. - pub async fn score_mode( - &self, - ids: Vec, - n_options: usize, - mode: Mode, - ) -> Result> { - let last = self.last_hidden(ids, mode).await?; + .await??; if last.len() != self.hidden { bail!( "hidden size {} does not match the head ({})", diff --git a/src/models/cua_s1/native/src/lib.rs b/src/models/cua_s1/native/src/lib.rs index c9e478bb..3be5ae17 100644 --- a/src/models/cua_s1/native/src/lib.rs +++ b/src/models/cua_s1/native/src/lib.rs @@ -6,6 +6,5 @@ pub mod contract; pub mod cuda; pub mod engine; pub mod model; -pub mod printable; pub mod pyjson; pub mod server; diff --git a/src/models/cua_s1/native/src/main.rs b/src/models/cua_s1/native/src/main.rs index 81240c64..3ade9523 100644 --- a/src/models/cua_s1/native/src/main.rs +++ b/src/models/cua_s1/native/src/main.rs @@ -2,11 +2,10 @@ //! //! omni-cua-s1-native --model [--port 8000] //! -//! See recipe/cua_s1/native.md for building the CUDA library, exporting the merged -//! checkpoint and checking the worker. +//! See recipe/cua_s1/native.md for building the CUDA library and exporting the merged +//! checkpoint. -use std::io::{BufRead, Write}; -use std::path::{Path, PathBuf}; +use std::path::PathBuf; use std::sync::Arc; use std::time::Instant; @@ -14,10 +13,11 @@ use anyhow::{Result, ensure}; use clap::Parser; use clap::builder::RangedU64ValueParser; -use omni_cua_s1_native::contract::{self, map_request, parse_body}; +use omni_cua_s1_native::contract; use omni_cua_s1_native::cuda; -use omni_cua_s1_native::engine::{self, Engine, Mode, Options, Prompter}; -use omni_cua_s1_native::server::{self, App, DecideError, Limits}; +use omni_cua_s1_native::engine::Engine; +use omni_cua_s1_native::model::Options; +use omni_cua_s1_native::server::{self, App, Limits}; #[derive(Parser)] #[command(about = "Cua-S1 4B 0.2 text worker on native CUDA kernels")] @@ -62,23 +62,6 @@ struct Args { /// minute; use it with --gemm-plans so that it runs once. #[arg(long, env = "CUA_S1_GEMM_SEARCH")] gemm_search: bool, - /// Read request bodies on stdin (one JSON string per line), print the prompt - /// token ids or the rejection for each, and exit. Loads only the tokenizer. - #[arg(long)] - encode_only: bool, - /// Score every question of the request bodies in this JSON file (name -> body), - /// eagerly and as served, print one JSON line per run, and exit. - #[arg(long)] - score_all: Option, - /// Time the forward pass over the token ids in each of these JSON files (lists - /// of integers), without HTTP, and exit. - #[arg(long, num_args = 1..)] - bench: Vec, - #[arg(long, default_value_t = 50, value_parser = RangedU64ValueParser::::new().range(1..))] - bench_repeat: usize, - /// Idle time before each timed forward pass, as between separate requests. - #[arg(long, default_value_t = 0)] - bench_gap_ms: u64, } impl Args { @@ -100,135 +83,24 @@ impl Args { } } -fn encode_only(args: &Args) -> Result<()> { - let prompter = Prompter::load(&args.model)?; - let mut out = std::io::stdout().lock(); - for line in std::io::stdin().lock().lines() { - let body: String = serde_json::from_str(&line?)?; - let result = parse_body(body.as_bytes()) - .and_then(|b| map_request(&b, args.max_questions)) - .map_err(DecideError::Request) - .and_then(|r| { - let ids = server::encode_all(&prompter, &r, args.max_prompt_tokens)?; - Ok((r, ids)) - }); - let record = match result { - Ok((request, ids)) => serde_json::json!({ - "status": 200, - "questions": request.questions.iter().map(|q| &q.name).collect::>(), - "ids": ids, - }), - Err(DecideError::Request(e)) => { - serde_json::json!({"status": e.status, "detail": e.message}) - } - Err(DecideError::Internal(e)) => return Err(e), - }; - writeln!(out, "{record}")?; - } - Ok(()) -} - -async fn score_all(args: &Args, inputs: &Path) -> Result<()> { - let cases: serde_json::Map = - serde_json::from_str(&std::fs::read_to_string(inputs)?)?; - let engine = Engine::load(&args.model, &args.options()?).await?; - let graph_max = engine.graph_max_tokens(); - let mut out = std::io::stdout().lock(); - for (case, body) in &cases { - let raw = serde_json::to_vec(body)?; - let request = parse_body(&raw) - .and_then(|b| map_request(&b, args.max_questions)) - .map_err(|e| anyhow::anyhow!("{case}: {}", e.message))?; - for question in &request.questions { - let ids = engine.prompter.encode(&request.state, question)?; - for (mode, label) in [(Mode::Eager, "eager"), (Mode::Auto, "served")] { - let probs = engine - .score_mode(ids.clone(), question.keys.len(), mode) - .await?; - let probabilities: serde_json::Map = question - .keys - .iter() - .zip(&probs) - .map(|(k, p)| (k.clone(), serde_json::json!(p))) - .collect(); - writeln!( - out, - "{}", - serde_json::json!({ - "case": case, - "question": question.name, - "tokens": ids.len(), - "mode": label, - "graph": mode == Mode::Auto && ids.len() <= graph_max, - "probabilities": probabilities, - }) - )?; - } - } - } - Ok(()) -} - -async fn bench(args: &Args) -> Result<()> { - let engine = Engine::load(&args.model, &args.options()?).await?; - println!( - "loaded in {:.1} s, graphs up to {} tokens", - engine.load_seconds, - engine.graph_max_tokens() - ); - for file in &args.bench { - let ids: Vec = serde_json::from_str(&std::fs::read_to_string(file)?)?; - let mut times = Vec::with_capacity(args.bench_repeat); - for i in 0..args.bench_repeat + 3 { - std::thread::sleep(std::time::Duration::from_millis(args.bench_gap_ms)); - let started = Instant::now(); - engine.last_hidden_blocking(&ids)?; - if i >= 3 { - times.push(started.elapsed().as_secs_f64() * 1e3); - } - } - times.sort_by(f64::total_cmp); - let at = |q: f64| times[((times.len() - 1) as f64 * q).round() as usize]; - println!( - "{}: {} tokens, forward p50 {:.2} ms p95 {:.2} ms min {:.2} ms", - file.display(), - ids.len(), - at(0.5), - at(0.95), - times[0] - ); - } - Ok(()) -} - #[tokio::main] async fn main() -> Result<()> { let args = Args::parse(); - if args.encode_only { - return encode_only(&args); - } - if let Some(inputs) = &args.score_all { - return score_all(&args, inputs).await; - } - if !args.bench.is_empty() { - return bench(&args).await; - } - let provenance = engine::provenance(&args.model)?; - if provenance.adapter_revision != contract::ADAPTER_REVISION { + let engine = Engine::load(&args.model, &args.options()?).await?; + if engine.provenance.adapter_revision != contract::ADAPTER_REVISION { eprintln!( "warning: adapter revision {} is not the pinned {}", - provenance.adapter_revision, + engine.provenance.adapter_revision, contract::ADAPTER_REVISION ); } - if provenance.base_revision != contract::BASE_REVISION { + if engine.provenance.base_revision != contract::BASE_REVISION { eprintln!( "warning: base revision {} is not the pinned {}", - provenance.base_revision, + engine.provenance.base_revision, contract::BASE_REVISION ); } - let engine = Engine::load(&args.model, &args.options()?).await?; println!( "loaded in {:.1} s on {} (bfloat16, graphs up to {} tokens)", engine.load_seconds, @@ -241,12 +113,8 @@ async fn main() -> Result<()> { max_questions: args.max_questions, max_prompt_tokens: args.max_prompt_tokens, }; - let app = Arc::new(App::new( - engine, - limits, - api_key, - &provenance.adapter_revision, - )); + let revision = engine.provenance.adapter_revision.clone(); + let app = Arc::new(App::new(engine, limits, api_key, &revision)); let started = Instant::now(); server::warmup(&app).await?; println!("warmed up in {:.1} s", started.elapsed().as_secs_f64()); diff --git a/src/models/cua_s1/native/src/model.rs b/src/models/cua_s1/native/src/model.rs index e8a583f9..9321c6e8 100644 --- a/src/models/cua_s1/native/src/model.rs +++ b/src/models/cua_s1/native/src/model.rs @@ -558,14 +558,6 @@ impl Scratch { } } -/// How to run one forward pass. -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum Mode { - /// From the graph for this length when graphs are on and it fits, else eagerly. - Auto, - Eager, -} - pub struct Model { pub cfg: Config, _weights: Weights, @@ -878,7 +870,8 @@ impl Model { } /// The final-norm hidden state at the last position, as float32. - pub fn forward(&mut self, ids: &[u32], mode: Mode) -> Result> { + /// Runs from the graph for this length when graphs are on and it fits, else eagerly. + pub fn forward(&mut self, ids: &[u32]) -> Result> { let t = ids.len(); ensure!(t > 0, "empty prompt"); let (vocab, h) = (self.embed.shape[0], self.cfg.hidden); @@ -896,7 +889,6 @@ impl Model { self.stream, )?); } - let use_graph = mode == Mode::Auto && in_graph_scratch; let s = if in_graph_scratch { self.graph_scratch.as_ref() } else { @@ -906,7 +898,7 @@ impl Model { let ids32: Vec = ids.iter().flat_map(|&i| (i as i32).to_le_bytes()).collect(); // SAFETY: the ids buffer holds at least t int32 values. unsafe { cuda::upload(s.at(s.ids), &ids32, self.stream)? }; - if use_graph { + if in_graph_scratch { self.graph_for(t)?; self.graphs[&t].launch(self.stream)?; } else { diff --git a/src/models/cua_s1/native/src/printable.rs b/src/models/cua_s1/native/src/printable.rs deleted file mode 100644 index 9bcd2f9a..00000000 --- a/src/models/cua_s1/native/src/printable.rs +++ /dev/null @@ -1,717 +0,0 @@ -//! Generated from Python 3.12.13 (Unicode 15.0.0): code points >= 0x80 for which -//! `str.isprintable()` is false, as inclusive ranges. Python's `repr` escapes these. -//! Regenerate with tests/make_printable.py. - -pub const NON_PRINTABLE: &[(u32, u32)] = &[ - (0x80, 0xa0), - (0xad, 0xad), - (0x378, 0x379), - (0x380, 0x383), - (0x38b, 0x38b), - (0x38d, 0x38d), - (0x3a2, 0x3a2), - (0x530, 0x530), - (0x557, 0x558), - (0x58b, 0x58c), - (0x590, 0x590), - (0x5c8, 0x5cf), - (0x5eb, 0x5ee), - (0x5f5, 0x605), - (0x61c, 0x61c), - (0x6dd, 0x6dd), - (0x70e, 0x70f), - (0x74b, 0x74c), - (0x7b2, 0x7bf), - (0x7fb, 0x7fc), - (0x82e, 0x82f), - (0x83f, 0x83f), - (0x85c, 0x85d), - (0x85f, 0x85f), - (0x86b, 0x86f), - (0x88f, 0x897), - (0x8e2, 0x8e2), - (0x984, 0x984), - (0x98d, 0x98e), - (0x991, 0x992), - (0x9a9, 0x9a9), - (0x9b1, 0x9b1), - (0x9b3, 0x9b5), - (0x9ba, 0x9bb), - (0x9c5, 0x9c6), - (0x9c9, 0x9ca), - (0x9cf, 0x9d6), - (0x9d8, 0x9db), - (0x9de, 0x9de), - (0x9e4, 0x9e5), - (0x9ff, 0xa00), - (0xa04, 0xa04), - (0xa0b, 0xa0e), - (0xa11, 0xa12), - (0xa29, 0xa29), - (0xa31, 0xa31), - (0xa34, 0xa34), - (0xa37, 0xa37), - (0xa3a, 0xa3b), - (0xa3d, 0xa3d), - (0xa43, 0xa46), - (0xa49, 0xa4a), - (0xa4e, 0xa50), - (0xa52, 0xa58), - (0xa5d, 0xa5d), - (0xa5f, 0xa65), - (0xa77, 0xa80), - (0xa84, 0xa84), - (0xa8e, 0xa8e), - (0xa92, 0xa92), - (0xaa9, 0xaa9), - (0xab1, 0xab1), - (0xab4, 0xab4), - (0xaba, 0xabb), - (0xac6, 0xac6), - (0xaca, 0xaca), - (0xace, 0xacf), - (0xad1, 0xadf), - (0xae4, 0xae5), - (0xaf2, 0xaf8), - (0xb00, 0xb00), - (0xb04, 0xb04), - (0xb0d, 0xb0e), - (0xb11, 0xb12), - (0xb29, 0xb29), - (0xb31, 0xb31), - (0xb34, 0xb34), - (0xb3a, 0xb3b), - (0xb45, 0xb46), - (0xb49, 0xb4a), - (0xb4e, 0xb54), - (0xb58, 0xb5b), - (0xb5e, 0xb5e), - (0xb64, 0xb65), - (0xb78, 0xb81), - (0xb84, 0xb84), - (0xb8b, 0xb8d), - (0xb91, 0xb91), - (0xb96, 0xb98), - (0xb9b, 0xb9b), - (0xb9d, 0xb9d), - (0xba0, 0xba2), - (0xba5, 0xba7), - (0xbab, 0xbad), - (0xbba, 0xbbd), - (0xbc3, 0xbc5), - (0xbc9, 0xbc9), - (0xbce, 0xbcf), - (0xbd1, 0xbd6), - (0xbd8, 0xbe5), - (0xbfb, 0xbff), - (0xc0d, 0xc0d), - (0xc11, 0xc11), - (0xc29, 0xc29), - (0xc3a, 0xc3b), - (0xc45, 0xc45), - (0xc49, 0xc49), - (0xc4e, 0xc54), - (0xc57, 0xc57), - (0xc5b, 0xc5c), - (0xc5e, 0xc5f), - (0xc64, 0xc65), - (0xc70, 0xc76), - (0xc8d, 0xc8d), - (0xc91, 0xc91), - (0xca9, 0xca9), - (0xcb4, 0xcb4), - (0xcba, 0xcbb), - (0xcc5, 0xcc5), - (0xcc9, 0xcc9), - (0xcce, 0xcd4), - (0xcd7, 0xcdc), - (0xcdf, 0xcdf), - (0xce4, 0xce5), - (0xcf0, 0xcf0), - (0xcf4, 0xcff), - (0xd0d, 0xd0d), - (0xd11, 0xd11), - (0xd45, 0xd45), - (0xd49, 0xd49), - (0xd50, 0xd53), - (0xd64, 0xd65), - (0xd80, 0xd80), - (0xd84, 0xd84), - (0xd97, 0xd99), - (0xdb2, 0xdb2), - (0xdbc, 0xdbc), - (0xdbe, 0xdbf), - (0xdc7, 0xdc9), - (0xdcb, 0xdce), - (0xdd5, 0xdd5), - (0xdd7, 0xdd7), - (0xde0, 0xde5), - (0xdf0, 0xdf1), - (0xdf5, 0xe00), - (0xe3b, 0xe3e), - (0xe5c, 0xe80), - (0xe83, 0xe83), - (0xe85, 0xe85), - (0xe8b, 0xe8b), - (0xea4, 0xea4), - (0xea6, 0xea6), - (0xebe, 0xebf), - (0xec5, 0xec5), - (0xec7, 0xec7), - (0xecf, 0xecf), - (0xeda, 0xedb), - (0xee0, 0xeff), - (0xf48, 0xf48), - (0xf6d, 0xf70), - (0xf98, 0xf98), - (0xfbd, 0xfbd), - (0xfcd, 0xfcd), - (0xfdb, 0xfff), - (0x10c6, 0x10c6), - (0x10c8, 0x10cc), - (0x10ce, 0x10cf), - (0x1249, 0x1249), - (0x124e, 0x124f), - (0x1257, 0x1257), - (0x1259, 0x1259), - (0x125e, 0x125f), - (0x1289, 0x1289), - (0x128e, 0x128f), - (0x12b1, 0x12b1), - (0x12b6, 0x12b7), - (0x12bf, 0x12bf), - (0x12c1, 0x12c1), - (0x12c6, 0x12c7), - (0x12d7, 0x12d7), - (0x1311, 0x1311), - (0x1316, 0x1317), - (0x135b, 0x135c), - (0x137d, 0x137f), - (0x139a, 0x139f), - (0x13f6, 0x13f7), - (0x13fe, 0x13ff), - (0x1680, 0x1680), - (0x169d, 0x169f), - (0x16f9, 0x16ff), - (0x1716, 0x171e), - (0x1737, 0x173f), - (0x1754, 0x175f), - (0x176d, 0x176d), - (0x1771, 0x1771), - (0x1774, 0x177f), - (0x17de, 0x17df), - (0x17ea, 0x17ef), - (0x17fa, 0x17ff), - (0x180e, 0x180e), - (0x181a, 0x181f), - (0x1879, 0x187f), - (0x18ab, 0x18af), - (0x18f6, 0x18ff), - (0x191f, 0x191f), - (0x192c, 0x192f), - (0x193c, 0x193f), - (0x1941, 0x1943), - (0x196e, 0x196f), - (0x1975, 0x197f), - (0x19ac, 0x19af), - (0x19ca, 0x19cf), - (0x19db, 0x19dd), - (0x1a1c, 0x1a1d), - (0x1a5f, 0x1a5f), - (0x1a7d, 0x1a7e), - (0x1a8a, 0x1a8f), - (0x1a9a, 0x1a9f), - (0x1aae, 0x1aaf), - (0x1acf, 0x1aff), - (0x1b4d, 0x1b4f), - (0x1b7f, 0x1b7f), - (0x1bf4, 0x1bfb), - (0x1c38, 0x1c3a), - (0x1c4a, 0x1c4c), - (0x1c89, 0x1c8f), - (0x1cbb, 0x1cbc), - (0x1cc8, 0x1ccf), - (0x1cfb, 0x1cff), - (0x1f16, 0x1f17), - (0x1f1e, 0x1f1f), - (0x1f46, 0x1f47), - (0x1f4e, 0x1f4f), - (0x1f58, 0x1f58), - (0x1f5a, 0x1f5a), - (0x1f5c, 0x1f5c), - (0x1f5e, 0x1f5e), - (0x1f7e, 0x1f7f), - (0x1fb5, 0x1fb5), - (0x1fc5, 0x1fc5), - (0x1fd4, 0x1fd5), - (0x1fdc, 0x1fdc), - (0x1ff0, 0x1ff1), - (0x1ff5, 0x1ff5), - (0x1fff, 0x200f), - (0x2028, 0x202f), - (0x205f, 0x206f), - (0x2072, 0x2073), - (0x208f, 0x208f), - (0x209d, 0x209f), - (0x20c1, 0x20cf), - (0x20f1, 0x20ff), - (0x218c, 0x218f), - (0x2427, 0x243f), - (0x244b, 0x245f), - (0x2b74, 0x2b75), - (0x2b96, 0x2b96), - (0x2cf4, 0x2cf8), - (0x2d26, 0x2d26), - (0x2d28, 0x2d2c), - (0x2d2e, 0x2d2f), - (0x2d68, 0x2d6e), - (0x2d71, 0x2d7e), - (0x2d97, 0x2d9f), - (0x2da7, 0x2da7), - (0x2daf, 0x2daf), - (0x2db7, 0x2db7), - (0x2dbf, 0x2dbf), - (0x2dc7, 0x2dc7), - (0x2dcf, 0x2dcf), - (0x2dd7, 0x2dd7), - (0x2ddf, 0x2ddf), - (0x2e5e, 0x2e7f), - (0x2e9a, 0x2e9a), - (0x2ef4, 0x2eff), - (0x2fd6, 0x2fef), - (0x2ffc, 0x3000), - (0x3040, 0x3040), - (0x3097, 0x3098), - (0x3100, 0x3104), - (0x3130, 0x3130), - (0x318f, 0x318f), - (0x31e4, 0x31ef), - (0x321f, 0x321f), - (0xa48d, 0xa48f), - (0xa4c7, 0xa4cf), - (0xa62c, 0xa63f), - (0xa6f8, 0xa6ff), - (0xa7cb, 0xa7cf), - (0xa7d2, 0xa7d2), - (0xa7d4, 0xa7d4), - (0xa7da, 0xa7f1), - (0xa82d, 0xa82f), - (0xa83a, 0xa83f), - (0xa878, 0xa87f), - (0xa8c6, 0xa8cd), - (0xa8da, 0xa8df), - (0xa954, 0xa95e), - (0xa97d, 0xa97f), - (0xa9ce, 0xa9ce), - (0xa9da, 0xa9dd), - (0xa9ff, 0xa9ff), - (0xaa37, 0xaa3f), - (0xaa4e, 0xaa4f), - (0xaa5a, 0xaa5b), - (0xaac3, 0xaada), - (0xaaf7, 0xab00), - (0xab07, 0xab08), - (0xab0f, 0xab10), - (0xab17, 0xab1f), - (0xab27, 0xab27), - (0xab2f, 0xab2f), - (0xab6c, 0xab6f), - (0xabee, 0xabef), - (0xabfa, 0xabff), - (0xd7a4, 0xd7af), - (0xd7c7, 0xd7ca), - (0xd7fc, 0xf8ff), - (0xfa6e, 0xfa6f), - (0xfada, 0xfaff), - (0xfb07, 0xfb12), - (0xfb18, 0xfb1c), - (0xfb37, 0xfb37), - (0xfb3d, 0xfb3d), - (0xfb3f, 0xfb3f), - (0xfb42, 0xfb42), - (0xfb45, 0xfb45), - (0xfbc3, 0xfbd2), - (0xfd90, 0xfd91), - (0xfdc8, 0xfdce), - (0xfdd0, 0xfdef), - (0xfe1a, 0xfe1f), - (0xfe53, 0xfe53), - (0xfe67, 0xfe67), - (0xfe6c, 0xfe6f), - (0xfe75, 0xfe75), - (0xfefd, 0xff00), - (0xffbf, 0xffc1), - (0xffc8, 0xffc9), - (0xffd0, 0xffd1), - (0xffd8, 0xffd9), - (0xffdd, 0xffdf), - (0xffe7, 0xffe7), - (0xffef, 0xfffb), - (0xfffe, 0xffff), - (0x1000c, 0x1000c), - (0x10027, 0x10027), - (0x1003b, 0x1003b), - (0x1003e, 0x1003e), - (0x1004e, 0x1004f), - (0x1005e, 0x1007f), - (0x100fb, 0x100ff), - (0x10103, 0x10106), - (0x10134, 0x10136), - (0x1018f, 0x1018f), - (0x1019d, 0x1019f), - (0x101a1, 0x101cf), - (0x101fe, 0x1027f), - (0x1029d, 0x1029f), - (0x102d1, 0x102df), - (0x102fc, 0x102ff), - (0x10324, 0x1032c), - (0x1034b, 0x1034f), - (0x1037b, 0x1037f), - (0x1039e, 0x1039e), - (0x103c4, 0x103c7), - (0x103d6, 0x103ff), - (0x1049e, 0x1049f), - (0x104aa, 0x104af), - (0x104d4, 0x104d7), - (0x104fc, 0x104ff), - (0x10528, 0x1052f), - (0x10564, 0x1056e), - (0x1057b, 0x1057b), - (0x1058b, 0x1058b), - (0x10593, 0x10593), - (0x10596, 0x10596), - (0x105a2, 0x105a2), - (0x105b2, 0x105b2), - (0x105ba, 0x105ba), - (0x105bd, 0x105ff), - (0x10737, 0x1073f), - (0x10756, 0x1075f), - (0x10768, 0x1077f), - (0x10786, 0x10786), - (0x107b1, 0x107b1), - (0x107bb, 0x107ff), - (0x10806, 0x10807), - (0x10809, 0x10809), - (0x10836, 0x10836), - (0x10839, 0x1083b), - (0x1083d, 0x1083e), - (0x10856, 0x10856), - (0x1089f, 0x108a6), - (0x108b0, 0x108df), - (0x108f3, 0x108f3), - (0x108f6, 0x108fa), - (0x1091c, 0x1091e), - (0x1093a, 0x1093e), - (0x10940, 0x1097f), - (0x109b8, 0x109bb), - (0x109d0, 0x109d1), - (0x10a04, 0x10a04), - (0x10a07, 0x10a0b), - (0x10a14, 0x10a14), - (0x10a18, 0x10a18), - (0x10a36, 0x10a37), - (0x10a3b, 0x10a3e), - (0x10a49, 0x10a4f), - (0x10a59, 0x10a5f), - (0x10aa0, 0x10abf), - (0x10ae7, 0x10aea), - (0x10af7, 0x10aff), - (0x10b36, 0x10b38), - (0x10b56, 0x10b57), - (0x10b73, 0x10b77), - (0x10b92, 0x10b98), - (0x10b9d, 0x10ba8), - (0x10bb0, 0x10bff), - (0x10c49, 0x10c7f), - (0x10cb3, 0x10cbf), - (0x10cf3, 0x10cf9), - (0x10d28, 0x10d2f), - (0x10d3a, 0x10e5f), - (0x10e7f, 0x10e7f), - (0x10eaa, 0x10eaa), - (0x10eae, 0x10eaf), - (0x10eb2, 0x10efc), - (0x10f28, 0x10f2f), - (0x10f5a, 0x10f6f), - (0x10f8a, 0x10faf), - (0x10fcc, 0x10fdf), - (0x10ff7, 0x10fff), - (0x1104e, 0x11051), - (0x11076, 0x1107e), - (0x110bd, 0x110bd), - (0x110c3, 0x110cf), - (0x110e9, 0x110ef), - (0x110fa, 0x110ff), - (0x11135, 0x11135), - (0x11148, 0x1114f), - (0x11177, 0x1117f), - (0x111e0, 0x111e0), - (0x111f5, 0x111ff), - (0x11212, 0x11212), - (0x11242, 0x1127f), - (0x11287, 0x11287), - (0x11289, 0x11289), - (0x1128e, 0x1128e), - (0x1129e, 0x1129e), - (0x112aa, 0x112af), - (0x112eb, 0x112ef), - (0x112fa, 0x112ff), - (0x11304, 0x11304), - (0x1130d, 0x1130e), - (0x11311, 0x11312), - (0x11329, 0x11329), - (0x11331, 0x11331), - (0x11334, 0x11334), - (0x1133a, 0x1133a), - (0x11345, 0x11346), - (0x11349, 0x1134a), - (0x1134e, 0x1134f), - (0x11351, 0x11356), - (0x11358, 0x1135c), - (0x11364, 0x11365), - (0x1136d, 0x1136f), - (0x11375, 0x113ff), - (0x1145c, 0x1145c), - (0x11462, 0x1147f), - (0x114c8, 0x114cf), - (0x114da, 0x1157f), - (0x115b6, 0x115b7), - (0x115de, 0x115ff), - (0x11645, 0x1164f), - (0x1165a, 0x1165f), - (0x1166d, 0x1167f), - (0x116ba, 0x116bf), - (0x116ca, 0x116ff), - (0x1171b, 0x1171c), - (0x1172c, 0x1172f), - (0x11747, 0x117ff), - (0x1183c, 0x1189f), - (0x118f3, 0x118fe), - (0x11907, 0x11908), - (0x1190a, 0x1190b), - (0x11914, 0x11914), - (0x11917, 0x11917), - (0x11936, 0x11936), - (0x11939, 0x1193a), - (0x11947, 0x1194f), - (0x1195a, 0x1199f), - (0x119a8, 0x119a9), - (0x119d8, 0x119d9), - (0x119e5, 0x119ff), - (0x11a48, 0x11a4f), - (0x11aa3, 0x11aaf), - (0x11af9, 0x11aff), - (0x11b0a, 0x11bff), - (0x11c09, 0x11c09), - (0x11c37, 0x11c37), - (0x11c46, 0x11c4f), - (0x11c6d, 0x11c6f), - (0x11c90, 0x11c91), - (0x11ca8, 0x11ca8), - (0x11cb7, 0x11cff), - (0x11d07, 0x11d07), - (0x11d0a, 0x11d0a), - (0x11d37, 0x11d39), - (0x11d3b, 0x11d3b), - (0x11d3e, 0x11d3e), - (0x11d48, 0x11d4f), - (0x11d5a, 0x11d5f), - (0x11d66, 0x11d66), - (0x11d69, 0x11d69), - (0x11d8f, 0x11d8f), - (0x11d92, 0x11d92), - (0x11d99, 0x11d9f), - (0x11daa, 0x11edf), - (0x11ef9, 0x11eff), - (0x11f11, 0x11f11), - (0x11f3b, 0x11f3d), - (0x11f5a, 0x11faf), - (0x11fb1, 0x11fbf), - (0x11ff2, 0x11ffe), - (0x1239a, 0x123ff), - (0x1246f, 0x1246f), - (0x12475, 0x1247f), - (0x12544, 0x12f8f), - (0x12ff3, 0x12fff), - (0x13430, 0x1343f), - (0x13456, 0x143ff), - (0x14647, 0x167ff), - (0x16a39, 0x16a3f), - (0x16a5f, 0x16a5f), - (0x16a6a, 0x16a6d), - (0x16abf, 0x16abf), - (0x16aca, 0x16acf), - (0x16aee, 0x16aef), - (0x16af6, 0x16aff), - (0x16b46, 0x16b4f), - (0x16b5a, 0x16b5a), - (0x16b62, 0x16b62), - (0x16b78, 0x16b7c), - (0x16b90, 0x16e3f), - (0x16e9b, 0x16eff), - (0x16f4b, 0x16f4e), - (0x16f88, 0x16f8e), - (0x16fa0, 0x16fdf), - (0x16fe5, 0x16fef), - (0x16ff2, 0x16fff), - (0x187f8, 0x187ff), - (0x18cd6, 0x18cff), - (0x18d09, 0x1afef), - (0x1aff4, 0x1aff4), - (0x1affc, 0x1affc), - (0x1afff, 0x1afff), - (0x1b123, 0x1b131), - (0x1b133, 0x1b14f), - (0x1b153, 0x1b154), - (0x1b156, 0x1b163), - (0x1b168, 0x1b16f), - (0x1b2fc, 0x1bbff), - (0x1bc6b, 0x1bc6f), - (0x1bc7d, 0x1bc7f), - (0x1bc89, 0x1bc8f), - (0x1bc9a, 0x1bc9b), - (0x1bca0, 0x1ceff), - (0x1cf2e, 0x1cf2f), - (0x1cf47, 0x1cf4f), - (0x1cfc4, 0x1cfff), - (0x1d0f6, 0x1d0ff), - (0x1d127, 0x1d128), - (0x1d173, 0x1d17a), - (0x1d1eb, 0x1d1ff), - (0x1d246, 0x1d2bf), - (0x1d2d4, 0x1d2df), - (0x1d2f4, 0x1d2ff), - (0x1d357, 0x1d35f), - (0x1d379, 0x1d3ff), - (0x1d455, 0x1d455), - (0x1d49d, 0x1d49d), - (0x1d4a0, 0x1d4a1), - (0x1d4a3, 0x1d4a4), - (0x1d4a7, 0x1d4a8), - (0x1d4ad, 0x1d4ad), - (0x1d4ba, 0x1d4ba), - (0x1d4bc, 0x1d4bc), - (0x1d4c4, 0x1d4c4), - (0x1d506, 0x1d506), - (0x1d50b, 0x1d50c), - (0x1d515, 0x1d515), - (0x1d51d, 0x1d51d), - (0x1d53a, 0x1d53a), - (0x1d53f, 0x1d53f), - (0x1d545, 0x1d545), - (0x1d547, 0x1d549), - (0x1d551, 0x1d551), - (0x1d6a6, 0x1d6a7), - (0x1d7cc, 0x1d7cd), - (0x1da8c, 0x1da9a), - (0x1daa0, 0x1daa0), - (0x1dab0, 0x1deff), - (0x1df1f, 0x1df24), - (0x1df2b, 0x1dfff), - (0x1e007, 0x1e007), - (0x1e019, 0x1e01a), - (0x1e022, 0x1e022), - (0x1e025, 0x1e025), - (0x1e02b, 0x1e02f), - (0x1e06e, 0x1e08e), - (0x1e090, 0x1e0ff), - (0x1e12d, 0x1e12f), - (0x1e13e, 0x1e13f), - (0x1e14a, 0x1e14d), - (0x1e150, 0x1e28f), - (0x1e2af, 0x1e2bf), - (0x1e2fa, 0x1e2fe), - (0x1e300, 0x1e4cf), - (0x1e4fa, 0x1e7df), - (0x1e7e7, 0x1e7e7), - (0x1e7ec, 0x1e7ec), - (0x1e7ef, 0x1e7ef), - (0x1e7ff, 0x1e7ff), - (0x1e8c5, 0x1e8c6), - (0x1e8d7, 0x1e8ff), - (0x1e94c, 0x1e94f), - (0x1e95a, 0x1e95d), - (0x1e960, 0x1ec70), - (0x1ecb5, 0x1ed00), - (0x1ed3e, 0x1edff), - (0x1ee04, 0x1ee04), - (0x1ee20, 0x1ee20), - (0x1ee23, 0x1ee23), - (0x1ee25, 0x1ee26), - (0x1ee28, 0x1ee28), - (0x1ee33, 0x1ee33), - (0x1ee38, 0x1ee38), - (0x1ee3a, 0x1ee3a), - (0x1ee3c, 0x1ee41), - (0x1ee43, 0x1ee46), - (0x1ee48, 0x1ee48), - (0x1ee4a, 0x1ee4a), - (0x1ee4c, 0x1ee4c), - (0x1ee50, 0x1ee50), - (0x1ee53, 0x1ee53), - (0x1ee55, 0x1ee56), - (0x1ee58, 0x1ee58), - (0x1ee5a, 0x1ee5a), - (0x1ee5c, 0x1ee5c), - (0x1ee5e, 0x1ee5e), - (0x1ee60, 0x1ee60), - (0x1ee63, 0x1ee63), - (0x1ee65, 0x1ee66), - (0x1ee6b, 0x1ee6b), - (0x1ee73, 0x1ee73), - (0x1ee78, 0x1ee78), - (0x1ee7d, 0x1ee7d), - (0x1ee7f, 0x1ee7f), - (0x1ee8a, 0x1ee8a), - (0x1ee9c, 0x1eea0), - (0x1eea4, 0x1eea4), - (0x1eeaa, 0x1eeaa), - (0x1eebc, 0x1eeef), - (0x1eef2, 0x1efff), - (0x1f02c, 0x1f02f), - (0x1f094, 0x1f09f), - (0x1f0af, 0x1f0b0), - (0x1f0c0, 0x1f0c0), - (0x1f0d0, 0x1f0d0), - (0x1f0f6, 0x1f0ff), - (0x1f1ae, 0x1f1e5), - (0x1f203, 0x1f20f), - (0x1f23c, 0x1f23f), - (0x1f249, 0x1f24f), - (0x1f252, 0x1f25f), - (0x1f266, 0x1f2ff), - (0x1f6d8, 0x1f6db), - (0x1f6ed, 0x1f6ef), - (0x1f6fd, 0x1f6ff), - (0x1f777, 0x1f77a), - (0x1f7da, 0x1f7df), - (0x1f7ec, 0x1f7ef), - (0x1f7f1, 0x1f7ff), - (0x1f80c, 0x1f80f), - (0x1f848, 0x1f84f), - (0x1f85a, 0x1f85f), - (0x1f888, 0x1f88f), - (0x1f8ae, 0x1f8af), - (0x1f8b2, 0x1f8ff), - (0x1fa54, 0x1fa5f), - (0x1fa6e, 0x1fa6f), - (0x1fa7d, 0x1fa7f), - (0x1fa89, 0x1fa8f), - (0x1fabe, 0x1fabe), - (0x1fac6, 0x1facd), - (0x1fadc, 0x1fadf), - (0x1fae9, 0x1faef), - (0x1faf9, 0x1faff), - (0x1fb93, 0x1fb93), - (0x1fbcb, 0x1fbef), - (0x1fbfa, 0x1ffff), - (0x2a6e0, 0x2a6ff), - (0x2b73a, 0x2b73f), - (0x2b81e, 0x2b81f), - (0x2cea2, 0x2ceaf), - (0x2ebe1, 0x2f7ff), - (0x2fa1e, 0x2ffff), - (0x3134b, 0x3134f), - (0x323b0, 0xe00ff), - (0xe01f0, 0x10ffff), -]; diff --git a/src/models/cua_s1/native/src/pyjson.rs b/src/models/cua_s1/native/src/pyjson.rs index 4cda98e4..7c60e1d5 100644 --- a/src/models/cua_s1/native/src/pyjson.rs +++ b/src/models/cua_s1/native/src/pyjson.rs @@ -4,15 +4,13 @@ //! `contract.parse_body` adds: duplicate keys, `NaN`/`Infinity`, numbers that are //! out of range for a float, lone surrogates and very deep nesting are all rejected, //! and errors surface in the same order as in Python. -//! - [`dumps`] follows `json.dumps(value, ensure_ascii=False)` (and the compact form -//! Starlette uses for responses). -//! - [`repr`] follows Python's `repr`, which the error messages quote. +//! - [`dumps`] follows `json.dumps(value, ensure_ascii=False)`. +//! - [`repr`] follows Python's `repr`, which the error messages quote, except that +//! characters outside ASCII are written as they are unless they are lone surrogates. use std::collections::HashSet; use std::fmt::Write as _; -use crate::printable::NON_PRINTABLE; - /// Nesting limits of the Python worker, measured on CPython 3.12.13 under uvicorn. /// Both come from interpreter recursion limits, so they depend on the call stack and /// are not documented constants. @@ -106,15 +104,6 @@ pub enum Value { Object(Vec<(PyStr, Value)>), } -impl Value { - pub fn get(&self, key: &str) -> Option<&Value> { - match self { - Value::Object(pairs) => pairs.iter().find(|(k, _)| k == key).map(|(_, v)| v), - _ => None, - } - } -} - /// Values can nest thousands of levels deep before the depth check rejects them, so /// they are dropped without recursion. impl Drop for Value { @@ -602,17 +591,15 @@ pub fn write_json_str(s: &str, out: &mut String) { out.push('"'); } -/// `json.dumps(value, ensure_ascii=False)`, with the default separators when -/// `compact` is false and `(",", ":")` when it is true. The value must not hold -/// lone surrogates (`parse` rejects them). -pub fn dumps(value: &Value, compact: bool) -> String { +/// `json.dumps(value, ensure_ascii=False)`. The value must not hold lone surrogates +/// (`parse` rejects them). +pub fn dumps(value: &Value) -> String { let mut out = String::new(); - write_value(value, compact, &mut out); + write_value(value, &mut out); out } -fn write_value(value: &Value, compact: bool, out: &mut String) { - let (item_sep, key_sep) = if compact { (",", ":") } else { (", ", ": ") }; +fn write_value(value: &Value, out: &mut String) { match value { Value::Null => out.push_str("null"), Value::Bool(b) => out.push_str(if *b { "true" } else { "false" }), @@ -623,9 +610,9 @@ fn write_value(value: &Value, compact: bool, out: &mut String) { out.push('['); for (i, item) in items.iter().enumerate() { if i > 0 { - out.push_str(item_sep); + out.push_str(", "); } - write_value(item, compact, out); + write_value(item, out); } out.push(']'); } @@ -633,32 +620,19 @@ fn write_value(value: &Value, compact: bool, out: &mut String) { out.push('{'); for (i, (key, item)) in pairs.iter().enumerate() { if i > 0 { - out.push_str(item_sep); + out.push_str(", "); } write_json_str(key.as_str().expect("checked UTF-8"), out); - out.push_str(key_sep); - write_value(item, compact, out); + out.push_str(": "); + write_value(item, out); } out.push('}'); } } } -fn is_printable(cp: u32) -> bool { - NON_PRINTABLE - .binary_search_by(|&(lo, hi)| { - if hi < cp { - std::cmp::Ordering::Less - } else if lo > cp { - std::cmp::Ordering::Greater - } else { - std::cmp::Ordering::Equal - } - }) - .is_err() -} - -/// Python's `repr(str)`. +/// Python's `repr(str)`, except that characters outside ASCII are kept as they are +/// unless they are lone surrogates. pub fn repr_str(s: &PyStr) -> String { let cps: Vec = s.code_points().collect(); let squote = cps.contains(&('\'' as u32)); @@ -676,11 +650,11 @@ pub fn repr_str(s: &PyStr) -> String { 0x0a => out.push_str("\\n"), 0x0d => out.push_str("\\r"), 0..=0x1f | 0x7f => write!(out, "\\x{cp:02x}").unwrap(), - 0x20..=0x7e => out.push(char::from_u32(cp).unwrap()), - _ if is_printable(cp) => out.push(char::from_u32(cp).unwrap()), - 0x80..=0xff => write!(out, "\\x{cp:02x}").unwrap(), - 0x100..=0xffff => write!(out, "\\u{cp:04x}").unwrap(), - _ => write!(out, "\\U{cp:08x}").unwrap(), + _ => match char::from_u32(cp) { + Some(c) => out.push(c), + // a lone surrogate + None => write!(out, "\\u{cp:04x}").unwrap(), + }, } } out.push(quote); @@ -761,42 +735,25 @@ mod tests { } } - #[test] - #[ignore = "needs CUA_S1_FLOAT_VECTORS"] - fn float_repr_matches_python_vectors() { - // tests/make_float_vectors.py writes " " lines - let path = std::env::var("CUA_S1_FLOAT_VECTORS") - .expect("CUA_S1_FLOAT_VECTORS must name a file from tests/make_float_vectors.py"); - let text = std::fs::read_to_string(path).unwrap(); - let mut n = 0; - for line in text.lines() { - let (bits, want) = line.split_once(' ').unwrap(); - let x = f64::from_bits(u64::from_str_radix(bits, 16).unwrap()); - assert_eq!(float_repr(x), want, "bits {bits}"); - n += 1; - } - assert!(n > 1000); - } - #[test] fn dumps_matches_python() { let v = parse(br#"{"a": [1.0, 1e16, 1e-5, 0.0001, -0.0, 1e22, 123456789012345678, 3.14e-07, -0, true, null]}"#).unwrap(); let v = Value::Object(v); assert_eq!( - dumps(&v, false), + dumps(&v), r#"{"a": [1.0, 1e+16, 1e-05, 0.0001, -0.0, 1e+22, 123456789012345678, 3.14e-07, 0, true, null]}"# ); let s = Value::Str(PyStr::new("\u{0}\u{1f}\u{7f}\u{2028}\"\\/\t\u{8}\u{c}é😀")); assert_eq!( - dumps(&s, false), + dumps(&s), "\"\\u0000\\u001f\u{7f}\u{2028}\\\"\\\\/\\t\\b\\fé😀\"" ); } #[test] - fn repr_matches_python() { - let s = PyStr::new("a\u{7f}\u{a0}\u{2028}😀é"); - assert_eq!(repr_str(&s), "'a\\x7f\\xa0\\u2028😀é'"); + fn repr_str_quotes_and_escapes() { + let s = PyStr::new("a\u{7f}\u{a0}😀é"); + assert_eq!(repr_str(&s), "'a\\x7f\u{a0}😀é'"); assert_eq!(repr_str(&PyStr::new("it's")), "\"it's\""); assert_eq!(repr_str(&PyStr::new("both'\"")), "'both\\'\"'"); let v = Value::Object(parse(br#"{"k": [1, 2.5, null, true, "x"], "e": {}}"#).unwrap()); diff --git a/src/models/cua_s1/native/src/server.rs b/src/models/cua_s1/native/src/server.rs index c087ac23..8dc87e12 100644 --- a/src/models/cua_s1/native/src/server.rs +++ b/src/models/cua_s1/native/src/server.rs @@ -1,5 +1,5 @@ //! HTTP routes, matching the Python worker: `GET /health` and `POST /v1/systemone`, -//! one decision at a time, the same status codes and the same response bytes. +//! one decision at a time, with the same status codes and response format. use std::sync::Arc; @@ -12,7 +12,7 @@ use axum::routing::{get, post}; use http_body_util::BodyExt; use crate::contract::{self, Request, RequestError, detail_json, map_request, parse_body}; -use crate::engine::{Engine, Prompter}; +use crate::engine::Engine; use crate::pyjson::{PyStr, repr_str, write_json_str}; pub struct Limits { @@ -59,28 +59,27 @@ fn error(status: u16, message: &str) -> Response { ) } -pub enum DecideError { +enum DecideError { Request(RequestError), Internal(anyhow::Error), } -/// Token ids for every question, checking the per-question prompt limit before any -/// forward pass runs. -pub fn encode_all( - prompter: &Prompter, - request: &Request, - max_prompt_tokens: usize, -) -> Result>, DecideError> { +/// Score each question and build the response body. Every prompt is tokenized and +/// checked against the prompt limit before any forward pass runs. +async fn decide(app: &App, request: &Request) -> Result { + let limit = app.limits.max_prompt_tokens; let mut encoded = Vec::with_capacity(request.questions.len()); for question in &request.questions { - let ids = prompter + let ids = app + .engine + .prompter .encode(&request.state, question) .map_err(DecideError::Internal)?; - if max_prompt_tokens > 0 && ids.len() > max_prompt_tokens { + if limit > 0 && ids.len() > limit { return Err(DecideError::Request(RequestError::new( 413, format!( - "question {}: prompt is {} tokens, over the {max_prompt_tokens}-token limit", + "question {}: prompt is {} tokens, over the {limit}-token limit", repr_str(&PyStr::new(&question.name)), ids.len() ), @@ -88,12 +87,6 @@ pub fn encode_all( } encoded.push(ids); } - Ok(encoded) -} - -/// Score each question and build the response body. -pub async fn decide(app: &App, request: &Request) -> Result { - let encoded = encode_all(&app.engine.prompter, request, app.limits.max_prompt_tokens)?; let mut answers = String::from("{"); let mut prompt_tokens = 0; for (i, (question, ids)) in request.questions.iter().zip(encoded).enumerate() { diff --git a/src/models/cua_s1/native/tests/make_float_vectors.py b/src/models/cua_s1/native/tests/make_float_vectors.py deleted file mode 100644 index 329f13fd..00000000 --- a/src/models/cua_s1/native/tests/make_float_vectors.py +++ /dev/null @@ -1,64 +0,0 @@ -"""Write float formatting cases for the pyjson tests: " " lines. - - python3 src/models/cua_s1/native/tests/make_float_vectors.py floats.txt - CUA_S1_FLOAT_VECTORS=floats.txt cargo test -p omni-cua-s1-native float_repr -- --ignored - -Special values, powers of two and ten, values next to the points where repr switches -between plain and exponent notation, random bit patterns, and random short decimals. -""" - -import math -import random -import struct -import sys - - -def bits(x: float) -> str: - return struct.pack(">d", x).hex() - - -def main() -> None: - rng = random.Random(20260928) - values = [ - 0.0, - -0.0, - 1.0, - -1.0, - 0.5, - 0.1, - 0.2, - 0.3, - 1 / 3, - 2 / 3, - math.pi, - math.e, - 5e-324, - 2.2250738585072014e-308, - 1.7976931348623157e308, - 2.0**53, - 2.0**53 + 2, - ] - values += [2.0**e for e in range(-1074, 1024, 7)] - values += [10.0**e for e in range(-323, 309)] - for e in (-5, -4, 15, 16, 17): - base = 10.0**e - values += [math.nextafter(base, 0.0), base, math.nextafter(base, math.inf)] - while len(values) < 30000: - pick = rng.random() - if pick < 0.5: - x = struct.unpack(">d", rng.getrandbits(64).to_bytes(8, "big"))[0] - elif pick < 0.8: - x = float( - f"{rng.randint(1, 10 ** rng.randint(1, 17))}e{rng.randint(-30, 30)}" - ) - else: - x = rng.uniform(-1e6, 1e6) - if math.isfinite(x): - values.append(x) - with open(sys.argv[1], "w") as f: - for x in values: - f.write(f"{bits(x)} {x!r}\n") - - -if __name__ == "__main__": - main() diff --git a/src/models/cua_s1/native/tests/make_printable.py b/src/models/cua_s1/native/tests/make_printable.py deleted file mode 100644 index 93c285fb..00000000 --- a/src/models/cua_s1/native/tests/make_printable.py +++ /dev/null @@ -1,31 +0,0 @@ -"""Write src/printable.rs: the code points >= 0x80 that Python's repr escapes. - - python3.12 src/models/cua_s1/native/tests/make_printable.py > src/models/cua_s1/native/src/printable.rs - -Use the Python version the reference worker runs on; the table follows its Unicode -database. -""" - -import sys -import unicodedata - -ranges = [] -for cp in range(0x80, sys.maxunicode + 1): - if not chr(cp).isprintable(): - if ranges and ranges[-1][1] == cp - 1: - ranges[-1][1] = cp - else: - ranges.append([cp, cp]) -version = ".".join(map(str, sys.version_info[:3])) -print( - f"//! Generated from Python {version} (Unicode {unicodedata.unidata_version}): code points >= 0x80 for which" -) -print( - "//! `str.isprintable()` is false, as inclusive ranges. Python's `repr` escapes these." -) -print("//! Regenerate with tests/make_printable.py.") -print() -print("pub const NON_PRINTABLE: &[(u32, u32)] = &[") -for lo, hi in ranges: - print(f" ({lo:#x}, {hi:#x}),") -print("];") From cfbccc4d2694f5dfdd7993a685482347e541e851 Mon Sep 17 00:00:00 2001 From: Tianyao Wu Date: Wed, 30 Sep 2026 17:12:01 +0800 Subject: [PATCH 09/10] cua_s1: trim the text worker to the reference path Keep the model, its HTTP worker and the tests that exercise them. The checked-in input set, revision detection, extra flags, bearer auth and duplicated validation go; the limits become constants. The README now describes the worker as the correctness reference for native execution. Signed-off-by: Tianyao Wu --- recipe/cua_s1/text.md | 54 ++---- src/frontend/cua_s1_text.py | 239 ++++++----------------- src/models/cua_s1/README.md | 4 +- src/models/cua_s1/text/contract.py | 270 +++++++------------------- src/models/cua_s1/text/model.py | 140 +++----------- tests/cua_s1/data/text_inputs.json | 289 ---------------------------- tests/cua_s1/test_text_contract.py | 271 ++++++++++++-------------- tests/cua_s1/test_text_model.py | 36 ---- tests/cua_s1/test_text_server.py | 180 ++++------------- tests/cua_s1/test_text_tokenizer.py | 45 ----- 10 files changed, 323 insertions(+), 1205 deletions(-) delete mode 100644 tests/cua_s1/data/text_inputs.json delete mode 100644 tests/cua_s1/test_text_model.py delete mode 100644 tests/cua_s1/test_text_tokenizer.py diff --git a/recipe/cua_s1/text.md b/recipe/cua_s1/text.md index 53be6947..1fcfcb1d 100644 --- a/recipe/cua_s1/text.md +++ b/recipe/cua_s1/text.md @@ -1,78 +1,46 @@ # Cua-S1 4B 0.2 text worker -This recipe runs the Cua-S1 4B 0.2 `text` adapter behind the Rust frontend. The model lives in [`src/models/cua_s1/text/`](../../src/models/cua_s1/text/) and the HTTP worker in [`src/frontend/cua_s1_text.py`](../../src/frontend/cua_s1_text.py); [`src/models/cua_s1/README.md`](../../src/models/cua_s1/README.md) documents the inference contract and the request mapping. Only `choice` questions are supported. +This recipe runs the Cua-S1 4B 0.2 `text` adapter through Transformers and PEFT behind the Rust frontend. It is the correctness reference for native execution. The model is in [`src/models/cua_s1/text/`](../../src/models/cua_s1/text/), the HTTP worker in [`src/frontend/cua_s1_text.py`](../../src/frontend/cua_s1_text.py), and [`src/models/cua_s1/README.md`](../../src/models/cua_s1/README.md) documents the contract. Only `choice` questions are supported. -Run all commands from the repository root, on Linux with an NVIDIA GPU. - -## Install - -Use Python 3.12. The pinned versions match the upstream reference environment: +Run the commands from the repository root, on Linux with an NVIDIA GPU and Python 3.12. The pinned versions match the upstream reference environment: ```sh python3.12 -m venv .venv .venv/bin/python -m pip install -r recipe/cua_s1/requirements-text.txt -``` - -## Download the weights - -Download the pinned revisions (about 9.5 GB) into `weights/`: - -```sh .venv/bin/hf download Qwen/Qwen3.5-4B \ --revision 851bf6e806efd8d0a36b00ddf55e13ccb7b8cd0a --local-dir weights/Qwen3.5-4B .venv/bin/hf download cua-ai/cua-s1-4b-0.2 \ --revision 16818868b0cc7813808aae4e87b417657046ab79 --local-dir weights/cua-s1-4b-0.2 ``` -To verify every file against upstream's lock, clone [trycua/cua](https://github.com/trycua/cua) next to this repository, check out `0e75660ce4c2edda519e0c795fa3ad98abf4e76f`, and run: - -```sh -.venv/bin/python ../cua/libs/cua-s1/ci/fetch_pinned_weights.py --dest weights --verify-only -``` +Upstream's `libs/cua-s1/ci/fetch_pinned_weights.py --dest weights --verify-only` (in [trycua/cua](https://github.com/trycua/cua) at `0e75660ce4c2edda519e0c795fa3ad98abf4e76f`) checks every downloaded file against upstream's lock. -## Start the worker +Start the worker, which runs one warmup decision before it listens, then the frontend: ```sh PYTHONPATH=src .venv/bin/python -m frontend.cua_s1_text \ - --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text \ - --device cuda --dtype bfloat16 --host 127.0.0.1 --port 8000 -``` - -The worker loads the model and runs one warmup decision before it starts listening, so `GET /health` answers only once requests can be served; it then returns `{"status": "ready", "modality": "text", ...}`. The log reports load and warmup times separately. The worker refuses the `multimodal/` adapter and reports the adapter revision that `hf download` recorded. Every flag can also be set through an environment variable: `CUA_S1_BASE`, `CUA_S1_ADAPTER`, `CUA_S1_ADAPTER_REVISION`, `CUA_S1_DEVICE`, `CUA_S1_DTYPE`, `CUA_S1_HOST`, `CUA_S1_PORT`, `CUA_S1_MAX_BODY_BYTES`, `CUA_S1_MAX_QUESTIONS` and `CUA_S1_MAX_PROMPT_TOKENS`. `--adapter-revision` only sets the revision reported in `model` when the download metadata is missing; a value that contradicts the metadata stops the worker. Set `CUA_S1_API_KEY` to require `Authorization: Bearer ` on `/v1/systemone`. - -Oversized requests get `413`: bodies over 4 MiB, more than 64 questions, or a question whose prompt is over 16,384 tokens (`--max-body-bytes`, `--max-questions`, `--max-prompt-tokens`). The worker computes logits for every prompt position, as upstream does, so memory grows with prompt length: serving the 15,446-token test input in bfloat16 peaked at about 21.3 GiB in use on the card. Requests run one at a time, and the frontend gives up after 60 seconds. - -## Start the frontend - -Build and start the frontend from the repository root, with stable Rust installed: - -```sh + --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text --port 8000 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 +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 +Requests run one at a time. Bodies over 4 MiB, more than 64 questions, or a prompt over 16,384 tokens get `413`. The worker computes logits for every position, as upstream does, so memory grows with prompt length: the 15,446-token test input peaked at about 21.3 GiB in bfloat16. ```sh -curl http://127.0.0.1:8080/health curl http://127.0.0.1:8080/v1/systemone \ -H 'Content-Type: application/json' \ -d '{"model":"cua-s1-4b-0.2","state":"Dialog: Delete 3 files permanently? Buttons: Delete, Cancel","questions":{"pick":{"type":"choice","instructions":"Keep the files.","criteria":{"delete":"Click Delete","cancel":"Click Cancel"}}}}' ``` -The answer has the Jev choice shape. On an RTX 6000 Ada in bfloat16, the response is: +On an RTX 6000 Ada in bfloat16, the response is: ```json {"model":"cua-ai/cua-s1-4b-0.2@16818868b0cc7813808aae4e87b417657046ab79:text","answers":{"pick":{"type":"choice","choice":"cancel","probabilities":{"delete":0.0024726232513785362,"cancel":0.9975274205207825},"confidence":0.9750249565060322}},"usage":{"input_tokens":153,"output_tokens":0}} ``` -## Tests - -The contract and HTTP tests need neither weights nor a GPU. The tokenizer tests also run when `CUA_S1_BASE` points to the downloaded base model; they read only its tokenizer files: +The tests need neither weights nor a GPU; with `CUA_S1_BASE=weights/Qwen3.5-4B` they also check the tokenizer: ```sh -.venv/bin/python -m pip install pytest -CUA_S1_BASE=weights/Qwen3.5-4B PYTHONPATH=src .venv/bin/python -m pytest tests/cua_s1/test_text_*.py +.venv/bin/python -m pip install pytest httpx +PYTHONPATH=src .venv/bin/python -m pytest tests/cua_s1 ``` diff --git a/src/frontend/cua_s1_text.py b/src/frontend/cua_s1_text.py index eb2295a1..ef3a330c 100644 --- a/src/frontend/cua_s1_text.py +++ b/src/frontend/cua_s1_text.py @@ -1,21 +1,13 @@ -"""HTTP worker for Cua-S1 4B 0.2 (`text` adapter) behind the Rust frontend. +"""HTTP worker for Cua-S1 4B 0.2 (`text` adapter): `GET /health` and `POST /v1/systemone`. -Routes: `GET /health` and `POST /v1/systemone`. The model is loaded before the -server starts listening, and one forward pass runs at a time. The model itself -is in `src/models/cua_s1/text/`. - - PYTHONPATH=src python -m frontend.cua_s1_text --base --adapter +PYTHONPATH=src python -m frontend.cua_s1_text --base --adapter /text """ from __future__ import annotations import argparse import asyncio -import hmac -import json -import os import sys -import time import traceback from concurrent.futures import ThreadPoolExecutor from typing import Any @@ -26,214 +18,91 @@ from fastapi.responses import JSONResponse from models.cua_s1.text.contract import ( - ADAPTER_REVISION, - MODEL_NAME, + MODEL_ID, RequestError, answer, map_request, - model_identity, parse_body, ) -WARMUP_REQUEST = { - "model": MODEL_NAME, - "state": "Dialog: 'Update installed.' Button: OK", - "questions": { - "warmup": { - "type": "choice", - "instructions": "Close the dialog.", - "criteria": {"ok": "Click OK", "wait": "Wait"}, - } - }, -} - - -def build_app( - model: Any, - *, - api_key: str | None, - max_body_bytes: int, - max_questions: int, - max_prompt_tokens: int, - revision: str, -): - app = FastAPI() - pool = ThreadPoolExecutor(max_workers=1) - identity = model_identity(revision) - expected_auth = ( - f"Bearer {api_key}".encode("utf-8", "surrogateescape") if api_key else b"" - ) - - def error(status: int, message: str) -> JSONResponse: - return JSONResponse({"detail": message}, status_code=status) - - def authorized(request: Request) -> bool: - if not api_key: - return True - supplied = request.headers.get("authorization", "").encode( - "utf-8", "surrogateescape" - ) - return hmac.compare_digest(supplied, expected_auth) +MAX_BODY_BYTES = 4 << 20 +MAX_PROMPT_TOKENS = 16384 +WARMUP = ( + b'{"model": "cua-s1-4b-0.2", "state": "Dialog: Update installed.", "questions": {"q":' + b' {"type": "choice", "instructions": "Close it.", "criteria": {"ok": "OK", "wait": "Wait"}}}}' +) - @app.get("/health") - def health(): - return { - "status": "ready", - "modality": "text", - "model": identity, - "device": model.device, - "dtype": model.dtype, - } - def decide(mapped): - # Tokenize every question first, so an over-long prompt is rejected - # before any forward pass runs. - encoded = [] - for question in mapped.questions: - inputs = model.encode(mapped.state, question) - n = int(inputs["input_ids"].shape[1]) - if max_prompt_tokens and n > max_prompt_tokens: +def build_app(model: Any) -> FastAPI: + app = FastAPI() + pool = ThreadPoolExecutor(max_workers=1) # one forward pass at a time + + def decide(raw: bytes) -> dict[str, Any]: + state, questions = map_request(parse_body(raw)) + encoded = [model.encode(state, q) for q in questions] + tokens = [int(x["input_ids"].shape[1]) for x in encoded] + for q, n in zip(questions, tokens): # before any forward pass + if n > MAX_PROMPT_TOKENS: raise RequestError( - f"question {question.name!r}: prompt is {n} tokens, " - f"over the {max_prompt_tokens}-token limit", - status=413, + f"question {q.name!r}: {n} prompt tokens, over {MAX_PROMPT_TOKENS}", + 413, ) - encoded.append((question, inputs)) - answers, prompt_tokens = {}, 0 - for question, inputs in encoded: - scored = model.score_encoded(inputs, len(question.keys)) - answers[question.name] = answer(question, scored.probabilities) - prompt_tokens += scored.prompt_tokens + answers = { + q.name: answer(q, model.score(x, len(q.keys))) + for q, x in zip(questions, encoded) + } return { - "model": identity, + "model": MODEL_ID, "answers": answers, - "usage": {"input_tokens": prompt_tokens, "output_tokens": 0}, + "usage": {"input_tokens": sum(tokens), "output_tokens": 0}, } + @app.get("/health") + def health(): + return {"status": "ready", "model": MODEL_ID} + @app.post("/v1/systemone") async def systemone(request: Request): - if not authorized(request): - return error(401, "invalid or missing bearer token") - length = request.headers.get("content-length") - if length and length.isdigit() and int(length) > max_body_bytes: - return error(413, "request body too large") raw = bytearray() async for chunk in request.stream(): - raw.extend(chunk) - if len(raw) > max_body_bytes: - return error(413, "request body too large") + raw += chunk + if len(raw) > MAX_BODY_BYTES: + return JSONResponse({"detail": "request body too large"}, 413) try: - mapped = map_request(parse_body(bytes(raw)), max_questions=max_questions) - loop = asyncio.get_running_loop() - # Returned as a JSONResponse: FastAPI's default encoder would drop - # every key that starts with "_sa", and question names and option - # keys come from the request. - return JSONResponse(await loop.run_in_executor(pool, decide, mapped)) - except RequestError as exc: - return error(exc.status, str(exc)) + # JSONResponse, not FastAPI's encoder, which drops keys starting with "_sa". + return JSONResponse( + await asyncio.get_running_loop().run_in_executor(pool, decide, raw) + ) + except RequestError as error: + return JSONResponse({"detail": str(error)}, error.status) except Exception: traceback.print_exc(file=sys.stderr) - return error(500, "inference failed") + return JSONResponse({"detail": "inference failed"}, 500) - def warmup() -> None: - """Run one decision on the worker thread through the full request path.""" - mapped = map_request(WARMUP_REQUEST) - json.dumps(pool.submit(decide, mapped).result(), allow_nan=False) - - app.state.warmup = warmup + # One decision on the worker thread before listening, so the first request does not + # pay for first-call setup there. + app.state.warmup = lambda: pool.submit(decide, WARMUP).result() return app -def main(argv: list[str] | None = None) -> None: - env = os.environ.get +def main() -> None: parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("--base", required=True, help="local Qwen/Qwen3.5-4B directory") parser.add_argument( - "--base", - default=env("CUA_S1_BASE"), - help="local Qwen/Qwen3.5-4B directory (env CUA_S1_BASE)", - ) - parser.add_argument( - "--adapter", - default=env("CUA_S1_ADAPTER"), - help="local cua-ai/cua-s1-4b-0.2 directory (env CUA_S1_ADAPTER)", - ) - parser.add_argument( - "--adapter-revision", - default=env("CUA_S1_ADAPTER_REVISION"), - help="adapter revision to report when the download metadata is missing", - ) - parser.add_argument("--device", default=env("CUA_S1_DEVICE", "cuda")) - parser.add_argument( - "--dtype", - default=env("CUA_S1_DTYPE", "bfloat16"), - choices=["bfloat16", "float16", "float32"], - ) - parser.add_argument("--host", default=env("CUA_S1_HOST", "127.0.0.1")) - parser.add_argument("--port", type=int, default=int(env("CUA_S1_PORT", "8000"))) - parser.add_argument( - "--max-body-bytes", - type=int, - default=int(env("CUA_S1_MAX_BODY_BYTES", str(4 << 20))), - ) - parser.add_argument( - "--max-questions", type=int, default=int(env("CUA_S1_MAX_QUESTIONS", "64")) - ) - parser.add_argument( - "--max-prompt-tokens", - type=int, - default=int(env("CUA_S1_MAX_PROMPT_TOKENS", "16384")), - help="per question; 0 disables the check", + "--adapter", required=True, help="local text/ adapter directory" ) - args = parser.parse_args(argv) - if not args.base or not args.adapter: - parser.error("--base and --adapter are required") + parser.add_argument("--device", default="cuda") + parser.add_argument("--dtype", default="bfloat16", choices=["bfloat16", "float32"]) + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--port", type=int, default=8000) + args = parser.parse_args() import uvicorn - from models.cua_s1.text.model import ( - TextModel, - downloaded_revision, - text_adapter_dir, - ) + from models.cua_s1.text.model import TextModel - # Fail before loading weights if this is not the text adapter. - text_adapter_dir(args.adapter) - detected = downloaded_revision(args.adapter) - if detected and args.adapter_revision and detected != args.adapter_revision: - parser.error( - f"--adapter-revision {args.adapter_revision} does not match the " - f"downloaded revision {detected}" - ) - revision = detected or args.adapter_revision or ADAPTER_REVISION - if revision != ADAPTER_REVISION: - print( - f"warning: adapter revision {revision} is not the pinned {ADAPTER_REVISION}", - flush=True, - ) - if not detected: - print( - "note: no download metadata under --adapter; the adapter revision is not verified", - flush=True, - ) - - model = TextModel(args.base, args.adapter, args.device, args.dtype) - print( - f"loaded in {model.load_seconds:.1f} s on {args.device} ({args.dtype})", - flush=True, - ) - app = build_app( - model, - api_key=env("CUA_S1_API_KEY") or None, - max_body_bytes=args.max_body_bytes, - max_questions=args.max_questions, - max_prompt_tokens=args.max_prompt_tokens, - revision=revision, - ) - # One decision before listening, so the first real request does not pay - # for lazy weight loading or first-call kernel setup on the worker thread. - started = time.perf_counter() + app = build_app(TextModel(args.base, args.adapter, args.device, args.dtype)) app.state.warmup() - print(f"warmed up in {time.perf_counter() - started:.1f} s", flush=True) uvicorn.run(app, host=args.host, port=args.port, log_level="warning") diff --git a/src/models/cua_s1/README.md b/src/models/cua_s1/README.md index 8711b17e..aa0e4be7 100644 --- a/src/models/cua_s1/README.md +++ b/src/models/cua_s1/README.md @@ -2,7 +2,7 @@ This directory owns Cua-S1 4B 0.2 ([#10](https://github.com/ThinkFlowLab/system1-omni/issues/10)): request mapping, prompt construction, adapter selection, execution, and the answer-letter readout. This page records the pinned upstream revisions, the inference contract an implementation must match, and how its outputs will be compared with the upstream reference. -Status: a reference worker for the `text` adapter loads the model directly through Hugging Face Transformers and PEFT and serves `/v1/systemone`. The model is in [`text/`](text/) (`contract.py` for request mapping, prompts and answers; `model.py` for loading and the answer-letter readout), the HTTP worker is [`src/frontend/cua_s1_text.py`](../../frontend/cua_s1_text.py), and setup is in [`recipe/cua_s1/text.md`](../../../recipe/cua_s1/text.md). A worker for the `multimodal` adapter is proposed in [#12](https://github.com/ThinkFlowLab/system1-omni/pull/12). +Status: a reference worker for the `text` adapter loads the model through Hugging Face Transformers and PEFT: [`text/`](text/), served by [`src/frontend/cua_s1_text.py`](../../frontend/cua_s1_text.py), with setup in [`recipe/cua_s1/text.md`](../../../recipe/cua_s1/text.md). It is the correctness reference for native execution, which starts with a minimal `text` prefill path on CUDA. The `multimodal` adapter is deferred; see [Not covered yet](#not-covered-yet). ## Pinned revisions @@ -88,7 +88,7 @@ The status is `422` when a well-formed request cannot be answered: ## Validation -**Inputs.** The fixed input set is upstream's two checked-in fixtures, converted to `/v1/systemone` requests with the chooser's rendered regions as `state`, plus `/v1/systemone` choice requests, in `tests/cua_s1/data/text_inputs.json`. These cover 1 to 26 options, short and long states, string and structured `state`, `instructions` and `criteria`, `null` criteria, non-ASCII text, and text that spells a special token. Each input is scored once per configuration. +**Inputs.** The fixed input set is upstream's two checked-in fixtures, converted to `/v1/systemone` requests with the chooser's rendered regions as `state`, plus `/v1/systemone` choice requests. These cover 1 to 26 options, short and long states, string and structured `state`, `instructions` and `criteria`, `null` criteria, non-ASCII text, and text that spells a special token. Each input is scored once per configuration. **Tolerances.** These are declared before any comparison is run: diff --git a/src/models/cua_s1/text/contract.py b/src/models/cua_s1/text/contract.py index 5fcfcb3e..326e758d 100644 --- a/src/models/cua_s1/text/contract.py +++ b/src/models/cua_s1/text/contract.py @@ -1,31 +1,22 @@ -"""Request mapping, prompt construction and answers for Cua-S1 4B 0.2. - -This module follows the contract in `src/models/cua_s1/README.md`. It has no -torch or Transformers imports, so it can be tested without weights. +"""Request mapping, prompts and answers for Cua-S1 4B 0.2, following +`src/models/cua_s1/README.md`. No torch imports, so it can be tested without weights. """ from __future__ import annotations import json import math -import string from dataclasses import dataclass from typing import Any MODEL_NAME = "cua-s1-4b-0.2" -ADAPTER_REPO = "cua-ai/cua-s1-4b-0.2" -ADAPTER_REVISION = "16818868b0cc7813808aae4e87b417657046ab79" -BASE_REPO = "Qwen/Qwen3.5-4B" - -LETTERS = string.ascii_uppercase -MAX_OPTIONS = len(LETTERS) +MODEL_ID = "cua-ai/cua-s1-4b-0.2@16818868b0cc7813808aae4e87b417657046ab79:text" +LETTERS = "ABCDEFGHIJKLMNOPQRSTUVWXYZ" +MAX_QUESTIONS = 64 -# The system message, the user message layout and the fixed values below are -# copied from trycua/cua at 0e75660ce4c2edda519e0c795fa3ad98abf4e76f: -# `libs/cua-s1/python/src/cua_s1/four_b.py` (SYSTEM_PROMPT, build_prompt, -# _describe_option) and `libs/cua-driver/examples/jev-use/python/ -# decision_models.py` (S1DecisionModel.score), as are the two upstream fixtures -# converted in `tests/cua_s1/data/text_inputs.json`. +# The system message and the user message layout are copied from trycua/cua at +# 0e75660ce4c2edda519e0c795fa3ad98abf4e76f (`libs/cua-s1/python/src/cua_s1/four_b.py` +# and `libs/cua-driver/examples/jev-use/python/decision_models.py`). # # MIT License # @@ -55,205 +46,100 @@ "one option: the single best next action to take. Answer with ONLY that " "option's letter -- no words, no punctuation, no explanation." ) -APP = "Cua Driver" -TASK_FAMILY = "closed-candidate decision" -ROLE = "Decision" -ACTION = "select" class RequestError(ValueError): - """A request the worker rejects. `status` is the HTTP status to return.""" - def __init__(self, message: str, status: int = 422) -> None: super().__init__(message) self.status = status -def _object_pairs(pairs: list[tuple[str, Any]]) -> dict[str, Any]: - obj: dict[str, Any] = {} +@dataclass(frozen=True) +class Question: + name: str + goal: str + keys: tuple[str, ...] + labels: tuple[str, ...] + + +def _unique_keys(pairs: list[tuple[str, Any]]) -> dict[str, Any]: + obj = {} for key, value in pairs: if key in obj: - raise RequestError(f"duplicate key {key!r} in a JSON object", status=400) + raise RequestError(f"duplicate key {key!r}", 400) obj[key] = value return obj -def _reject_constant(name: str) -> Any: - raise RequestError(f"{name} is not valid JSON", status=400) - - -def _finite_float(text: str) -> float: - value = float(text) - if not math.isfinite(value): - raise RequestError(f"number {text} is out of range", status=400) - return value - - -def _check_unicode(value: Any) -> None: - """Reject lone surrogates (for example a `\\ud800` escape): they cannot be - encoded as UTF-8, so they cannot be tokenized or echoed back.""" - if isinstance(value, str): - value.encode("utf-8") - elif isinstance(value, dict): - for key, item in value.items(): - key.encode("utf-8") - _check_unicode(item) - elif isinstance(value, list): - for item in value: - _check_unicode(item) - - def parse_body(raw: bytes) -> dict[str, Any]: - """Decode a request body, keeping key order and rejecting duplicate keys.""" try: - text = raw.decode("utf-8") - body = json.loads( - text, - object_pairs_hook=_object_pairs, - parse_constant=_reject_constant, - parse_float=_finite_float, - ) - _check_unicode(body) - except RequestError: - raise - except RecursionError as error: - raise RequestError("request body is nested too deeply", status=400) from error - except UnicodeError as error: - raise RequestError( - "request body must be valid UTF-8 text", status=400 - ) from error - except ValueError as error: - # json.JSONDecodeError, a UTF-8 byte order mark, or an integer too long - # for Python to convert. - raise RequestError("request body must be valid JSON", status=400) from error + body = json.loads(raw.decode(), object_pairs_hook=_unique_keys) + # NaN, Infinity, numbers out of range and lone surrogates fail here. + json.dumps(body, ensure_ascii=False, allow_nan=False).encode() + except (ValueError, RecursionError) as error: + if isinstance(error, RequestError): + raise + raise RequestError(f"request body is not valid JSON: {error}", 400) from error if not isinstance(body, dict): - raise RequestError("request body must be a JSON object", status=400) + raise RequestError("request body must be a JSON object", 400) return body -def as_text(value: Any) -> str: - """Render `state` or `instructions` as prompt text. - - A string is used as is; an object or array is serialized the way Python's - `json.dumps(value, ensure_ascii=False)` does. - """ - if isinstance(value, str): - return value - return json.dumps(value, ensure_ascii=False) - - -def escape_label(value: str) -> str: - """Escape an option label the way upstream's chooser does.""" - return json.dumps(value, ensure_ascii=False)[1:-1] - - -@dataclass(frozen=True) -class Question: - """One `choice` question mapped onto the prompt fields.""" - - name: str - goal: str - keys: tuple[str, ...] - labels: tuple[str, ...] - - -@dataclass(frozen=True) -class Request: - state: str - questions: tuple[Question, ...] - - -def _check_json_value(value: Any, where: str, allow_null: bool) -> None: - if value is None: - if not allow_null: - raise RequestError(f"{where} must not be null") - return - if isinstance(value, bool) or isinstance(value, (int, float)): - raise RequestError(f"{where} must be a string, an object or an array") +def _text(value: Any, where: str) -> str: + """A string as is; an object or array as Python's json.dumps writes it.""" if not isinstance(value, (str, dict, list)): raise RequestError(f"{where} must be a string, an object or an array") + return value if isinstance(value, str) else json.dumps(value, ensure_ascii=False) -def map_request(body: dict[str, Any], *, max_questions: int = 64) -> Request: - """Validate a `/v1/systemone` body and map it onto prompt fields.""" - model = body.get("model") - if model != MODEL_NAME: +def map_request(body: dict[str, Any]) -> tuple[str, list[Question]]: + if body.get("model") != MODEL_NAME: raise RequestError(f"'model' must be {MODEL_NAME!r}") - - if "state" not in body: - raise RequestError("'state' is required") - state_value = body["state"] - _check_json_value(state_value, "'state'", allow_null=False) - if state_value in ("", {}, []): + if body.get("state") in ("", {}, []): raise RequestError("'state' must not be empty") - state = as_text(state_value) - + state = _text(body.get("state"), "'state'") questions = body.get("questions") if not isinstance(questions, dict) or not questions: raise RequestError("'questions' must be a non-empty object") - if len(questions) > max_questions: - raise RequestError( - f"too many questions ({len(questions)} > {max_questions})", status=413 - ) - - # Check every question type before the per-question checks, so a `score` - # or `noul` question anywhere rejects the whole request with that reason. - for name, question in questions.items(): - if not isinstance(question, dict): - raise RequestError(f"question {name!r} must be an object") - kind = question.get("type") - if kind in ("score", "noul"): - raise RequestError( - f"question {name!r}: type {kind!r} is not supported; " - "Cua-S1 4B 0.2 answers 'choice' questions only" - ) - if kind != "choice": - raise RequestError(f"question {name!r}: unknown type {kind!r}") - + if len(questions) > MAX_QUESTIONS: + raise RequestError(f"more than {MAX_QUESTIONS} questions", 413) mapped = [] - for name, question in questions.items(): + for name, q in questions.items(): where = f"question {name!r}" - if "instructions" not in question: + if not isinstance(q, dict): + raise RequestError(f"{where} must be an object") + if q.get("type") in ("score", "noul"): + raise RequestError(f"{where}: type {q['type']!r} is not supported") + if q.get("type") != "choice": + raise RequestError(f"{where}: unknown type {q.get('type')!r}") + if "instructions" not in q: raise RequestError(f"{where}: 'instructions' is required") - instructions = question["instructions"] - _check_json_value(instructions, f"{where}: 'instructions'", allow_null=True) - goal = "" if instructions is None else as_text(instructions) - - criteria = question.get("criteria") - if not isinstance(criteria, dict): - raise RequestError(f"{where}: 'criteria' must be an object") - if not criteria: - raise RequestError(f"{where}: 'criteria' must have at least one option") - if len(criteria) > MAX_OPTIONS: + goal = "" if q["instructions"] is None else _text(q["instructions"], where) + criteria = q.get("criteria") + if not isinstance(criteria, dict) or not 1 <= len(criteria) <= len(LETTERS): raise RequestError( - f"{where}: {len(criteria)} options; at most {MAX_OPTIONS} are supported" + f"{where}: 'criteria' must be an object with 1 to 26 options" ) - keys, labels = [], [] - for key, value in criteria.items(): - _check_json_value(value, f"{where}: option {key!r}", allow_null=True) - if value is None: - text = key - else: - text = as_text(value) - keys.append(key) - labels.append(escape_label(text)) - mapped.append( - Question(name=name, goal=goal, keys=tuple(keys), labels=tuple(labels)) + labels = tuple( + json.dumps( + key if value is None else _text(value, f"{where}: {key!r}"), + ensure_ascii=False, + )[1:-1] + for key, value in criteria.items() ) - return Request(state=state, questions=tuple(mapped)) + mapped.append(Question(name, goal, tuple(criteria), labels)) + return state, mapped def build_messages(state: str, question: Question) -> list[dict[str, str]]: - """Chat messages for one question, matching upstream `build_prompt` (text).""" - option_lines = "\n".join( - f'{letter}. {ROLE} "{label}" -> {ACTION}' - for letter, label in zip(LETTERS, question.labels, strict=False) + options = "\n".join( + f'{letter}. Decision "{label}" -> select' + for letter, label in zip(LETTERS, question.labels) ) user = ( (f"Goal: {question.goal}\n\n" if question.goal else "") - + f"App: {APP}\nTask family: {TASK_FAMILY}\n\n" - + f"Accessibility tree:\n{state}\n\n" - + f"Options:\n{option_lines}\n\nAnswer with a single letter." + + "App: Cua Driver\nTask family: closed-candidate decision\n\n" + + f"Accessibility tree:\n{state}\n\nOptions:\n{options}\n\nAnswer with a single letter." ) return [ {"role": "system", "content": SYSTEM_PROMPT}, @@ -261,34 +147,14 @@ def build_messages(state: str, question: Question) -> list[dict[str, str]]: ] -def confidence(probabilities: list[float]) -> float: - """Normalized entropy, `1 - H(p) / ln(n)`, as the LAYA worker reports it.""" - n = len(probabilities) - if n < 2: - return 1.0 - entropy = -sum(p * math.log(min(max(p, 1e-12), 1.0)) for p in probabilities) - return min(max(1.0 - entropy / math.log(n), 0.0), 1.0) - - def answer(question: Question, probabilities: list[float]) -> dict[str, Any]: - """The Jev choice answer. Ties go to the earliest option.""" - if len(probabilities) != len(question.keys) or not all( - math.isfinite(p) and 0.0 <= p <= 1.0 for p in probabilities - ): - raise ValueError(f"model returned invalid probabilities: {probabilities}") - if not math.isclose(sum(probabilities), 1.0, abs_tol=1e-5): - raise ValueError(f"model probabilities do not sum to one: {probabilities}") - best = 0 - for index, p in enumerate(probabilities): - if p > probabilities[best]: - best = index + """The Jev choice answer; ties go to the earliest option. `confidence` is the + normalized entropy `1 - H(p) / ln(n)`, as the LAYA worker reports it.""" + n = len(probabilities) + entropy = -sum(p * math.log(p) for p in probabilities if p > 0) return { "type": "choice", - "choice": question.keys[best], - "probabilities": dict(zip(question.keys, probabilities, strict=True)), - "confidence": confidence(probabilities), + "choice": question.keys[max(range(n), key=probabilities.__getitem__)], + "probabilities": dict(zip(question.keys, probabilities)), + "confidence": max(0.0, 1 - entropy / math.log(n)) if n > 1 else 1.0, } - - -def model_identity(revision: str = ADAPTER_REVISION) -> str: - return f"{ADAPTER_REPO}@{revision}:text" diff --git a/src/models/cua_s1/text/model.py b/src/models/cua_s1/text/model.py index 9724c426..ea0415af 100644 --- a/src/models/cua_s1/text/model.py +++ b/src/models/cua_s1/text/model.py @@ -1,129 +1,45 @@ -"""Load Qwen3.5-4B with the Cua-S1 `text` adapter and score one prompt. - -The calls mirror upstream `cua_s1.four_b.FourBModel` (text modality): the same -model class, an unmerged PEFT adapter, the chat template with its default -generation prompt, full logits, and a fp32 softmax over the letter logits at -the last position. Keeping them the same is what makes the worker's -probabilities bitwise identical to the reference in the same environment. - -Torch, Transformers and PEFT are imported only when a model is loaded, so the -adapter checks can run without them. +"""Qwen3.5-4B with the Cua-S1 `text` adapter, loaded and scored as upstream +`cua_s1.four_b.FourBModel` does: an unmerged PEFT adapter, the chat template with its +generation prompt, full logits, and a float32 softmax over the option letters at the +last position. That keeps the probabilities bitwise identical to the reference. """ from __future__ import annotations import json -import re -import time -from dataclasses import dataclass from pathlib import Path -from .contract import BASE_REPO, LETTERS, Question, build_messages - - -def text_adapter_dir(adapter_root: str | Path) -> Path: - """Return the `text` adapter directory under the adapter root. - - Accepts the repository root (`/text`) or the `text/` directory - itself, and refuses the `multimodal/` adapter: PEFT only warns about keys - it cannot place, so loading the wrong adapter would otherwise go unnoticed. - """ - root = Path(adapter_root) - path = root / "text" if (root / "text" / "adapter_config.json").exists() else root - config_file = path / "adapter_config.json" - if not config_file.exists(): - raise RuntimeError(f"no adapter_config.json under {root}") - config = json.loads(config_file.read_text()) - if config.get("base_model_name_or_path") != BASE_REPO: - raise RuntimeError(f"{config_file}: base model is not {BASE_REPO}") - if {"linear_fc1", "linear_fc2"} & set(config.get("target_modules") or []): - raise RuntimeError( - f"{config_file}: this is the multimodal adapter, not the text adapter" - ) - return path - - -def downloaded_revision(adapter_root: str | Path) -> str | None: - """The commit that `hf download --local-dir` recorded for the text adapter, if any. +import torch +from peft import PeftModel +from transformers import AutoModelForCausalLM, AutoTokenizer - `hf download` keeps its metadata under the repository root, so this also - looks one level up when `adapter_root` is the `text/` directory itself. - """ - root = Path(adapter_root) - places = [(root, "text/"), (root, "")] - if root.name == "text": - places.insert(0, (root.parent, "text/")) - for base, prefix in places: - cache = base / ".cache" / "huggingface" / "download" - try: - meta = (cache / f"{prefix}adapter_model.safetensors.metadata").read_text() - first = meta.splitlines()[0].strip() - except (OSError, IndexError): - continue - if re.fullmatch(r"[0-9a-f]{40}", first): - return first - return None - - -@dataclass -class Scored: - probabilities: list[float] - prompt_tokens: int +from .contract import LETTERS, Question, build_messages class TextModel: - def __init__( - self, base_model: str, adapter_root: str, device: str, dtype: str - ) -> None: - import torch - from peft import PeftModel - from transformers import AutoModelForCausalLM, AutoTokenizer - - self.device = device - self.dtype = dtype - started = time.perf_counter() - self.tokenizer = AutoTokenizer.from_pretrained(base_model) + def __init__(self, base: str, adapter: str, device: str, dtype: str) -> None: + # PEFT only warns about keys it cannot place: refuse the multimodal adapter. + config = json.loads((Path(adapter) / "adapter_config.json").read_text()) + if "linear_fc1" in config["target_modules"]: + raise ValueError( + f"{adapter} is the multimodal adapter; pass its text/ directory" + ) + self.tokenizer = AutoTokenizer.from_pretrained(base) model = AutoModelForCausalLM.from_pretrained( - base_model, dtype=getattr(torch, dtype), device_map=device + base, dtype=getattr(torch, dtype), device_map=device ) - model = PeftModel.from_pretrained(model, str(text_adapter_dir(adapter_root))) - model.eval() - self.model = model - self.load_seconds = time.perf_counter() - started - self.letter_ids = self._letter_ids() - - def _letter_ids(self) -> list[int]: - ids = [] - for letter in LETTERS: - tokens = self.tokenizer.encode(letter, add_special_tokens=False) - if len(tokens) != 1: - raise RuntimeError(f"letter {letter!r} is not a single token: {tokens}") - ids.append(tokens[0]) - return ids + self.model = PeftModel.from_pretrained(model, adapter).eval() + self.letter_ids = self.tokenizer.convert_tokens_to_ids(list(LETTERS)) def encode(self, state: str, question: Question): - """Tokenized prompt for one question, on CPU. - - The Qwen3.5 tokenizer adds no special tokens here (contract point 4); - the chat template already contains them. - """ - chat_text = self.tokenizer.apply_chat_template( + text = self.tokenizer.apply_chat_template( build_messages(state, question), tokenize=False, add_generation_prompt=True ) - return self.tokenizer(chat_text, return_tensors="pt") - - def score_encoded(self, inputs, n_options: int) -> Scored: - import torch - - with torch.no_grad(): - inputs = inputs.to(self.model.device) - out = self.model(**inputs) - final_logits = out.logits[0, -1, :] - letter_ids = self.letter_ids[:n_options] - option_logits = final_logits[ - torch.tensor(letter_ids, device=final_logits.device) - ] - probabilities = torch.softmax(option_logits.float(), dim=-1).tolist() - return Scored( - probabilities=probabilities, prompt_tokens=int(inputs["input_ids"].shape[1]) - ) + return self.tokenizer(text, return_tensors="pt") + + @torch.no_grad() + def score(self, inputs, n_options: int) -> list[float]: + logits = self.model(**inputs.to(self.model.device)).logits[0, -1] + return torch.softmax( + logits[self.letter_ids[:n_options]].float(), dim=-1 + ).tolist() diff --git a/tests/cua_s1/data/text_inputs.json b/tests/cua_s1/data/text_inputs.json deleted file mode 100644 index 1a56460b..00000000 --- a/tests/cua_s1/data/text_inputs.json +++ /dev/null @@ -1,289 +0,0 @@ -{ - "fixture_positive": { - "model": "cua-s1-4b-0.2", - "state": "Visual-region-derived observation for capture \"capture-fixture-1\":\n\"submit-text\": text 'Submit' at (300,240,100,40) confidence=0.96 interactive=true", - "questions": { - "pick": { - "type": "choice", - "instructions": "Submit the verified form.", - "criteria": { - "submit-form": "Submit using the unique validated visual region.", - "reobserve": "Discard this decision set and obtain a fresh observation.", - "abstain": "Stop without acting if no supplied action is safe." - } - } - } - }, - "fixture_negative": { - "model": "cua-s1-4b-0.2", - "state": "Visual-region-derived observation for capture \"synthetic-negative-1\":\n\"save\": text 'Save' at (10,10,80,30) confidence=0.98 interactive=true", - "questions": { - "pick": { - "type": "choice", - "instructions": "Choose exactly one candidate by applying its condition to the supplied state. The host alone authorizes any selected action.", - "criteria": { - "region:save": "Select only if exactly one supplied region has id=save, kind=text, exact_text=\"Send\", confidence=0.98, and it is the sole exact Send match at or above 0.80.", - "reobserve": "Select only when no action candidate condition matches and no regions are supplied. Do not act; request one fresh bounded observation.", - "abstain": "Select only when no action candidate condition matches and one or more regions are supplied. Do not act; stop." - } - } - } - }, - "one_option": { - "model": "cua-s1-4b-0.2", - "state": "Dialog: 'Update installed.' Button: OK", - "questions": { - "pick": { - "type": "choice", - "instructions": "Close the dialog.", - "criteria": { - "ok": "Click OK" - } - } - } - }, - "two_options": { - "model": "cua-s1-4b-0.2", - "state": "Dialog: 'Delete 3 files permanently?' Buttons: Delete, Cancel", - "questions": { - "pick": { - "type": "choice", - "instructions": "Keep the files.", - "criteria": { - "delete": "Click Delete", - "cancel": "Click Cancel" - } - } - } - }, - "max_26_options": { - "model": "cua-s1-4b-0.2", - "state": "Toolbar of a document editor. Selected text: 'quarterly results'.\nbutton 'Undo' enabled=true\nbutton 'Redo' enabled=true\nbutton 'Cut' enabled=true\nbutton 'Copy' enabled=true\nbutton 'Paste' enabled=true\nbutton 'Bold' enabled=true\nbutton 'Italic' enabled=true\nbutton 'Underline' enabled=true\nbutton 'Strikethrough' enabled=true\nbutton 'Font color' enabled=true\nbutton 'Highlight' enabled=true\nbutton 'Align left' enabled=true\nbutton 'Center' enabled=true\nbutton 'Align right' enabled=true\nbutton 'Justify' enabled=true\nbutton 'Bullets' enabled=true\nbutton 'Numbering' enabled=true\nbutton 'Indent' enabled=true\nbutton 'Outdent' enabled=true\nbutton 'Insert link' enabled=true\nbutton 'Insert image' enabled=true\nbutton 'Insert table' enabled=true\nbutton 'Comment' enabled=true\nbutton 'Find' enabled=true\nbutton 'Replace' enabled=true\nbutton 'Print' enabled=true", - "questions": { - "pick": { - "type": "choice", - "instructions": "Make the selected text bold.", - "criteria": { - "undo": "Click the 'Undo' button", - "redo": "Click the 'Redo' button", - "cut": "Click the 'Cut' button", - "copy": "Click the 'Copy' button", - "paste": "Click the 'Paste' button", - "bold": "Click the 'Bold' button", - "italic": "Click the 'Italic' button", - "underline": "Click the 'Underline' button", - "strikethrough": "Click the 'Strikethrough' button", - "font-color": "Click the 'Font color' button", - "highlight": "Click the 'Highlight' button", - "align-left": "Click the 'Align left' button", - "center": "Click the 'Center' button", - "align-right": "Click the 'Align right' button", - "justify": "Click the 'Justify' button", - "bullets": "Click the 'Bullets' button", - "numbering": "Click the 'Numbering' button", - "indent": "Click the 'Indent' button", - "outdent": "Click the 'Outdent' button", - "insert-link": "Click the 'Insert link' button", - "insert-image": "Click the 'Insert image' button", - "insert-table": "Click the 'Insert table' button", - "comment": "Click the 'Comment' button", - "find": "Click the 'Find' button", - "replace": "Click the 'Replace' button", - "print": "Click the 'Print' button" - } - } - } - }, - "long_state": { - "model": "cua-s1-4b-0.2", - "state": "Orders table (web admin), 300 rows, sorted by order number.\nrow 0: cell 'Order #10000' | cell 'Customer 0' | cell '10.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 1: cell 'Order #10001' | cell 'Customer 1' | cell '17.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 2: cell 'Order #10002' | cell 'Customer 2' | cell '24.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 3: cell 'Order #10003' | cell 'Customer 3' | cell '31.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 4: cell 'Order #10004' | cell 'Customer 4' | cell '38.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 5: cell 'Order #10005' | cell 'Customer 5' | cell '45.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 6: cell 'Order #10006' | cell 'Customer 6' | cell '52.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 7: cell 'Order #10007' | cell 'Customer 7' | cell '59.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 8: cell 'Order #10008' | cell 'Customer 8' | cell '66.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 9: cell 'Order #10009' | cell 'Customer 9' | cell '73.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 10: cell 'Order #10010' | cell 'Customer 10' | cell '80.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 11: cell 'Order #10011' | cell 'Customer 11' | cell '87.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 12: cell 'Order #10012' | cell 'Customer 12' | cell '94.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 13: cell 'Order #10013' | cell 'Customer 13' | cell '11.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 14: cell 'Order #10014' | cell 'Customer 14' | cell '18.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 15: cell 'Order #10015' | cell 'Customer 15' | cell '25.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 16: cell 'Order #10016' | cell 'Customer 16' | cell '32.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 17: cell 'Order #10017' | cell 'Customer 17' | cell '39.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 18: cell 'Order #10018' | cell 'Customer 18' | cell '46.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 19: cell 'Order #10019' | cell 'Customer 19' | cell '53.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 20: cell 'Order #10020' | cell 'Customer 20' | cell '60.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 21: cell 'Order #10021' | cell 'Customer 21' | cell '67.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 22: cell 'Order #10022' | cell 'Customer 22' | cell '74.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 23: cell 'Order #10023' | cell 'Customer 23' | cell '81.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 24: cell 'Order #10024' | cell 'Customer 24' | cell '88.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 25: cell 'Order #10025' | cell 'Customer 25' | cell '95.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 26: cell 'Order #10026' | cell 'Customer 26' | cell '12.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 27: cell 'Order #10027' | cell 'Customer 27' | cell '19.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 28: cell 'Order #10028' | cell 'Customer 28' | cell '26.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 29: cell 'Order #10029' | cell 'Customer 29' | cell '33.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 30: cell 'Order #10030' | cell 'Customer 30' | cell '40.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 31: cell 'Order #10031' | cell 'Customer 31' | cell '47.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 32: cell 'Order #10032' | cell 'Customer 32' | cell '54.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 33: cell 'Order #10033' | cell 'Customer 33' | cell '61.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 34: cell 'Order #10034' | cell 'Customer 34' | cell '68.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 35: cell 'Order #10035' | cell 'Customer 35' | cell '75.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 36: cell 'Order #10036' | cell 'Customer 36' | cell '82.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 37: cell 'Order #10037' | cell 'Customer 0' | cell '89.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 38: cell 'Order #10038' | cell 'Customer 1' | cell '96.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 39: cell 'Order #10039' | cell 'Customer 2' | cell '13.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 40: cell 'Order #10040' | cell 'Customer 3' | cell '20.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 41: cell 'Order #10041' | cell 'Customer 4' | cell '27.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 42: cell 'Order #10042' | cell 'Customer 5' | cell '34.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 43: cell 'Order #10043' | cell 'Customer 6' | cell '41.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 44: cell 'Order #10044' | cell 'Customer 7' | cell '48.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 45: cell 'Order #10045' | cell 'Customer 8' | cell '55.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 46: cell 'Order #10046' | cell 'Customer 9' | cell '62.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 47: cell 'Order #10047' | cell 'Customer 10' | cell '69.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 48: cell 'Order #10048' | cell 'Customer 11' | cell '76.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 49: cell 'Order #10049' | cell 'Customer 12' | cell '83.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 50: cell 'Order #10050' | cell 'Customer 13' | cell '90.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 51: cell 'Order #10051' | cell 'Customer 14' | cell '97.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 52: cell 'Order #10052' | cell 'Customer 15' | cell '14.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 53: cell 'Order #10053' | cell 'Customer 16' | cell '21.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 54: cell 'Order #10054' | cell 'Customer 17' | cell '28.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 55: cell 'Order #10055' | cell 'Customer 18' | cell '35.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 56: cell 'Order #10056' | cell 'Customer 19' | cell '42.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 57: cell 'Order #10057' | cell 'Customer 20' | cell '49.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 58: cell 'Order #10058' | cell 'Customer 21' | cell '56.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 59: cell 'Order #10059' | cell 'Customer 22' | cell '63.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 60: cell 'Order #10060' | cell 'Customer 23' | cell '70.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 61: cell 'Order #10061' | cell 'Customer 24' | cell '77.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 62: cell 'Order #10062' | cell 'Customer 25' | cell '84.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 63: cell 'Order #10063' | cell 'Customer 26' | cell '91.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 64: cell 'Order #10064' | cell 'Customer 27' | cell '98.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 65: cell 'Order #10065' | cell 'Customer 28' | cell '15.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 66: cell 'Order #10066' | cell 'Customer 29' | cell '22.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 67: cell 'Order #10067' | cell 'Customer 30' | cell '29.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 68: cell 'Order #10068' | cell 'Customer 31' | cell '36.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 69: cell 'Order #10069' | cell 'Customer 32' | cell '43.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 70: cell 'Order #10070' | cell 'Customer 33' | cell '50.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 71: cell 'Order #10071' | cell 'Customer 34' | cell '57.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 72: cell 'Order #10072' | cell 'Customer 35' | cell '64.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 73: cell 'Order #10073' | cell 'Customer 36' | cell '71.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 74: cell 'Order #10074' | cell 'Customer 0' | cell '78.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 75: cell 'Order #10075' | cell 'Customer 1' | cell '85.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 76: cell 'Order #10076' | cell 'Customer 2' | cell '92.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 77: cell 'Order #10077' | cell 'Customer 3' | cell '99.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 78: cell 'Order #10078' | cell 'Customer 4' | cell '16.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 79: cell 'Order #10079' | cell 'Customer 5' | cell '23.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 80: cell 'Order #10080' | cell 'Customer 6' | cell '30.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 81: cell 'Order #10081' | cell 'Customer 7' | cell '37.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 82: cell 'Order #10082' | cell 'Customer 8' | cell '44.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 83: cell 'Order #10083' | cell 'Customer 9' | cell '51.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 84: cell 'Order #10084' | cell 'Customer 10' | cell '58.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 85: cell 'Order #10085' | cell 'Customer 11' | cell '65.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 86: cell 'Order #10086' | cell 'Customer 12' | cell '72.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 87: cell 'Order #10087' | cell 'Customer 13' | cell '79.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 88: cell 'Order #10088' | cell 'Customer 14' | cell '86.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 89: cell 'Order #10089' | cell 'Customer 15' | cell '93.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 90: cell 'Order #10090' | cell 'Customer 16' | cell '10.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 91: cell 'Order #10091' | cell 'Customer 17' | cell '17.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 92: cell 'Order #10092' | cell 'Customer 18' | cell '24.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 93: cell 'Order #10093' | cell 'Customer 19' | cell '31.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 94: cell 'Order #10094' | cell 'Customer 20' | cell '38.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 95: cell 'Order #10095' | cell 'Customer 21' | cell '45.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 96: cell 'Order #10096' | cell 'Customer 22' | cell '52.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 97: cell 'Order #10097' | cell 'Customer 23' | cell '59.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 98: cell 'Order #10098' | cell 'Customer 24' | cell '66.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 99: cell 'Order #10099' | cell 'Customer 25' | cell '73.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 100: cell 'Order #10100' | cell 'Customer 26' | cell '80.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 101: cell 'Order #10101' | cell 'Customer 27' | cell '87.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 102: cell 'Order #10102' | cell 'Customer 28' | cell '94.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 103: cell 'Order #10103' | cell 'Customer 29' | cell '11.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 104: cell 'Order #10104' | cell 'Customer 30' | cell '18.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 105: cell 'Order #10105' | cell 'Customer 31' | cell '25.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 106: cell 'Order #10106' | cell 'Customer 32' | cell '32.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 107: cell 'Order #10107' | cell 'Customer 33' | cell '39.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 108: cell 'Order #10108' | cell 'Customer 34' | cell '46.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 109: cell 'Order #10109' | cell 'Customer 35' | cell '53.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 110: cell 'Order #10110' | cell 'Customer 36' | cell '60.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 111: cell 'Order #10111' | cell 'Customer 0' | cell '67.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 112: cell 'Order #10112' | cell 'Customer 1' | cell '74.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 113: cell 'Order #10113' | cell 'Customer 2' | cell '81.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 114: cell 'Order #10114' | cell 'Customer 3' | cell '88.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 115: cell 'Order #10115' | cell 'Customer 4' | cell '95.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 116: cell 'Order #10116' | cell 'Customer 5' | cell '12.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 117: cell 'Order #10117' | cell 'Customer 6' | cell '19.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 118: cell 'Order #10118' | cell 'Customer 7' | cell '26.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 119: cell 'Order #10119' | cell 'Customer 8' | cell '33.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 120: cell 'Order #10120' | cell 'Customer 9' | cell '40.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 121: cell 'Order #10121' | cell 'Customer 10' | cell '47.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 122: cell 'Order #10122' | cell 'Customer 11' | cell '54.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 123: cell 'Order #10123' | cell 'Customer 12' | cell '61.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 124: cell 'Order #10124' | cell 'Customer 13' | cell '68.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 125: cell 'Order #10125' | cell 'Customer 14' | cell '75.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 126: cell 'Order #10126' | cell 'Customer 15' | cell '82.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 127: cell 'Order #10127' | cell 'Customer 16' | cell '89.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 128: cell 'Order #10128' | cell 'Customer 17' | cell '96.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 129: cell 'Order #10129' | cell 'Customer 18' | cell '13.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 130: cell 'Order #10130' | cell 'Customer 19' | cell '20.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 131: cell 'Order #10131' | cell 'Customer 20' | cell '27.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 132: cell 'Order #10132' | cell 'Customer 21' | cell '34.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 133: cell 'Order #10133' | cell 'Customer 22' | cell '41.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 134: cell 'Order #10134' | cell 'Customer 23' | cell '48.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 135: cell 'Order #10135' | cell 'Customer 24' | cell '55.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 136: cell 'Order #10136' | cell 'Customer 25' | cell '62.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 137: cell 'Order #10137' | cell 'Customer 26' | cell '69.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 138: cell 'Order #10138' | cell 'Customer 27' | cell '76.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 139: cell 'Order #10139' | cell 'Customer 28' | cell '83.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 140: cell 'Order #10140' | cell 'Customer 29' | cell '90.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 141: cell 'Order #10141' | cell 'Customer 30' | cell '97.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 142: cell 'Order #10142' | cell 'Customer 31' | cell '14.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 143: cell 'Order #10143' | cell 'Customer 32' | cell '21.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 144: cell 'Order #10144' | cell 'Customer 33' | cell '28.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 145: cell 'Order #10145' | cell 'Customer 34' | cell '35.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 146: cell 'Order #10146' | cell 'Customer 35' | cell '42.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 147: cell 'Order #10147' | cell 'Customer 36' | cell '49.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 148: cell 'Order #10148' | cell 'Customer 0' | cell '56.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 149: cell 'Order #10149' | cell 'Customer 1' | cell '63.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 150: cell 'Order #10150' | cell 'Customer 2' | cell '70.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 151: cell 'Order #10151' | cell 'Customer 3' | cell '77.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 152: cell 'Order #10152' | cell 'Customer 4' | cell '84.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 153: cell 'Order #10153' | cell 'Customer 5' | cell '91.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 154: cell 'Order #10154' | cell 'Customer 6' | cell '98.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 155: cell 'Order #10155' | cell 'Customer 7' | cell '15.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 156: cell 'Order #10156' | cell 'Customer 8' | cell '22.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 157: cell 'Order #10157' | cell 'Customer 9' | cell '29.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 158: cell 'Order #10158' | cell 'Customer 10' | cell '36.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 159: cell 'Order #10159' | cell 'Customer 11' | cell '43.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 160: cell 'Order #10160' | cell 'Customer 12' | cell '50.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 161: cell 'Order #10161' | cell 'Customer 13' | cell '57.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 162: cell 'Order #10162' | cell 'Customer 14' | cell '64.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 163: cell 'Order #10163' | cell 'Customer 15' | cell '71.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 164: cell 'Order #10164' | cell 'Customer 16' | cell '78.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 165: cell 'Order #10165' | cell 'Customer 17' | cell '85.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 166: cell 'Order #10166' | cell 'Customer 18' | cell '92.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 167: cell 'Order #10167' | cell 'Customer 19' | cell '99.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 168: cell 'Order #10168' | cell 'Customer 20' | cell '16.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 169: cell 'Order #10169' | cell 'Customer 21' | cell '23.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 170: cell 'Order #10170' | cell 'Customer 22' | cell '30.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 171: cell 'Order #10171' | cell 'Customer 23' | cell '37.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 172: cell 'Order #10172' | cell 'Customer 24' | cell '44.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 173: cell 'Order #10173' | cell 'Customer 25' | cell '51.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 174: cell 'Order #10174' | cell 'Customer 26' | cell '58.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 175: cell 'Order #10175' | cell 'Customer 27' | cell '65.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 176: cell 'Order #10176' | cell 'Customer 28' | cell '72.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 177: cell 'Order #10177' | cell 'Customer 29' | cell '79.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 178: cell 'Order #10178' | cell 'Customer 30' | cell '86.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 179: cell 'Order #10179' | cell 'Customer 31' | cell '93.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 180: cell 'Order #10180' | cell 'Customer 32' | cell '10.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 181: cell 'Order #10181' | cell 'Customer 33' | cell '17.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 182: cell 'Order #10182' | cell 'Customer 34' | cell '24.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 183: cell 'Order #10183' | cell 'Customer 35' | cell '31.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 184: cell 'Order #10184' | cell 'Customer 36' | cell '38.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 185: cell 'Order #10185' | cell 'Customer 0' | cell '45.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 186: cell 'Order #10186' | cell 'Customer 1' | cell '52.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 187: cell 'Order #10187' | cell 'Customer 2' | cell '59.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 188: cell 'Order #10188' | cell 'Customer 3' | cell '66.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 189: cell 'Order #10189' | cell 'Customer 4' | cell '73.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 190: cell 'Order #10190' | cell 'Customer 5' | cell '80.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 191: cell 'Order #10191' | cell 'Customer 6' | cell '87.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 192: cell 'Order #10192' | cell 'Customer 7' | cell '94.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 193: cell 'Order #10193' | cell 'Customer 8' | cell '11.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 194: cell 'Order #10194' | cell 'Customer 9' | cell '18.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 195: cell 'Order #10195' | cell 'Customer 10' | cell '25.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 196: cell 'Order #10196' | cell 'Customer 11' | cell '32.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 197: cell 'Order #10197' | cell 'Customer 12' | cell '39.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 198: cell 'Order #10198' | cell 'Customer 13' | cell '46.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 199: cell 'Order #10199' | cell 'Customer 14' | cell '53.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 200: cell 'Order #10200' | cell 'Customer 15' | cell '60.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 201: cell 'Order #10201' | cell 'Customer 16' | cell '67.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 202: cell 'Order #10202' | cell 'Customer 17' | cell '74.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 203: cell 'Order #10203' | cell 'Customer 18' | cell '81.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 204: cell 'Order #10204' | cell 'Customer 19' | cell '88.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 205: cell 'Order #10205' | cell 'Customer 20' | cell '95.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 206: cell 'Order #10206' | cell 'Customer 21' | cell '12.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 207: cell 'Order #10207' | cell 'Customer 22' | cell '19.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 208: cell 'Order #10208' | cell 'Customer 23' | cell '26.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 209: cell 'Order #10209' | cell 'Customer 24' | cell '33.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 210: cell 'Order #10210' | cell 'Customer 25' | cell '40.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 211: cell 'Order #10211' | cell 'Customer 26' | cell '47.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 212: cell 'Order #10212' | cell 'Customer 27' | cell '54.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 213: cell 'Order #10213' | cell 'Customer 28' | cell '61.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 214: cell 'Order #10214' | cell 'Customer 29' | cell '68.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 215: cell 'Order #10215' | cell 'Customer 30' | cell '75.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 216: cell 'Order #10216' | cell 'Customer 31' | cell '82.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 217 note: customer reported a duplicate charge on this order\nrow 217: cell 'Order #10217' | cell 'Customer 32' | cell '89.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 218: cell 'Order #10218' | cell 'Customer 33' | cell '96.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 219: cell 'Order #10219' | cell 'Customer 34' | cell '13.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 220: cell 'Order #10220' | cell 'Customer 35' | cell '20.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 221: cell 'Order #10221' | cell 'Customer 36' | cell '27.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 222: cell 'Order #10222' | cell 'Customer 0' | cell '34.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 223: cell 'Order #10223' | cell 'Customer 1' | cell '41.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 224: cell 'Order #10224' | cell 'Customer 2' | cell '48.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 225: cell 'Order #10225' | cell 'Customer 3' | cell '55.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 226: cell 'Order #10226' | cell 'Customer 4' | cell '62.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 227: cell 'Order #10227' | cell 'Customer 5' | cell '69.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 228: cell 'Order #10228' | cell 'Customer 6' | cell '76.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 229: cell 'Order #10229' | cell 'Customer 7' | cell '83.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 230: cell 'Order #10230' | cell 'Customer 8' | cell '90.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 231: cell 'Order #10231' | cell 'Customer 9' | cell '97.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 232: cell 'Order #10232' | cell 'Customer 10' | cell '14.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 233: cell 'Order #10233' | cell 'Customer 11' | cell '21.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 234: cell 'Order #10234' | cell 'Customer 12' | cell '28.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 235: cell 'Order #10235' | cell 'Customer 13' | cell '35.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 236: cell 'Order #10236' | cell 'Customer 14' | cell '42.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 237: cell 'Order #10237' | cell 'Customer 15' | cell '49.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 238: cell 'Order #10238' | cell 'Customer 16' | cell '56.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 239: cell 'Order #10239' | cell 'Customer 17' | cell '63.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 240: cell 'Order #10240' | cell 'Customer 18' | cell '70.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 241: cell 'Order #10241' | cell 'Customer 19' | cell '77.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 242: cell 'Order #10242' | cell 'Customer 20' | cell '84.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 243: cell 'Order #10243' | cell 'Customer 21' | cell '91.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 244: cell 'Order #10244' | cell 'Customer 22' | cell '98.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 245: cell 'Order #10245' | cell 'Customer 23' | cell '15.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 246: cell 'Order #10246' | cell 'Customer 24' | cell '22.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 247: cell 'Order #10247' | cell 'Customer 25' | cell '29.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 248: cell 'Order #10248' | cell 'Customer 26' | cell '36.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 249: cell 'Order #10249' | cell 'Customer 27' | cell '43.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 250: cell 'Order #10250' | cell 'Customer 28' | cell '50.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 251: cell 'Order #10251' | cell 'Customer 29' | cell '57.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 252: cell 'Order #10252' | cell 'Customer 30' | cell '64.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 253: cell 'Order #10253' | cell 'Customer 31' | cell '71.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 254: cell 'Order #10254' | cell 'Customer 32' | cell '78.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 255: cell 'Order #10255' | cell 'Customer 33' | cell '85.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 256: cell 'Order #10256' | cell 'Customer 34' | cell '92.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 257: cell 'Order #10257' | cell 'Customer 35' | cell '99.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 258: cell 'Order #10258' | cell 'Customer 36' | cell '16.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 259: cell 'Order #10259' | cell 'Customer 0' | cell '23.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 260: cell 'Order #10260' | cell 'Customer 1' | cell '30.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 261: cell 'Order #10261' | cell 'Customer 2' | cell '37.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 262: cell 'Order #10262' | cell 'Customer 3' | cell '44.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 263: cell 'Order #10263' | cell 'Customer 4' | cell '51.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 264: cell 'Order #10264' | cell 'Customer 5' | cell '58.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 265: cell 'Order #10265' | cell 'Customer 6' | cell '65.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 266: cell 'Order #10266' | cell 'Customer 7' | cell '72.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 267: cell 'Order #10267' | cell 'Customer 8' | cell '79.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 268: cell 'Order #10268' | cell 'Customer 9' | cell '86.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 269: cell 'Order #10269' | cell 'Customer 10' | cell '93.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 270: cell 'Order #10270' | cell 'Customer 11' | cell '10.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 271: cell 'Order #10271' | cell 'Customer 12' | cell '17.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 272: cell 'Order #10272' | cell 'Customer 13' | cell '24.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 273: cell 'Order #10273' | cell 'Customer 14' | cell '31.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 274: cell 'Order #10274' | cell 'Customer 15' | cell '38.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 275: cell 'Order #10275' | cell 'Customer 16' | cell '45.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 276: cell 'Order #10276' | cell 'Customer 17' | cell '52.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 277: cell 'Order #10277' | cell 'Customer 18' | cell '59.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 278: cell 'Order #10278' | cell 'Customer 19' | cell '66.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 279: cell 'Order #10279' | cell 'Customer 20' | cell '73.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 280: cell 'Order #10280' | cell 'Customer 21' | cell '80.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 281: cell 'Order #10281' | cell 'Customer 22' | cell '87.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 282: cell 'Order #10282' | cell 'Customer 23' | cell '94.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 283: cell 'Order #10283' | cell 'Customer 24' | cell '11.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 284: cell 'Order #10284' | cell 'Customer 25' | cell '18.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 285: cell 'Order #10285' | cell 'Customer 26' | cell '25.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 286: cell 'Order #10286' | cell 'Customer 27' | cell '32.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 287: cell 'Order #10287' | cell 'Customer 28' | cell '39.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 288: cell 'Order #10288' | cell 'Customer 29' | cell '46.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 289: cell 'Order #10289' | cell 'Customer 30' | cell '53.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 290: cell 'Order #10290' | cell 'Customer 31' | cell '60.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 291: cell 'Order #10291' | cell 'Customer 32' | cell '67.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 292: cell 'Order #10292' | cell 'Customer 33' | cell '74.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 293: cell 'Order #10293' | cell 'Customer 34' | cell '81.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 294: cell 'Order #10294' | cell 'Customer 35' | cell '88.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 295: cell 'Order #10295' | cell 'Customer 36' | cell '95.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 296: cell 'Order #10296' | cell 'Customer 0' | cell '12.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'\nrow 297: cell 'Order #10297' | cell 'Customer 1' | cell '19.00 USD' | cell 'Paid' | button 'Open' | button 'Refund'\nrow 298: cell 'Order #10298' | cell 'Customer 2' | cell '26.00 USD' | cell 'Pending' | button 'Open' | button 'Refund'\nrow 299: cell 'Order #10299' | cell 'Customer 3' | cell '33.00 USD' | cell 'Refunded' | button 'Open' | button 'Refund'", - "questions": { - "pick": { - "type": "choice", - "instructions": "Refund the order with the reported duplicate charge.", - "criteria": { - "refund-217": "Click 'Refund' in row 217", - "refund-216": "Click 'Refund' in row 216", - "open-217": "Click 'Open' in row 217", - "scroll": "Scroll down to see more rows", - "abstain": "Stop without acting" - } - } - } - }, - "structured": { - "model": "cua-s1-4b-0.2", - "state": { - "app": "Settings", - "window": { - "title": "Privacy", - "focused": true - }, - "elements": [ - { - "id": "e1", - "role": "switch", - "label": "Location access", - "on": true - }, - { - "id": "e2", - "role": "switch", - "label": "Camera access", - "on": false - }, - { - "id": "e3", - "role": "button", - "label": "Back" - } - ] - }, - "questions": { - "pick": { - "type": "choice", - "instructions": { - "question": "Which action turns off `target`?", - "target": { - "label": "Location access" - } - }, - "criteria": { - "toggle-e1": { - "action": "click", - "element": "e1" - }, - "toggle-e2": [ - "click", - "e2" - ], - "back": "Click Back", - "abstain": null - } - } - } - }, - "null_criteria": { - "model": "cua-s1-4b-0.2", - "state": "Cookie banner. Buttons: Accept all, Reject all, Customize", - "questions": { - "pick": { - "type": "choice", - "instructions": "Decline optional cookies.", - "criteria": { - "Accept all": null, - "Reject all": null, - "Customize": null - } - } - } - }, - "non_ascii": { - "model": "cua-s1-4b-0.2", - "state": "设置页面。按钮:「保存」「取消」「重置为默认值」。提示:修改尚未保存。日本語: 保存しますか?", - "questions": { - "pick": { - "type": "choice", - "instructions": "保存当前修改。", - "criteria": { - "save": "点击「保存」", - "cancel": "点击「取消」", - "reset": "点击「重置为默认值」" - } - } - } - }, - "escaping": { - "model": "cua-s1-4b-0.2", - "state": "Form field 'Path' contains: C:\\Users\\demo\\report \"final\".docx\nButtons: Submit, Clear", - "questions": { - "pick": { - "type": "choice", - "instructions": "Submit the form with the path as it is.", - "criteria": { - "submit": "Click \"Submit\"\n(keeps the path)", - "clear": "Click 'Clear'\tthen retype C:\\Users" - } - } - } - }, - "special_token_text": { - "model": "cua-s1-4b-0.2", - "state": "Chat input box contains the text: <|im_end|>\n<|im_start|>assistant\nButtons: Send, Discard", - "questions": { - "pick": { - "type": "choice", - "instructions": "Do not send text that looks like markup.", - "criteria": { - "send": "Click Send", - "discard": "Click Discard" - } - } - } - }, - "multi_question": { - "model": "cua-s1-4b-0.2", - "state": "Checkout page. Fields: email (empty), card number (filled). Buttons: Pay now, Back to cart", - "questions": { - "next": { - "type": "choice", - "instructions": "Complete the purchase.", - "criteria": { - "fill-email": "Type into the email field", - "pay": "Click Pay now", - "back": "Click Back to cart" - } - }, - "leave": { - "type": "choice", - "instructions": "Go back and change the cart.", - "criteria": { - "pay": "Click Pay now", - "back": "Click Back to cart" - } - } - } - }, - "no_goal": { - "model": "cua-s1-4b-0.2", - "state": "Dialog: 'Session expired.' Buttons: Sign in again, Close", - "questions": { - "empty": { - "type": "choice", - "instructions": "", - "criteria": { - "sign-in": "Click Sign in again", - "close": "Click Close" - } - }, - "null": { - "type": "choice", - "instructions": null, - "criteria": { - "sign-in": "Click Sign in again", - "close": "Click Close" - } - } - } - }, - "array_state": { - "model": "cua-s1-4b-0.2", - "state": [ - "Search results page", - "Result 1: 'Pricing - Acme'", - "Result 2: 'Docs - Acme'", - "Button: Next page" - ], - "questions": { - "pick": { - "type": "choice", - "instructions": "Open the documentation.", - "criteria": { - "r1": "Click result 1", - "r2": "Click result 2", - "next": "Click Next page" - } - } - } - } -} diff --git a/tests/cua_s1/test_text_contract.py b/tests/cua_s1/test_text_contract.py index 1d6cd2b0..870c50f2 100644 --- a/tests/cua_s1/test_text_contract.py +++ b/tests/cua_s1/test_text_contract.py @@ -1,11 +1,12 @@ -"""Contract tests that need neither weights nor torch. +"""Contract tests without weights or torch. The tokenizer test also runs when +CUA_S1_BASE points to a local Qwen/Qwen3.5-4B directory (tokenizer files only). PYTHONPATH=src python -m pytest tests/cua_s1 """ import json import math -from pathlib import Path +import os import pytest @@ -13,19 +14,46 @@ RequestError, answer, build_messages, - confidence, map_request, parse_body, ) -INPUTS = json.loads( - (Path(__file__).parent / "data" / "text_inputs.json").read_text(encoding="utf-8") -) -# The user message upstream's chooser builds for -# libs/cua-driver/examples/jev-use/fixtures/jev-choice-request-v1.json at the -# pinned revision (FourBModel text modality). -FIXTURE_POSITIVE_USER = ( +def body(state="Screen", **question): + q = { + "type": "choice", + "instructions": "Pick one.", + "criteria": {"a": "A", "b": "B"}, + } + q.update(question) + return {"model": "cua-s1-4b-0.2", "state": state, "questions": {"q": q}} + + +def mapped(request): + return map_request(parse_body(json.dumps(request).encode())) + + +def reject(request, status=422): + raw = request if isinstance(request, bytes) else json.dumps(request).encode() + with pytest.raises(RequestError) as info: + map_request(parse_body(raw)) + assert info.value.status == status + return str(info.value) + + +# upstream's libs/cua-driver/examples/jev-use/fixtures/jev-choice-request-v1.json, +# with the chooser's rendered region as `state`, and the user message upstream builds for it +FIXTURE = body( + 'Visual-region-derived observation for capture "capture-fixture-1":\n' + "\"submit-text\": text 'Submit' at (300,240,100,40) confidence=0.96 interactive=true", + instructions="Submit the verified form.", + criteria={ + "submit-form": "Submit using the unique validated visual region.", + "reobserve": "Discard this decision set and obtain a fresh observation.", + "abstain": "Stop without acting if no supplied action is safe.", + }, +) +FIXTURE_USER = ( "Goal: Submit the verified form.\n\n" "App: Cua Driver\nTask family: closed-candidate decision\n\n" "Accessibility tree:\n" @@ -39,178 +67,117 @@ ) -def mapped(name): - return map_request(parse_body(json.dumps(INPUTS[name]).encode())) - - -def reject(body, status=422): - raw = body if isinstance(body, bytes) else json.dumps(body).encode() - with pytest.raises(RequestError) as info: - map_request(parse_body(raw)) - assert info.value.status == status - return str(info.value) - - -def base(**question): - q = { - "type": "choice", - "instructions": "Pick one.", - "criteria": {"a": "A", "b": "B"}, - } - q.update(question) - return {"model": "cua-s1-4b-0.2", "state": "Screen", "questions": {"q": q}} - - def test_fixture_prompt_matches_upstream(): - request = mapped("fixture_positive") - messages = build_messages(request.state, request.questions[0]) - assert messages[1] == {"role": "user", "content": FIXTURE_POSITIVE_USER} - assert messages[0]["role"] == "system" - assert messages[0]["content"].startswith( + state, (question,) = mapped(FIXTURE) + system, user = build_messages(state, question) + assert user == {"role": "user", "content": FIXTURE_USER} + assert system["content"].startswith( "You are a one-pass computer-use decision model." ) -def test_every_input_maps(): - for name in INPUTS: - request = mapped(name) - assert request.questions - for question in request.questions: - assert 1 <= len(question.keys) <= 26 - - -def test_goal_line_left_out_when_empty_or_null(): - request = mapped("no_goal") - for question in request.questions: - user = build_messages(request.state, question)[1]["content"] - assert user.startswith("App: Cua Driver\n") - +@pytest.mark.skipif(not os.environ.get("CUA_S1_BASE"), reason="set CUA_S1_BASE to run") +def test_tokenizer(): + from transformers import AutoTokenizer -def test_structured_values_and_null_label(): - request = mapped("structured") - state = INPUTS["structured"]["state"] - assert request.state == json.dumps(state, ensure_ascii=False) - question = request.questions[0] - assert question.goal.startswith('{"question": "Which action turns off `target`?"') - assert question.labels[0] == '{\\"action\\": \\"click\\", \\"element\\": \\"e1\\"}' - assert question.labels[1] == '[\\"click\\", \\"e2\\"]' - assert question.labels[3] == "abstain" + tokenizer = AutoTokenizer.from_pretrained(os.environ["CUA_S1_BASE"]) + assert tokenizer.convert_tokens_to_ids(list("ABCDEFGHIJKLMNOPQRSTUVWXYZ")) == list( + range(32, 58) + ) + state, (question,) = mapped(FIXTURE) + text = tokenizer.apply_chat_template( + build_messages(state, question), tokenize=False, add_generation_prompt=True + ) + assert text.endswith("<|im_start|>assistant\n\n") + ids = tokenizer(text)["input_ids"] + assert len(ids) == 218 + assert ids == tokenizer(text, add_special_tokens=False)["input_ids"] -def test_label_escaping_matches_chooser(): - question = mapped("escaping").questions[0] - assert question.labels[0] == 'Click \\"Submit\\"\\n(keeps the path)' - assert question.labels[1] == "Click 'Clear'\\tthen retype C:\\\\Users" - assert mapped("non_ascii").questions[0].labels[0] == "点击「保存」" +def test_goal_left_out_when_empty_or_null(): + for goal in ["", None]: + state, (question,) = mapped(body(instructions=goal)) + assert build_messages(state, question)[1]["content"].startswith( + "App: Cua Driver\n" + ) -def test_score_or_noul_rejects_the_whole_request(): - body = base() - body["questions"]["s"] = { - "type": "score", - "instructions": "Rate it.", - "criteria": ["low", "high"], +def test_structured_values_escaping_and_null_label(): + tree = { + "app": "Settings", + "elements": [{"id": "e1", "label": "Location", "on": True}], } - assert "'score' is not supported" in reject(body) - body = base() - body["questions"]["n"] = {"type": "noul", "instructions": "Is it red?"} - assert "'noul' is not supported" in reject(body) - - -def test_option_count_limits(): - assert "at least one option" in reject(base(criteria={})) - many = {f"o{i}": f"Option {i}" for i in range(27)} - assert "27 options" in reject(base(criteria=many)) - assert ( - len( - map_request(base(criteria={f"o{i}": "x" for i in range(26)})) - .questions[0] - .keys + state, (question,) = mapped( + body( + tree, + instructions={"question": "Which one?"}, + criteria={ + "e1": {"action": "click", "element": "e1"}, + "e2": ["click", "e2"], + "quote": 'Click "Submit"\n(tab\there) C:\\Users', + "save": "点击「保存」", + "abstain": None, + }, ) - == 26 + ) + assert state == json.dumps(tree, ensure_ascii=False) + assert question.goal == '{"question": "Which one?"}' + assert question.labels == ( + '{\\"action\\": \\"click\\", \\"element\\": \\"e1\\"}', + '[\\"click\\", \\"e2\\"]', + 'Click \\"Submit\\"\\n(tab\\there) C:\\\\Users', + "点击「保存」", + "abstain", ) -def test_duplicate_keys_anywhere(): - raw = ( - b'{"model": "cua-s1-4b-0.2", "state": "S", "questions": {"q": {"type": "choice",' - b' "instructions": "I", "criteria": {"a": "A", "a": "B"}}}}' - ) - assert "duplicate key 'a'" in reject(raw, status=400) - raw = ( - b'{"model": "cua-s1-4b-0.2", "state": {"x": 1, "x": 2}, "questions": {"q": {"type":' - b' "choice", "instructions": "I", "criteria": {"a": "A"}}}}' +def test_request_errors(): + assert "not supported" in reject(body(type="score", criteria=["low", "high"])) + assert "not supported" in reject(body(type="noul")) + assert "unknown type" in reject(body(type="rank")) + assert "1 to 26 options" in reject(body(criteria={})) + assert "1 to 26 options" in reject(body(criteria={f"o{i}": "x" for i in range(27)})) + assert ( + len(mapped(body(criteria={f"o{i}": "x" for i in range(26)}))[1][0].keys) == 26 ) - assert "duplicate key 'x'" in reject(raw, status=400) + no_instructions = body() + del no_instructions["questions"]["q"]["instructions"] + assert "'instructions' is required" in reject(no_instructions) + for value in [1, 2.5, True]: + reject(body(criteria={"a": value})) + for state in ["", {}, [], None, 3, True]: + reject(body(state)) + assert "'model'" in reject({**body(), "model": "english"}) + many = body() + many["questions"] = {f"q{i}": many["questions"]["q"] for i in range(65)} + reject(many, status=413) @pytest.mark.parametrize( "raw", [ + b'{"model": "cua-s1-4b-0.2", "state": {"x": 1, "x": 2}}', b'{"model": "cua-s1-4b-0.2", "state": NaN}', b'{"model": "cua-s1-4b-0.2", "state": {"x": 1e400}}', b'{"model": "cua-s1-4b-0.2", "state": {"x": ' + b"9" * 5000 + b"}}", b'{"model": "cua-s1-4b-0.2", "state": "\\ud800"}', b"[" * 100000 + b"]" * 100000, b"\xff\xfe", - '{"model": "cua-s1-4b-0.2", "state": "S"}'.encode("utf-16"), - b"\xef\xbb\xbf" + b'{"model": "cua-s1-4b-0.2", "state": "S"}', + b"\xef\xbb\xbf{}", + b"not json", + b"[1, 2]", ], ) def test_malformed_bodies_are_400(raw): reject(raw, status=400) -def test_question_shape_errors(): - body = base() - body["questions"]["q"] = "not an object" - assert "must be an object" in reject(body) - assert "'criteria' must be an object" in reject(base(criteria=["a", "b"])) - body = base() - del body["questions"]["q"]["instructions"] - assert "'instructions' is required" in reject(body) - assert "unknown type 'rank'" in reject(base(type="rank")) - body = base() - body["questions"] = {f"q{i}": body["questions"]["q"] for i in range(3)} - with pytest.raises(RequestError) as info: - map_request(body, max_questions=2) - assert info.value.status == 413 - - -@pytest.mark.parametrize("value", [1, 2.5, True, False]) -def test_number_or_boolean_criteria_value(value): - assert "must be a string, an object or an array" in reject( - base(criteria={"a": value, "b": "B"}) - ) - - -@pytest.mark.parametrize("state", ["", {}, [], None, 3, True]) -def test_bad_state(state): - body = base() - body["state"] = state - reject(body) - - -def test_model_name_and_body_shape(): - body = base() - body["model"] = "english" - assert "'model' must be" in reject(body) - reject(b"not json", status=400) - reject(b"[1, 2]", status=400) - - -def test_confidence_is_normalized_entropy(): - assert confidence([1.0]) == 1.0 - assert confidence([0.5, 0.5]) == pytest.approx(0.0, abs=1e-12) - p = [0.88, 0.12, 0.0] - h = -(0.88 * math.log(0.88) + 0.12 * math.log(0.12)) - assert confidence(p) == pytest.approx(1 - h / math.log(3)) - - -def test_answer_shape_and_ties(): - question = mapped("two_options").questions[0] - result = answer(question, [0.5, 0.5]) - assert result["choice"] == "delete" - assert result["type"] == "choice" - assert list(result["probabilities"]) == ["delete", "cancel"] - assert answer(question, [0.2, 0.8])["choice"] == "cancel" +def test_answer(): + _, (question,) = mapped(body()) + tie = answer(question, [0.5, 0.5]) + assert tie["choice"] == "a" and tie["confidence"] == pytest.approx(0.0, abs=1e-12) + result = answer(question, [0.12, 0.88]) + assert result["choice"] == "b" + assert list(result["probabilities"]) == ["a", "b"] + h = -(0.12 * math.log(0.12) + 0.88 * math.log(0.88)) + assert result["confidence"] == pytest.approx(1 - h / math.log(2)) diff --git a/tests/cua_s1/test_text_model.py b/tests/cua_s1/test_text_model.py deleted file mode 100644 index 8dd0fd83..00000000 --- a/tests/cua_s1/test_text_model.py +++ /dev/null @@ -1,36 +0,0 @@ -"""Adapter directory checks: no weights, no torch.""" - -import json - -import pytest - -from models.cua_s1.text.model import downloaded_revision, text_adapter_dir - -REV = "16818868b0cc7813808aae4e87b417657046ab79" - - -def write_adapter(path, targets): - path.mkdir(parents=True) - config = {"base_model_name_or_path": "Qwen/Qwen3.5-4B", "target_modules": targets} - (path / "adapter_config.json").write_text(json.dumps(config)) - - -def test_text_adapter_dir(tmp_path): - write_adapter(tmp_path / "text", ["q_proj", "down_proj"]) - write_adapter(tmp_path / "multimodal", ["q_proj", "linear_fc1"]) - assert text_adapter_dir(tmp_path) == tmp_path / "text" - assert text_adapter_dir(tmp_path / "text") == tmp_path / "text" - with pytest.raises(RuntimeError, match="multimodal adapter"): - text_adapter_dir(tmp_path / "multimodal") - - -def test_downloaded_revision_from_repository_root(tmp_path): - write_adapter(tmp_path / "text", ["q_proj"]) - meta = ( - tmp_path / ".cache/huggingface/download/text/adapter_model.safetensors.metadata" - ) - meta.parent.mkdir(parents=True) - meta.write_text(f"{REV}\nabc\n1\n") - assert downloaded_revision(tmp_path) == REV - assert downloaded_revision(tmp_path / "text") == REV - assert downloaded_revision(tmp_path / "missing") is None diff --git a/tests/cua_s1/test_text_server.py b/tests/cua_s1/test_text_server.py index ed5cb7fd..bc4be9c9 100644 --- a/tests/cua_s1/test_text_server.py +++ b/tests/cua_s1/test_text_server.py @@ -1,12 +1,9 @@ """HTTP tests for the worker with a fake model: no weights, no torch.""" import json -from dataclasses import dataclass import pytest -# The worker's own requirements include fastapi and httpx; skip where only the -# contract tests' dependencies are installed. pytest.importorskip("fastapi") pytest.importorskip("httpx") from fastapi.testclient import TestClient # noqa: E402 @@ -14,170 +11,75 @@ from frontend.cua_s1_text import build_app # noqa: E402 -@dataclass -class _Ids: - shape: tuple +class Ids: + def __init__(self, n): + self.shape = (1, n) class FakeModel: - device = "cpu" - dtype = "float32" - - def __init__(self, tokens=100, fail=False, nan=False): - self.tokens = tokens - self.fail = fail - self.nan = nan - self.forward_calls = 0 + def __init__(self, tokens=100, error=None, nan=False): + self.tokens, self.error, self.nan, self.calls = tokens, error, nan, 0 def encode(self, state, question): - return {"input_ids": _Ids(shape=(1, self.tokens))} - - def score_encoded(self, inputs, n_options): - self.forward_calls += 1 - if self.fail: - raise RuntimeError("CUDA out of memory") - - @dataclass - class Scored: - probabilities: list - prompt_tokens: int + return {"input_ids": Ids(self.tokens)} - probabilities = [0.1] * n_options - probabilities[-1] = 1.0 - 0.1 * (n_options - 1) - if self.nan: - probabilities[0] = float("nan") - return Scored(probabilities, inputs["input_ids"].shape[1]) - - -def client(model=None, api_key=None, max_body_bytes=4 << 20, max_prompt_tokens=32768): - app = build_app( - model or FakeModel(), - api_key=api_key, - max_body_bytes=max_body_bytes, - max_questions=64, - max_prompt_tokens=max_prompt_tokens, - revision="r", - ) - return TestClient(app) + def score(self, inputs, n_options): + self.calls += 1 + if self.error: + raise self.error + p = [0.1] * (n_options - 1) + [1 - 0.1 * (n_options - 1)] + return [float("nan")] + p[1:] if self.nan else p BODY = { "model": "cua-s1-4b-0.2", "state": "Screen", "questions": { - "q": { + "_sa": { "type": "choice", "instructions": "Pick.", - "criteria": {"a": "A", "b": "B"}, + "criteria": {"_x": "A", "b": "B"}, } }, } -def test_health(): - response = client().get("/health") - assert response.status_code == 200 - assert response.json()["model"] == "cua-ai/cua-s1-4b-0.2@r:text" - assert response.json()["status"] == "ready" - assert response.json()["modality"] == "text" - - -def test_choice_answer(): - response = client().post("/v1/systemone", json=BODY) - assert response.status_code == 200, response.text - body = response.json() - assert body["answers"]["q"]["type"] == "choice" - assert body["answers"]["q"]["choice"] == "b" - assert body["usage"] == {"input_tokens": 100, "output_tokens": 0} +def post(model=None, **kwargs): + return TestClient(build_app(model or FakeModel())).post("/v1/systemone", **kwargs) -def test_keys_come_back_as_sent(): - body = json.loads(json.dumps(BODY)) - body["questions"] = { - "_sample": { - "type": "choice", - "instructions": "Pick.", - "criteria": {"_save": "Save", "b": "B"}, - } - } - response = client().post("/v1/systemone", json=body) - assert response.status_code == 200, response.text - answers = response.json()["answers"] - assert list(answers) == ["_sample"] - assert list(answers["_sample"]["probabilities"]) == ["_save", "b"] - - -def test_chunked_upload(): - raw = json.dumps(BODY).encode() - response = client().post( - "/v1/systemone", - content=iter([raw[:10], raw[10:]]), - headers={"content-type": "application/json"}, - ) +def test_health_and_answer(): + app = build_app(FakeModel()) + assert TestClient(app).get("/health").json()["status"] == "ready" + response = post(json=BODY) assert response.status_code == 200, response.text + reply = response.json() + assert reply["answers"]["_sa"]["choice"] == "b" + assert list(reply["answers"]["_sa"]["probabilities"]) == ["_x", "b"] + assert reply["usage"] == {"input_tokens": 100, "output_tokens": 0} def test_errors(): - c = client() bad = json.loads(json.dumps(BODY)) - bad["questions"]["q"]["type"] = "noul" - response = c.post("/v1/systemone", json=bad) - assert response.status_code == 422 - assert "'noul' is not supported" in response.json()["detail"] - assert c.post("/v1/systemone", content=b"{").status_code == 400 - - -def test_limits(): - assert client(max_body_bytes=50).post("/v1/systemone", json=BODY).status_code == 413 - raw = json.dumps(BODY).encode() - streamed = client(max_body_bytes=50).post( - "/v1/systemone", - content=iter([raw[:40], raw[40:]]), - headers={"content-type": "application/json"}, - ) - assert streamed.status_code == 413 - model = FakeModel(tokens=40000) - body = json.loads(json.dumps(BODY)) - body["questions"]["r"] = body["questions"]["q"] - response = client(model).post("/v1/systemone", json=body) - assert response.status_code == 413 - assert "token limit" in response.json()["detail"] - assert model.forward_calls == 0 - - -@pytest.mark.parametrize("model", [FakeModel(fail=True), FakeModel(nan=True)]) -def test_model_failure_is_json_500(model): - response = client(model).post("/v1/systemone", json=BODY) + bad["questions"]["_sa"]["type"] = "noul" + assert post(json=bad).status_code == 422 + assert post(content=b"{").status_code == 400 + raw = b" " * (4 << 20) + json.dumps(BODY).encode() + assert post(content=iter([raw[:10], raw[10:]])).status_code == 413 + model = FakeModel(tokens=20000) + assert post(model, json=BODY).status_code == 413 and model.calls == 0 + + +@pytest.mark.parametrize( + "model", [FakeModel(error=RuntimeError("CUDA out of memory")), FakeModel(nan=True)] +) +def test_model_failure_is_500(model): + response = post(model, json=BODY) assert response.status_code == 500 assert response.json() == {"detail": "inference failed"} -def test_warmup_runs_the_request_path(): +def test_warmup(): model = FakeModel() - app = build_app( - model, - api_key=None, - max_body_bytes=1 << 20, - max_questions=64, - max_prompt_tokens=32768, - revision="r", - ) - app.state.warmup() - assert model.forward_calls == 1 - with pytest.raises(ValueError): - build_app( - FakeModel(nan=True), - api_key=None, - max_body_bytes=1 << 20, - max_questions=64, - max_prompt_tokens=32768, - revision="r", - ).state.warmup() - - -def test_bearer_token(): - c = client(api_key="secret") - assert c.post("/v1/systemone", json=BODY).status_code == 401 - assert c.get("/health").status_code == 200 - ok = c.post("/v1/systemone", json=BODY, headers={"Authorization": "Bearer secret"}) - assert ok.status_code == 200 + build_app(model).state.warmup() + assert model.calls == 1 diff --git a/tests/cua_s1/test_text_tokenizer.py b/tests/cua_s1/test_text_tokenizer.py deleted file mode 100644 index 90bcd941..00000000 --- a/tests/cua_s1/test_text_tokenizer.py +++ /dev/null @@ -1,45 +0,0 @@ -"""Tokenizer checks against the pinned base model (tokenizer files only, no weights). - -Set CUA_S1_BASE to a local Qwen/Qwen3.5-4B directory to run them. -""" - -import json -import os -from pathlib import Path - -import pytest - -from models.cua_s1.text.contract import LETTERS, build_messages, map_request, parse_body - -BASE = os.environ.get("CUA_S1_BASE") -pytestmark = pytest.mark.skipif(not BASE, reason="set CUA_S1_BASE to run") - - -@pytest.fixture(scope="module") -def tokenizer(): - from transformers import AutoTokenizer - - return AutoTokenizer.from_pretrained(BASE) - - -def test_letter_ids(tokenizer): - ids = [tokenizer.encode(letter, add_special_tokens=False) for letter in LETTERS] - assert ids == [[32 + i] for i in range(26)] - - -def test_fixture_prompt(tokenizer): - inputs = json.loads( - (Path(__file__).parent / "data" / "text_inputs.json").read_text( - encoding="utf-8" - ) - ) - request = map_request(parse_body(json.dumps(inputs["fixture_positive"]).encode())) - text = tokenizer.apply_chat_template( - build_messages(request.state, request.questions[0]), - tokenize=False, - add_generation_prompt=True, - ) - assert text.endswith("<|im_start|>assistant\n\n") - ids = tokenizer(text)["input_ids"] - assert len(ids) == 218 - assert ids == tokenizer(text, add_special_tokens=False)["input_ids"] From 8367333a5d0e115c4c4ad366b67a91a0553449ce Mon Sep 17 00:00:00 2001 From: Tianyao Wu Date: Wed, 30 Sep 2026 17:53:07 +0800 Subject: [PATCH 10/10] cua_s1: trim the native worker to a minimal eager path Run each prompt eagerly with cuBLASLt's first-choice GEMM algorithms; CUDA Graphs and GEMM tuning move to a follow-up. Parse requests with serde_json instead of emulating CPython's json module, serve from main.rs without the extra configuration, and check the attention kernel against a float64 reference instead of a second kernel. Signed-off-by: Tianyao Wu --- Cargo.lock | 191 +---- recipe/cua_s1/export_text_merged.py | 108 +-- recipe/cua_s1/native.md | 33 +- src/backends/cuda/qwen3_5/README.md | 16 +- src/backends/cuda/qwen3_5/attention.cu | 95 +-- src/backends/cuda/qwen3_5/gemm.cu | 356 +-------- src/backends/cuda/qwen3_5/ops.h | 36 +- src/backends/cuda/qwen3_5/runtime.cu | 32 - src/models/cua_s1/native/Cargo.toml | 6 +- src/models/cua_s1/native/README.md | 38 - src/models/cua_s1/native/src/contract.rs | 468 ++++-------- src/models/cua_s1/native/src/cuda.rs | 93 --- src/models/cua_s1/native/src/engine.rs | 279 ++----- src/models/cua_s1/native/src/json.rs | 245 ++++++ src/models/cua_s1/native/src/lib.rs | 3 +- src/models/cua_s1/native/src/main.rs | 200 +++-- src/models/cua_s1/native/src/model.rs | 319 +------- src/models/cua_s1/native/src/pyjson.rs | 858 ---------------------- src/models/cua_s1/native/src/server.rs | 227 ------ src/models/cua_s1/native/tests/kernels.rs | 168 ++--- 20 files changed, 671 insertions(+), 3100 deletions(-) delete mode 100644 src/models/cua_s1/native/README.md create mode 100644 src/models/cua_s1/native/src/json.rs delete mode 100644 src/models/cua_s1/native/src/pyjson.rs delete mode 100644 src/models/cua_s1/native/src/server.rs diff --git a/Cargo.lock b/Cargo.lock index 41ba0a9b..16864a16 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -31,56 +31,6 @@ version = "0.2.21" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "683d7910e743518b0e34f1186f92494becacb047c7b6bf616c96772180fef923" -[[package]] -name = "anstream" -version = "1.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "824a212faf96e9acacdbd09febd34438f8f711fb84e09a8916013cd7815ca28d" -dependencies = [ - "anstyle", - "anstyle-parse", - "anstyle-query", - "anstyle-wincon", - "colorchoice", - "is_terminal_polyfill", - "utf8parse", -] - -[[package]] -name = "anstyle" -version = "1.0.14" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "940b3a0ca603d1eade50a4846a2afffd5ef57a9feac2c0e2ec2e14f9ead76000" - -[[package]] -name = "anstyle-parse" -version = "1.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "52ce7f38b242319f7cabaa6813055467063ecdc9d355bbb4ce0c68908cd8130e" -dependencies = [ - "utf8parse", -] - -[[package]] -name = "anstyle-query" -version = "1.1.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" -dependencies = [ - "windows-sys 0.61.2", -] - -[[package]] -name = "anstyle-wincon" -version = "3.0.11" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" -dependencies = [ - "anstyle", - "once_cell_polyfill", - "windows-sys 0.61.2", -] - [[package]] name = "anyhow" version = "1.0.104" @@ -169,15 +119,6 @@ version = "2.13.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3ded4057c258ba199e2d26386d3af3780957ecaee6c4ef4041c6b4b8b97c0b06" -[[package]] -name = "block-buffer" -version = "0.10.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71" -dependencies = [ - "generic-array", -] - [[package]] name = "bumpalo" version = "3.20.3" @@ -228,56 +169,10 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06" dependencies = [ "cfg-if", - "cpufeatures 0.3.1", + "cpufeatures", "rand_core 0.10.1", ] -[[package]] -name = "clap" -version = "4.6.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "aa8876b300ab35ba921adea3dfd70157a46249b33f95c9084ae5709785478946" -dependencies = [ - "clap_builder", - "clap_derive", -] - -[[package]] -name = "clap_builder" -version = "4.6.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ec0797fb7aeb1406c84efac526901f7ec3ead2124f946b494e72879d4b54704d" -dependencies = [ - "anstream", - "anstyle", - "clap_lex", - "strsim", -] - -[[package]] -name = "clap_derive" -version = "4.6.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f9c751b79415d4e559e3d1fcf128e09e720eb673a06d26cf6f392d37d75b66e0" -dependencies = [ - "heck", - "proc-macro2", - "quote", - "syn 3.0.6", -] - -[[package]] -name = "clap_lex" -version = "1.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1c133bc6a41be0d194c306b5506d15e6feeea7b1d6604bd3f8310dfb2ca96486" - -[[package]] -name = "colorchoice" -version = "1.0.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" - [[package]] name = "compact_str" version = "0.9.1" @@ -293,15 +188,6 @@ dependencies = [ "static_assertions", ] -[[package]] -name = "cpufeatures" -version = "0.2.17" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280" -dependencies = [ - "libc", -] - [[package]] name = "cpufeatures" version = "0.3.1" @@ -342,16 +228,6 @@ version = "0.2.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5" -[[package]] -name = "crypto-common" -version = "0.1.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" -dependencies = [ - "generic-array", - "typenum", -] - [[package]] name = "darling" version = "0.20.11" @@ -427,16 +303,6 @@ dependencies = [ "syn 2.0.119", ] -[[package]] -name = "digest" -version = "0.10.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" -dependencies = [ - "block-buffer", - "crypto-common", -] - [[package]] name = "displaydoc" version = "0.2.7" @@ -569,16 +435,6 @@ dependencies = [ "slab", ] -[[package]] -name = "generic-array" -version = "0.14.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" -dependencies = [ - "typenum", - "version_check", -] - [[package]] name = "getrandom" version = "0.2.17" @@ -648,12 +504,6 @@ version = "0.17.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" -[[package]] -name = "heck" -version = "0.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" - [[package]] name = "http" version = "1.5.0" @@ -886,12 +736,6 @@ version = "2.12.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "791930b43c0d5973160d90a8f3894509f2b273430f5c5c73b668636d0287c5c0" -[[package]] -name = "is_terminal_polyfill" -version = "1.70.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695" - [[package]] name = "itertools" version = "0.14.0" @@ -1056,14 +900,12 @@ version = "0.1.0" dependencies = [ "anyhow", "axum", - "clap", "half", - "http-body-util", "libloading", "memmap2", "safetensors", + "serde", "serde_json", - "sha2", "tokenizers", "tokio", ] @@ -1084,12 +926,6 @@ version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" -[[package]] -name = "once_cell_polyfill" -version = "1.70.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" - [[package]] name = "onig" version = "6.5.3" @@ -1562,17 +1398,6 @@ dependencies = [ "serde", ] -[[package]] -name = "sha2" -version = "0.10.9" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" -dependencies = [ - "cfg-if", - "cpufeatures 0.2.17", - "digest", -] - [[package]] name = "shlex" version = "2.0.1" @@ -1893,12 +1718,6 @@ version = "0.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" -[[package]] -name = "typenum" -version = "1.20.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b6f5e870be6c3b371b77fe0ee0bafb859fa4964b4404c27de1d380043c4dda20" - [[package]] name = "unicode-ident" version = "1.0.26" @@ -1950,12 +1769,6 @@ version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" -[[package]] -name = "utf8parse" -version = "0.2.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" - [[package]] name = "version_check" version = "0.9.5" diff --git a/recipe/cua_s1/export_text_merged.py b/recipe/cua_s1/export_text_merged.py index 4036bcd2..a2741560 100644 --- a/recipe/cua_s1/export_text_merged.py +++ b/recipe/cua_s1/export_text_merged.py @@ -1,98 +1,30 @@ -"""Export Qwen3.5-4B with the Cua-S1 `text` adapter merged, for the native worker. +"""Export Qwen3.5-4B with the Cua-S1 `text` adapter merged into the bfloat16 weights, +for the native worker (recipe/cua_s1/native.md). Run it in the reference worker's +environment (recipe/cua_s1/text.md): PYTHONPATH=src .venv/bin/python recipe/cua_s1/export_text_merged.py \ --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text \ --out weights/cua-s1-4b-0.2-text-merged - -Run it in the environment of the text worker (recipe/cua_s1/text.md). It loads the -model as that worker does, merges the adapter into the bfloat16 weights with PEFT's -`merge_and_unload` on the given device, and writes the checkpoint and tokenizer -files to --out, plus `cua_s1_export.json` with the revisions it was made from (read -from the Hugging Face download metadata) and the SHA-256 of the tokenizer.json it -wrote. The native worker reads that record and checks the tokenizer against it: -Transformers writes the pre-tokenizer rule it actually uses into this file, which -is not the one in the base repository's tokenizer.json. """ -from __future__ import annotations - import argparse -import hashlib import json -import re -import time from pathlib import Path -import peft -import torch -import transformers - -from models.cua_s1.text.contract import ADAPTER_REPO, BASE_REPO, LETTERS -from models.cua_s1.text.model import TextModel, downloaded_revision - - -def base_revision(base: Path) -> str | None: - """The commit Hugging Face recorded when it downloaded config.json, if any.""" - meta = base / ".cache/huggingface/download/config.json.metadata" - try: - first = meta.read_text().splitlines()[0].strip() - except (OSError, IndexError): - return None - return first if re.fullmatch(r"[0-9a-f]{40}", first) else None - - -def main() -> None: - parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) - parser.add_argument("--base", required=True, type=Path) - parser.add_argument("--adapter", required=True, type=Path) - parser.add_argument("--out", required=True, type=Path) - parser.add_argument("--device", default="cuda") - args = parser.parse_args() - # The native worker refuses an export that does not record both revisions. - revisions = { - "base": base_revision(args.base), - "adapter": downloaded_revision(args.adapter), - } - missing = [name for name, revision in revisions.items() if revision is None] - if missing: - parser.error( - f"no download metadata for the {' and '.join(missing)} weights; " - "download them with `hf download --revision ... --local-dir ...`" - ) - - started = time.perf_counter() - loaded = TextModel(str(args.base), str(args.adapter), args.device, "bfloat16") - model = loaded.model.merge_and_unload().eval() - model.save_pretrained(args.out, safe_serialization=True, max_shard_size="5GB") - loaded.tokenizer.save_pretrained(args.out) - tokenizer = (args.out / "tokenizer.json").read_bytes() - record = { - "format": "cua-s1-text-merged/1", - "base": {"repo": BASE_REPO, "revision": revisions["base"]}, - "adapter": { - "repo": ADAPTER_REPO, - "revision": revisions["adapter"], - "subfolder": "text", - }, - "merge": { - "method": "peft merge_and_unload", - "device": args.device, - "dtype": "bfloat16", - "torch": torch.__version__, - "transformers": transformers.__version__, - "peft": peft.__version__, - }, - "letters": LETTERS, - "tokenizer": { - "file": "tokenizer.json", - "sha256": hashlib.sha256(tokenizer).hexdigest(), - "saved_by": f"transformers {transformers.__version__}", - }, - } - (args.out / "cua_s1_export.json").write_text(json.dumps(record, indent=2) + "\n") - print(f"exported to {args.out} in {time.perf_counter() - started:.1f} s") - print(json.dumps(record)) - - -if __name__ == "__main__": - main() +from models.cua_s1.text.model import TextModel + +parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) +parser.add_argument("--base", required=True) +parser.add_argument("--adapter", required=True) +parser.add_argument("--out", required=True, type=Path) +args = parser.parse_args() + +loaded = TextModel(args.base, args.adapter, "cuda", "bfloat16") +loaded.model.merge_and_unload().save_pretrained(args.out, max_shard_size="5GB") +# Transformers writes the pre-tokenizer rule it uses into tokenizer.json; the native +# worker tokenizes with that file. +loaded.tokenizer.save_pretrained(args.out) +# The native worker refuses a directory without this marker, such as the base model. +(args.out / "cua_s1_export.json").write_text( + json.dumps({"format": "cua-s1-text-merged/1"}) +) diff --git a/recipe/cua_s1/native.md b/recipe/cua_s1/native.md index e4f8da45..cf9e9961 100644 --- a/recipe/cua_s1/native.md +++ b/recipe/cua_s1/native.md @@ -1,21 +1,15 @@ # Cua-S1 4B 0.2 native text worker -The native worker ([`src/models/cua_s1/native/`](../../src/models/cua_s1/native/)) serves the `text` adapter like the reference worker in [`text.md`](text.md), with the forward pass on the CUDA kernels in [`src/backends/cuda/qwen3_5/`](../../src/backends/cuda/qwen3_5/). It needs an NVIDIA GPU with compute capability 8.0 or newer and was measured on an RTX 6000 Ada (sm_89). +The native worker ([`src/models/cua_s1/native/`](../../src/models/cua_s1/native/)) serves the `text` adapter like the reference worker in [`text.md`](text.md), with the Qwen3.5-4B forward pass on the CUDA kernels in [`src/backends/cuda/qwen3_5/`](../../src/backends/cuda/qwen3_5/) and no Python or PyTorch. It needs an NVIDIA GPU with compute capability 8.0 or newer; only an RTX 6000 Ada (sm_89) with CUDA 13.2 has been run. -Run the commands from the repository root. The reference worker's setup from `text.md` is needed once, to export the merged weights. - -## Build +Run the commands from the repository root. Pass your GPU's compute capability to `build.sh` (89 for Ada, 80 for A100, 90 for H100); the worker finds `libqwen3_5_cuda.so` next to its executable, or at `CUA_S1_CUDA_LIB`: ```sh src/backends/cuda/qwen3_5/build.sh target/release 89 # needs nvcc and cuBLASLt cargo build --release --locked -p omni-cua-s1-native ``` -Pass your GPU's compute capability to `build.sh` (89 for Ada, 80 for A100, 90 for H100). Only CUDA 13.2 on sm_89 has been run. The worker finds `libqwen3_5_cuda.so` next to its executable; `--cuda-lib` (or `CUA_S1_CUDA_LIB`) points elsewhere. - -## Export the merged weights - -The worker loads Qwen3.5-4B with the `text` adapter already merged into the bfloat16 weights. With the weights downloaded as in `text.md`, in the reference worker's environment: +The worker loads the weights with the `text` adapter merged in. With the reference worker's environment and weights from `text.md`, export them once (about 8.5 GB): ```sh PYTHONPATH=src .venv/bin/python recipe/cua_s1/export_text_merged.py \ @@ -23,31 +17,18 @@ PYTHONPATH=src .venv/bin/python recipe/cua_s1/export_text_merged.py \ --out weights/cua-s1-4b-0.2-text-merged ``` -This writes about 8.5 GB: the checkpoint, the tokenizer files, and `cua_s1_export.json`, which records the revisions and the SHA-256 of `tokenizer.json`. The worker refuses a `tokenizer.json` that does not match it, and warns if the revisions are not the pinned ones. - -## Start the worker +Start the worker (`CUA_S1_HOST` and `CUA_S1_PORT` default to `127.0.0.1` and `8000`), then the frontend and requests as in `text.md`: ```sh -target/release/omni-cua-s1-native --model weights/cua-s1-4b-0.2-text-merged \ - --gemm-plans weights/gemm-plans.json --gemm-search --port 8000 +CUA_S1_MODEL=weights/cua-s1-4b-0.2-text-merged target/release/omni-cua-s1-native ``` -The first start tunes the GEMMs for this GPU and writes the choices to `--gemm-plans`; with `--gemm-search` that takes about a minute. Later starts read the file and are ready in about 2 seconds, and give the same results each time. Without `--gemm-search`, tuning takes a few seconds and the worker is somewhat slower for prompts of 200 to 2,000 tokens. The file records the GPU, the cuBLASLt version and `--graph-max-tokens`; the worker refuses a file that does not match and leaves it alone, so keep one per GPU model and CUDA version, and remove it to tune again. +Each question is one eager forward pass over its prompt; the final hidden state at the last position times the 26 letter rows of the output projection gives the option probabilities. The probabilities are not bitwise identical to the reference worker's, since the adapter is merged and the kernels differ; they are held to the tolerance in [`src/models/cua_s1/README.md`](../../src/models/cua_s1/README.md#validation). Error messages are worded differently, and bodies nested more than 127 levels deep are refused. -The same options exist as environment variables (`CUA_S1_MODEL`, `CUA_S1_PORT`, `CUA_S1_GEMM_PLANS`, ...; see `--help`), as do the reference worker's request limits and `CUA_S1_API_KEY`. The Rust frontend and the requests are as in `text.md`. - -## Tests - -Tests without a GPU, then the kernel checks (attention against a float32 kernel, the gated delta rule against a float64 token-by-token reference, and the GEMM plans): +The request tests need no GPU; the kernel tests compare attention and the chunked Gated DeltaNet prefill with float64 references: ```sh cargo test -p omni-cua-s1-native CUA_S1_CUDA_LIB=$PWD/target/release/libqwen3_5_cuda.so \ cargo test --release -p omni-cua-s1-native --test kernels -- --ignored ``` - -## Not covered - -- `score` and `noul` questions, and the `multimodal` adapter, as in the reference worker. -- More than one request at a time: the worker answers one decision at a time, like the reference worker. -- GPUs other than sm_89, and CUDA versions other than 13.2, have not been run. The GEMM plans are per GPU and cuBLASLt version. diff --git a/src/backends/cuda/qwen3_5/README.md b/src/backends/cuda/qwen3_5/README.md index 8863ab7e..b84f50b5 100644 --- a/src/backends/cuda/qwen3_5/README.md +++ b/src/backends/cuda/qwen3_5/README.md @@ -1,21 +1,9 @@ # Qwen3.5 prefill operations -CUDA kernels for a prefill-only Qwen3.5 forward pass, built into `libqwen3_5_cuda.so`: +CUDA kernels for a prefill-only Qwen3.5 forward pass, built into `libqwen3_5_cuda.so` with a C interface ([`ops.h`](ops.h)), so that a Rust model engine loads it at run time and builds without a CUDA toolkit. The Cua-S1 native worker ([`src/models/cua_s1/native/`](../../../models/cua_s1/native/)) uses it and keeps the layer loop and buffers. ```sh src/backends/cuda/qwen3_5/build.sh [compute capability, default 89] ``` -The library has a C interface ([`ops.h`](ops.h)): the operations, plus the few CUDA runtime calls a caller needs (allocation, copies, streams, graph capture), so that a Rust model engine can load it at run time and build without a CUDA toolkit. The Cua-S1 native worker ([`src/models/cua_s1/native/`](../../../models/cua_s1/native/)) uses it; the layer loop, buffers, CUDA graphs and GEMM algorithm choice stay in that model engine. - -| File | Operations | -| --- | --- | -| `norm.cu` | Zero-centred RMSNorm, the residual add fused with the next norm, and the gated RMSNorm of the Gated DeltaNet output. | -| `elementwise.cu` | Embedding lookup; the Gated DeltaNet causal convolution with SiLU and its gates; the attention output gate; SiLU(gate) * up. | -| `attention.cu` | q/k RMSNorm and partial rotary embedding; causal attention with grouped KV heads (head dim 256) on tensor cores, FlashAttention-2 style. | -| `gdn_prefill.cu` | The chunked gated delta rule (chunks of 64) in three kernels, split the way flash-linear-attention splits it: per-chunk preparation, the state carried from chunk to chunk, and the per-chunk output. | -| `gemm.cu` | bfloat16 GEMMs through cuBLASLt with float32 accumulation, algorithm tuning, and saving and loading the tuned choices. | -| `runtime.cu` | The CUDA runtime calls. | -| `mma.cuh`, `common.cuh` | `mma.sync`, `ldmatrix` and `cp.async` helpers, and shared device helpers. | - -The norm, elementwise and q/k preparation kernels round to bfloat16 at the same points as the Transformers implementation (`modeling_qwen3_5.py`). Attention and the Gated DeltaNet prefill keep some intermediate results in bfloat16, as FlashAttention and flash-linear-attention do, where Transformers' float32 fallback for the gated delta rule does not; the model is checked end to end against the float32 reference worker. Tensor-core kernels need sm_80 or newer; they are measured on sm_89 (RTX 6000 Ada). Split-K reductions that accumulate into the output in place are not used, and plans that use them are refused on import, so a given GEMM algorithm always gives the same result. `build.sh` also embeds PTX, but only sm_89 has been run. +The norm, elementwise and q/k preparation kernels round to bfloat16 where Transformers (`modeling_qwen3_5.py`) does. Attention (FlashAttention-2 style, on tensor cores) and the chunked gated delta rule keep some intermediate results in bfloat16, as FlashAttention and flash-linear-attention do. GEMMs go through cuBLASLt with its first heuristic choice. Tensor-core kernels need sm_80 or newer; only sm_89 has been run. diff --git a/src/backends/cuda/qwen3_5/attention.cu b/src/backends/cuda/qwen3_5/attention.cu index 75489deb..8c03666e 100644 --- a/src/backends/cuda/qwen3_5/attention.cu +++ b/src/backends/cuda/qwen3_5/attention.cu @@ -5,7 +5,7 @@ // four warps of 16 rows each, and walks the keys up to its last query in tiles of // 32, keeping the output and the online softmax in registers. The probabilities are // rounded to bfloat16 for the P*V product, as in flash attention; the running sums -// stay float32. cs1_attention_simple is a plain float32 version kept for checking. +// stay float32. #include "common.cuh" #include "mma.cuh" #include "ops.h" @@ -15,9 +15,6 @@ namespace { constexpr int DH = 256; // head dim constexpr int PER = DH / 32; // values per lane -constexpr int QB = 16; // queries per block, two per warp -constexpr int KB = 32; // keys per shared-memory tile -constexpr int ATTN_THREADS = 256; // One warp per (token, head), q heads first, then k heads. Each lane holds 8 // consecutive dims, so the rotary partner of dim i < 32 (dim i + 32) sits in lane ^ 4. @@ -73,86 +70,6 @@ __global__ void attn_prep_kernel(const bf16* __restrict__ qg, const bf16* __rest } } -// Causal attention, float32 scores and online softmax. A block takes QB queries of -// one head and walks the keys up to its last query in shared-memory tiles. -__global__ void __launch_bounds__(ATTN_THREADS) - attention_kernel(const bf16* __restrict__ q, const bf16* __restrict__ k, const bf16* __restrict__ v, int ldv, - bf16* __restrict__ out, int T, int Hq, int Hk, float scale) { - __shared__ __align__(16) bf16 ks[KB * DH]; - __shared__ __align__(16) bf16 vs[KB * DH]; - const int h = blockIdx.y, hk = h / (Hq / Hk); - const int warp = threadIdx.x / 32, lane = threadIdx.x & 31, d0 = lane * PER; - const int first = blockIdx.x * QB + warp * 2; - - float qv[2][PER], acc[2][PER], m[2], l[2]; -#pragma unroll - for (int r = 0; r < 2; r++) { - const int t = first + r; - if (t < T) { - load8(q + ((size_t)t * Hq + h) * DH + d0, qv[r]); - } else { -#pragma unroll - for (int i = 0; i < PER; i++) qv[r][i] = 0.f; - } -#pragma unroll - for (int i = 0; i < PER; i++) acc[r][i] = 0.f; - m[r] = -INFINITY; - l[r] = 0.f; - } - - const int kv_end = min(T, (int)(blockIdx.x * QB + QB)); - for (int k0 = 0; k0 < kv_end; k0 += KB) { - __syncthreads(); - for (int x = threadIdx.x; x < KB * DH / 8; x += blockDim.x) { - const int j = x / (DH / 8), c = (x % (DH / 8)) * 8, s = k0 + j; - Pack8 kk{}, vv{}; - if (s < T) { - kk = *reinterpret_cast(k + ((size_t)s * Hk + hk) * DH + c); - vv = *reinterpret_cast(v + (size_t)s * ldv + (size_t)hk * DH + c); - } - *reinterpret_cast(ks + j * DH + c) = kk; - *reinterpret_cast(vs + j * DH + c) = vv; - } - __syncthreads(); - const int jn = min(KB, kv_end - k0); - for (int j = 0; j < jn; j++) { - float kv[PER]; - load8(ks + j * DH + d0, kv); - float dot[2] = {0.f, 0.f}; -#pragma unroll - for (int i = 0; i < PER; i++) { - dot[0] = fmaf(qv[0][i], kv[i], dot[0]); - dot[1] = fmaf(qv[1][i], kv[i], dot[1]); - } - dot[0] = warp_sum(dot[0]); - dot[1] = warp_sum(dot[1]); - float vx[PER]; - load8(vs + j * DH + d0, vx); - const int s = k0 + j; -#pragma unroll - for (int r = 0; r < 2; r++) { - if (s > first + r) continue; - const float score = dot[r] * scale; - const float mn = fmaxf(m[r], score); - const float corr = expf(m[r] - mn), p = expf(score - mn); - l[r] = l[r] * corr + p; -#pragma unroll - for (int i = 0; i < PER; i++) acc[r][i] = fmaf(p, vx[i], acc[r][i] * corr); - m[r] = mn; - } - } - } -#pragma unroll - for (int r = 0; r < 2; r++) { - const int t = first + r; - if (t >= T) continue; - float o[PER]; -#pragma unroll - for (int i = 0; i < PER; i++) o[i] = acc[r][i] / l[r]; - store8(out + ((size_t)t * Hq + h) * DH + d0, o); - } -} - // ---- flash attention ---- namespace flash { @@ -325,16 +242,6 @@ extern "C" int cs1_attn_prep(const void* qg, const void* kr, int ld, const void* return cudaGetLastError(); } -extern "C" int cs1_attention_simple(const void* q, const void* k, const void* v, int ldv, void* out, int T, int Hq, - int Hk, int Dh, float scale, void* stream) { - if (Dh != DH || Hk <= 0 || Hq % Hk != 0 || ldv % 8 != 0 || ldv < Hk * Dh || T < 0) return cudaErrorInvalidValue; - if (T == 0) return cudaSuccess; - attention_kernel<<(stream)>>>( - static_cast(q), static_cast(k), static_cast(v), ldv, - static_cast(out), T, Hq, Hk, scale); - return cudaGetLastError(); -} - extern "C" int cs1_attention(const void* q, const void* k, const void* v, int ldv, void* out, int T, int Hq, int Hk, int Dh, float scale, void* stream) { if (Dh != flash::D || Hk <= 0 || Hq % Hk != 0 || ldv % 8 != 0 || ldv < Hk * Dh || T < 0) diff --git a/src/backends/cuda/qwen3_5/gemm.cu b/src/backends/cuda/qwen3_5/gemm.cu index 3818be61..7de47f57 100644 --- a/src/backends/cuda/qwen3_5/gemm.cu +++ b/src/backends/cuda/qwen3_5/gemm.cu @@ -4,29 +4,13 @@ // y^T [N, M] = (w viewed as [K, N])^T * (x viewed as [K, M]); y's rows may be // strided (ldy >= N), so one GEMM can fill a slice of a wider buffer. // -// Algorithms: cs1_gemm_tune times cuBLASLt's candidates for a shape with L2 flushed -// before every call and keeps the fastest: the heuristic's shortlist, or (exhaustive) -// each algorithm id with each tile, stage count, custom option and swizzle it supports -// and split-K factors of 1 to 6, 8, 12 and 16, as far as cuBLASLt accepts them for the -// shape (a first pass of one call each keeps 12 to time properly). -// Many configurations are within noise of each other, so two searches often keep -// different ones with about the same speed. -// It replaces the heuristic's first choice only when it is more than 3% faster, so -// near ties rarely change between runs, and -// cs1_gemm_export / cs1_gemm_import let a caller keep the choices across runs, for the -// cuBLASLt version they were tuned with. A shape that was not tuned borrows the -// algorithm tuned for a nearby M with the same N, K and ldy (see plan_for), or takes -// the heuristic's first choice. Split-K reductions that accumulate into the output in place are -// excluded, since their order, and so the rounding, is not fixed. +// Each shape uses cuBLASLt's first heuristic choice, excluding split-K reductions that +// accumulate into the output in place, since their order, and so the rounding, is not +// fixed. #include -#include -#include -#include -#include #include #include -#include #include "ops.h" @@ -36,34 +20,17 @@ struct Plan { cublasLtMatmulDesc_t op = nullptr; cublasLtMatrixLayout_t a = nullptr, b = nullptr, c = nullptr; cublasLtMatmulAlgo_t algo{}; - bool tuned = false; }; using Key = std::tuple; // M, N, K, ldy -constexpr size_t FLUSH_BYTES = 256u << 20; - struct Gemm { cublasLtHandle_t handle = nullptr; void* workspace = nullptr; size_t workspace_bytes = 0; std::map plans; - // tuning only: a buffer larger than L2, a sink for its reads, and two events - void* flush = nullptr; - int* sink = nullptr; - cudaEvent_t e0 = nullptr, e1 = nullptr; }; -void release_tuning(Gemm& g) { - if (g.flush) cudaFree(g.flush); - if (g.sink) cudaFree(g.sink); - if (g.e0) cudaEventDestroy(g.e0); - if (g.e1) cudaEventDestroy(g.e1); - g.flush = nullptr; - g.sink = nullptr; - g.e0 = g.e1 = nullptr; -} - int status(cublasStatus_t s) { return s == CUBLAS_STATUS_SUCCESS ? 0 : 1000 + (int)s; } void destroy(Plan& p) { @@ -86,7 +53,8 @@ int describe(int M, int N, int K, int ldy, Plan& p) { return 0; } -int heuristics(Gemm& g, const Plan& p, int want, std::vector& out) { +// The heuristic's first choice, without in-place split-K reductions. +int first_choice(Gemm& g, Plan& p) { cublasLtMatmulPreference_t pref; cublasStatus_t s = cublasLtMatmulPreferenceCreate(&pref); if (s != CUBLAS_STATUS_SUCCESS) return status(s); @@ -95,100 +63,13 @@ int heuristics(Gemm& g, const Plan& p, int want, std::vector -std::vector cap_array(const cublasLtMatmulAlgo_t& algo, cublasLtMatmulAlgoCapAttributes_t attr) { - size_t bytes = 0; - cublasLtMatmulAlgoCapGetAttribute(&algo, attr, nullptr, 0, &bytes); - std::vector v(bytes / sizeof(T)); - if (bytes) cublasLtMatmulAlgoCapGetAttribute(&algo, attr, v.data(), bytes, &bytes); - return v; -} - -template -T cap(const cublasLtMatmulAlgo_t& algo, cublasLtMatmulAlgoCapAttributes_t attr) { - T v{}; - size_t n; - cublasLtMatmulAlgoCapGetAttribute(&algo, attr, &v, sizeof v, &n); - return v; -} - -// The configurations cuBLASLt accepts for the shape, within the workspace: each -// algorithm id with each tile, stage count, custom option and swizzle it supports, and -// split-K factors from `splits` (reduced in the compute or the output type, not in -// place). Other attributes stay at their defaults. -int every_config(Gemm& g, const Plan& p, std::vector& out) { - int ids[256], nids = 0; - const cublasStatus_t s = cublasLtMatmulAlgoGetIds(g.handle, CUBLAS_COMPUTE_32F, CUDA_R_32F, CUDA_R_16BF, - CUDA_R_16BF, CUDA_R_16BF, CUDA_R_16BF, 256, ids, &nids); - if (s != CUBLAS_STATUS_SUCCESS) return status(s); - const int splits[] = {1, 2, 3, 4, 5, 6, 8, 12, 16}; - const uint32_t schemes[] = {CUBLASLT_REDUCTION_SCHEME_NONE, CUBLASLT_REDUCTION_SCHEME_COMPUTE_TYPE, - CUBLASLT_REDUCTION_SCHEME_OUTPUT_TYPE}; - for (int i = 0; i < nids; i++) { - cublasLtMatmulAlgo_t base; - if (cublasLtMatmulAlgoInit(g.handle, CUBLAS_COMPUTE_32F, CUDA_R_32F, CUDA_R_16BF, CUDA_R_16BF, CUDA_R_16BF, - CUDA_R_16BF, ids[i], &base) != CUBLAS_STATUS_SUCCESS) - continue; - auto tiles = cap_array(base, CUBLASLT_ALGO_CAP_TILE_IDS); - auto stages = cap_array(base, CUBLASLT_ALGO_CAP_STAGES_IDS); - if (tiles.empty()) tiles.push_back(CUBLASLT_MATMUL_TILE_UNDEFINED); - if (stages.empty()) stages.push_back(CUBLASLT_MATMUL_STAGES_UNDEFINED); - const int splitk_ok = cap(base, CUBLASLT_ALGO_CAP_SPLITK_SUPPORT); - const uint32_t red_mask = cap(base, CUBLASLT_ALGO_CAP_REDUCTION_SCHEME_MASK); - const int swizzle_ok = cap(base, CUBLASLT_ALGO_CAP_CTA_SWIZZLING_SUPPORT); - const int custom_max = cap(base, CUBLASLT_ALGO_CAP_CUSTOM_OPTION_MAX); - for (uint32_t tile : tiles) - for (uint32_t stage : stages) - for (int custom = 0; custom <= custom_max; custom++) - for (int swz = 0; swz <= swizzle_ok; swz++) - for (int sk : splits) { - if (sk > 1 && !splitk_ok) break; - for (uint32_t red : schemes) { - if ((sk == 1) != (red == CUBLASLT_REDUCTION_SCHEME_NONE)) continue; - if (red != CUBLASLT_REDUCTION_SCHEME_NONE && !(red_mask & red)) continue; - cublasLtMatmulAlgo_t a = base; - cublasLtMatmulAlgoConfigSetAttribute(&a, CUBLASLT_ALGO_CONFIG_TILE_ID, &tile, - sizeof tile); - cublasLtMatmulAlgoConfigSetAttribute(&a, CUBLASLT_ALGO_CONFIG_STAGES_ID, &stage, - sizeof stage); - cublasLtMatmulAlgoConfigSetAttribute(&a, CUBLASLT_ALGO_CONFIG_CUSTOM_OPTION, &custom, - sizeof custom); - cublasLtMatmulAlgoConfigSetAttribute(&a, CUBLASLT_ALGO_CONFIG_CTA_SWIZZLING, &swz, - sizeof swz); - cublasLtMatmulAlgoConfigSetAttribute(&a, CUBLASLT_ALGO_CONFIG_SPLITK_NUM, &sk, - sizeof sk); - cublasLtMatmulAlgoConfigSetAttribute(&a, CUBLASLT_ALGO_CONFIG_REDUCTION_SCHEME, &red, - sizeof red); - if (usable(g, p, a)) out.push_back(a); - } - } - } + if (found == 0 || r.state != CUBLAS_STATUS_SUCCESS) return status(CUBLAS_STATUS_NOT_SUPPORTED); + p.algo = r.algo; return 0; } @@ -202,62 +83,15 @@ int plan_for(Gemm& g, int M, int N, int K, int ldy, Plan*& out) { } Plan p; int rc = describe(M, N, K, ldy, p); + if (rc == 0) rc = first_choice(g, p); if (rc != 0) { destroy(p); return rc; } - // Borrow a tuned algorithm: the one for the smallest tuned M above, if that M is at - // most twice this one; the one for the largest tuned M below, if no larger M was - // tuned; else take the heuristic's first choice. - const Plan* above = nullptr; - const Plan* below = nullptr; - int above_m = 0, below_m = 0; - for (auto& kv : g.plans) { - const auto [m, n, k, l] = kv.first; - if (n != N || k != K || l != ldy || !kv.second.tuned) continue; - if (m > M && (!above || m < above_m)) above = &kv.second, above_m = m; - if (m < M && (!below || m > below_m)) below = &kv.second, below_m = m; - } - if (above && above_m <= 2 * M && usable(g, p, above->algo)) { - p.algo = above->algo; - } else if (!above && below && usable(g, p, below->algo)) { - p.algo = below->algo; - } else { - std::vector cands; - rc = heuristics(g, p, 1, cands); - if (rc != 0) { - destroy(p); - return rc; - } - p.algo = cands[0].algo; - } out = &g.plans.emplace(key, p).first->second; return 0; } -// With CUA_S1_GEMM_LOG set, print each tuned choice to stderr. -void log_choice(int M, int N, int K, const cublasLtMatmulAlgo_t& a, float ms, float first_ms, size_t pick) { - static const bool on = std::getenv("CUA_S1_GEMM_LOG") != nullptr; - if (!on) return; - int tile = 0, stages = 0, splitk = 0, inner = 0, id = 0; - size_t n; - cublasLtMatmulAlgoConfigGetAttribute(&a, CUBLASLT_ALGO_CONFIG_ID, &id, sizeof(int), &n); - cublasLtMatmulAlgoConfigGetAttribute(&a, CUBLASLT_ALGO_CONFIG_TILE_ID, &tile, sizeof(int), &n); - cublasLtMatmulAlgoConfigGetAttribute(&a, CUBLASLT_ALGO_CONFIG_STAGES_ID, &stages, sizeof(int), &n); - cublasLtMatmulAlgoConfigGetAttribute(&a, CUBLASLT_ALGO_CONFIG_SPLITK_NUM, &splitk, sizeof(int), &n); - cublasLtMatmulAlgoConfigGetAttribute(&a, CUBLASLT_ALGO_CONFIG_INNER_SHAPE_ID, &inner, sizeof(int), &n); - fprintf(stderr, "gemm %5d x %5d x %5d: candidate %zu, algo %d tile %d stages %d splitK %d inner %d, %.1f us (first %.1f us)\n", - M, N, K, pick, id, tile, stages, splitk, inner, ms * 1e3f, first_ms * 1e3f); -} - -// Read a buffer larger than L2, so the next call finds none of its operands cached. -__global__ void flush_l2(const int4* p, size_t n, int* sink) { - int acc = 0; - for (size_t i = blockIdx.x * (size_t)blockDim.x + threadIdx.x; i < n; i += (size_t)gridDim.x * blockDim.x) - acc ^= p[i].x ^ p[i].w; - if (acc == 0x7fffffff) *sink = acc; -} - } // namespace extern "C" void* cs1_gemm_create(size_t workspace_bytes) { @@ -276,181 +110,11 @@ extern "C" void cs1_gemm_destroy(void* gemm) { Gemm* g = static_cast(gemm); if (!g) return; for (auto& kv : g->plans) destroy(kv.second); - release_tuning(*g); if (g->workspace) cudaFree(g->workspace); cublasLtDestroy(g->handle); delete g; } -extern "C" int cs1_gemm_tune(void* gemm, const void* x, const void* w, void* y, int M, int N, int K, int ldy, - int exhaustive, void* stream) { - Gemm* g = static_cast(gemm); - if (!g || M <= 0 || N <= 0 || K <= 0 || ldy < N) return cudaErrorInvalidValue; - const Key key{M, N, K, ldy}; - auto it = g->plans.find(key); - if (it != g->plans.end() && it->second.tuned) return 0; - if (it != g->plans.end()) { - destroy(it->second); - g->plans.erase(it); - } - Plan p; - int rc = describe(M, N, K, ldy, p); - std::vector shortlist; - if (rc == 0) rc = heuristics(*g, p, 16, shortlist); - if (rc != 0) { - destroy(p); - return rc; - } - // candidates: the heuristic's shortlist first, then (exhaustive) those of every_config - std::vector cands; - for (auto& r : shortlist) cands.push_back(r.algo); - if (exhaustive) { - std::vector all; - rc = every_config(*g, p, all); - if (rc != 0) { - destroy(p); - return rc; - } - for (auto& a : all) - if (std::none_of(cands.begin(), cands.end(), - [&](const cublasLtMatmulAlgo_t& c) { return std::memcmp(&c, &a, sizeof a) == 0; })) - cands.push_back(a); - } - cudaStream_t st = static_cast(stream); - if (!g->flush) { - if (cudaMalloc(&g->flush, FLUSH_BYTES) != cudaSuccess || cudaMalloc(&g->sink, sizeof(int)) != cudaSuccess || - cudaMemsetAsync(g->flush, 0, FLUSH_BYTES, st) != cudaSuccess || cudaEventCreate(&g->e0) != cudaSuccess || - cudaEventCreate(&g->e1) != cudaSuccess) { - release_tuning(*g); - destroy(p); - cudaGetLastError(); - return (int)cudaErrorMemoryAllocation; - } - } - const float alpha = 1.f, beta = 0.f; - // median of `reps` calls, each after an L2 flush; a huge value if the call fails - auto time = [&](const cublasLtMatmulAlgo_t& algo, int reps) { - if (cublasLtMatmul(g->handle, p.op, &alpha, w, p.a, x, p.b, &beta, y, p.c, y, p.c, &algo, g->workspace, - g->workspace_bytes, st) != CUBLAS_STATUS_SUCCESS) { - cudaGetLastError(); - return 1e30f; - } - std::vector times; - for (int r = 0; r < reps; r++) { - flush_l2<<<1024, 256, 0, st>>>(static_cast(g->flush), FLUSH_BYTES / sizeof(int4), g->sink); - cudaEventRecord(g->e0, st); - cublasLtMatmul(g->handle, p.op, &alpha, w, p.a, x, p.b, &beta, y, p.c, y, p.c, &algo, g->workspace, - g->workspace_bytes, st); - cudaEventRecord(g->e1, st); - if (cudaEventSynchronize(g->e1) != cudaSuccess) return 1e30f; - float ms = 0.f; - cudaEventElapsedTime(&ms, g->e0, g->e1); - times.push_back(ms); - } - std::sort(times.begin(), times.end()); - return times[reps / 2]; - }; - // with many candidates, one timed call each picks the 12 to time properly - std::vector keep; - if (cands.size() > 16) { - std::vector> quick; - for (size_t i = 0; i < cands.size(); i++) quick.push_back({time(cands[i], 1), i}); - std::sort(quick.begin(), quick.end()); - keep.push_back(0); // the heuristic's first choice, the baseline - for (size_t j = 0; j < quick.size() && keep.size() < 13; j++) - if (quick[j].second != 0 && quick[j].first < 1e30f) keep.push_back(quick[j].second); - } else { - for (size_t i = 0; i < cands.size(); i++) keep.push_back(i); - } - // nine timed calls, or three for shapes that take over 2 ms - const int reps = time(cands[0], 1) > 2.f ? 3 : 9; - std::vector median(cands.size(), 1e30f); - for (size_t i : keep) median[i] = time(cands[i], reps); - // the heuristic's first working choice, unless another is more than 3% faster - size_t pick = 0; - while (pick < median.size() && median[pick] >= 1e30f) pick++; - if (pick == median.size()) { - destroy(p); - return status(CUBLAS_STATUS_NOT_SUPPORTED); - } - const size_t first = pick; - const size_t fastest = std::min_element(median.begin(), median.end()) - median.begin(); - if (median[fastest] < 0.97f * median[pick]) pick = fastest; - p.algo = cands[pick]; - p.tuned = true; - log_choice(M, N, K, p.algo, median[pick], median[first], pick); - g->plans.emplace(key, p); - return (int)cudaGetLastError(); -} - -extern "C" void cs1_gemm_tune_done(void* gemm) { - if (gemm) release_tuning(*static_cast(gemm)); -} - -extern "C" size_t cs1_gemm_export(void* gemm, Cs1GemmPlan* out, size_t cap) { - Gemm* g = static_cast(gemm); - if (!g) return 0; - size_t n = 0; - for (auto& kv : g->plans) { - if (!kv.second.tuned) continue; - if (n < cap) { - const auto [m, nn, k, l] = kv.first; - out[n] = Cs1GemmPlan{m, nn, k, l, (uint64_t)cublasLtGetVersion(), {}}; - static_assert(sizeof(cublasLtMatmulAlgo_t) == sizeof(out[n].algo), "algo layout"); - std::memcpy(out[n].algo, &kv.second.algo, sizeof(out[n].algo)); - } - n++; - } - return n; -} - -extern "C" int cs1_gemm_import(void* gemm, const Cs1GemmPlan* plans, size_t n) { - Gemm* g = static_cast(gemm); - if (!g) return cudaErrorInvalidValue; - // check every plan before using any, so that a rejected file changes nothing - std::vector> checked; - int rc = 0; - for (size_t i = 0; i < n && rc == 0; i++) { - const Cs1GemmPlan& r = plans[i]; - if (r.m <= 0 || r.n <= 0 || r.k <= 0 || r.ldy < r.n) { - rc = cudaErrorInvalidValue; - break; - } - if (r.cublaslt_version != (uint64_t)cublasLtGetVersion()) { - rc = status(CUBLAS_STATUS_NOT_SUPPORTED); - break; - } - Plan p; - rc = describe(r.m, r.n, r.k, r.ldy, p); - if (rc == 0) { - std::memcpy(&p.algo, r.algo, sizeof(p.algo)); - if (reduces_in_place(p.algo) || !usable(*g, p, p.algo)) rc = status(CUBLAS_STATUS_NOT_SUPPORTED); - } - if (rc != 0) { - destroy(p); - break; - } - p.tuned = true; - checked.emplace_back(Key{r.m, r.n, r.k, r.ldy}, p); - } - if (rc != 0) { - for (auto& kv : checked) destroy(kv.second); - return rc; - } - for (auto& [key, p] : checked) { - auto it = g->plans.find(key); - if (it != g->plans.end()) { - destroy(it->second); - it->second = p; - } else { - g->plans.emplace(key, p); - } - } - return 0; -} - -extern "C" size_t cs1_gemm_version(void) { return cublasLtGetVersion(); } - extern "C" int cs1_gemm(void* gemm, const void* x, const void* w, void* y, int M, int N, int K, int ldy, void* stream) { Gemm* g = static_cast(gemm); diff --git a/src/backends/cuda/qwen3_5/ops.h b/src/backends/cuda/qwen3_5/ops.h index e625686c..6f63488e 100644 --- a/src/backends/cuda/qwen3_5/ops.h +++ b/src/backends/cuda/qwen3_5/ops.h @@ -24,8 +24,6 @@ extern "C" { uint32_t cs1_abi_version(void); const char* cs1_error_string(int code); int cs1_set_device(int device); -// Name, compute capability (major * 10 + minor) and SM count of the current device. -int cs1_device_info(char* name, size_t cap, int* compute_capability, int* sms); int cs1_malloc(void** ptr, size_t bytes); int cs1_free(void* ptr); int cs1_stream_create(void** stream); @@ -33,11 +31,6 @@ int cs1_stream_sync(void* stream); // Copy and wait for the copy. int cs1_upload(void* dst, const void* src, size_t bytes, void* stream); int cs1_download(void* dst, const void* src, size_t bytes, void* stream); -// Capture the work queued on `stream` between begin and end into an executable graph. -int cs1_graph_begin(void* stream); -int cs1_graph_end(void* stream, void** exec); -int cs1_graph_launch(void* exec, void* stream); -int cs1_graph_destroy(void* exec); // ---- operations ---- @@ -83,9 +76,6 @@ int cs1_attn_prep(const void* qg, const void* kr, int ld, const void* qw, const // [T, Hk, Dh] in rows of ldv; out [T, Hq, Dh]. int cs1_attention(const void* q, const void* k, const void* v, int ldv, void* out, int T, int Hq, int Hk, int Dh, float scale, void* stream); -// The same in float32 on CUDA cores, two queries per warp: slow, kept for checks. -int cs1_attention_simple(const void* q, const void* k, const void* v, int ldv, void* out, int T, - int Hq, int Hk, int Dh, float scale, void* stream); // x = x * sigmoid(gate), n elements. int cs1_sigmoid_gate(void* x, const void* gate, size_t n, void* stream); @@ -93,32 +83,10 @@ int cs1_sigmoid_gate(void* x, const void* gate, size_t n, void* stream); // out [T, I] = silu(gate) * up, from gate_up [T, 2*I] (gate first) in rows of ld. int cs1_silu_mul(const void* gate_up, int ld, void* out, int T, int I, void* stream); -// y [M, N] (rows of ldy) = x [M, K] * w [N, K]^T through cuBLASLt, float32 accumulation. -// cs1_gemm_tune picks the algorithm for one shape by timing, among the heuristic's -// shortlist or (exhaustive) a wider enumeration (see gemm.cu); it must not -// run during stream capture. cs1_gemm_tune_done frees the buffers tuning used. -// A tuned algorithm for one shape: `algo` holds a cublasLtMatmulAlgo_t, valid for the -// cuBLASLt version (cublasLtGetVersion) it was tuned with. -typedef struct { - int32_t m, n, k, ldy; - uint64_t cublaslt_version; - uint64_t algo[8]; -} Cs1GemmPlan; - +// y [M, N] (rows of ldy) = x [M, K] * w [N, K]^T through cuBLASLt, float32 accumulation, +// with cuBLASLt's first heuristic choice for each shape (see gemm.cu). void* cs1_gemm_create(size_t workspace_bytes); void cs1_gemm_destroy(void* gemm); -int cs1_gemm_tune(void* gemm, const void* x, const void* w, void* y, int M, int N, int K, int ldy, - int exhaustive, void* stream); -void cs1_gemm_tune_done(void* gemm); -// Copy up to `cap` tuned plans to `out`; returns how many there are. -size_t cs1_gemm_export(void* gemm, Cs1GemmPlan* out, size_t cap); -// Use these plans (from cs1_gemm_export, possibly of an earlier run): all of them, or -// none if one was tuned with another cuBLASLt version, fails cuBLASLt's check on this -// device, or reduces split-K in place. Whether a plan was tuned on this GPU model is -// not checked; the caller keeps that with the plans. -int cs1_gemm_import(void* gemm, const Cs1GemmPlan* plans, size_t n); -// cublasLtGetVersion(). -size_t cs1_gemm_version(void); int cs1_gemm(void* gemm, const void* x, const void* w, void* y, int M, int N, int K, int ldy, void* stream); diff --git a/src/backends/cuda/qwen3_5/runtime.cu b/src/backends/cuda/qwen3_5/runtime.cu index deaacc26..8361b83f 100644 --- a/src/backends/cuda/qwen3_5/runtime.cu +++ b/src/backends/cuda/qwen3_5/runtime.cu @@ -2,8 +2,6 @@ // (libqwen3_5_cuda.so) and never links CUDA itself. #include -#include - #include "ops.h" extern "C" { @@ -14,18 +12,6 @@ const char* cs1_error_string(int code) { return cudaGetErrorString(static_cast 0) std::snprintf(name, cap, "%s", prop.name); - *compute_capability = prop.major * 10 + prop.minor; - *sms = prop.multiProcessorCount; - return cudaSuccess; -} - int cs1_malloc(void** ptr, size_t bytes) { return cudaMalloc(ptr, bytes); } int cs1_free(void* ptr) { return cudaFree(ptr); } @@ -48,22 +34,4 @@ int cs1_download(void* dst, const void* src, size_t bytes, void* stream) { return e != cudaSuccess ? e : cudaStreamSynchronize(st); } -int cs1_graph_begin(void* stream) { - return cudaStreamBeginCapture(static_cast(stream), cudaStreamCaptureModeThreadLocal); -} - -int cs1_graph_end(void* stream, void** exec) { - cudaGraph_t graph = nullptr; - cudaError_t e = cudaStreamEndCapture(static_cast(stream), &graph); - if (e == cudaSuccess) e = cudaGraphInstantiate(reinterpret_cast(exec), graph, 0); - if (graph) cudaGraphDestroy(graph); - return e; -} - -int cs1_graph_launch(void* exec, void* stream) { - return cudaGraphLaunch(static_cast(exec), static_cast(stream)); -} - -int cs1_graph_destroy(void* exec) { return cudaGraphExecDestroy(static_cast(exec)); } - } // extern "C" diff --git a/src/models/cua_s1/native/Cargo.toml b/src/models/cua_s1/native/Cargo.toml index 35cf3a8e..2e8caf68 100644 --- a/src/models/cua_s1/native/Cargo.toml +++ b/src/models/cua_s1/native/Cargo.toml @@ -12,15 +12,13 @@ path = "src/main.rs" [dependencies] anyhow = "1.0.100" axum = "0.8.8" -clap = { version = "4.5.54", features = ["derive", "env"] } half = "2.7.1" -http-body-util = "0.1.3" # the CUDA kernels live in libqwen3_5_cuda.so, loaded at run time libloading = "0.8" memmap2 = "0.9.9" safetensors = "0.8.0" -serde_json = { version = "1.0.149", features = ["preserve_order"] } -sha2 = "0.10.9" +serde = "1" +serde_json = { version = "1.0.149", features = ["float_roundtrip", "preserve_order"] } # the onig regex backend, as in the Python tokenizers wheel tokenizers = { version = "=0.22.2", default-features = false, features = ["onig"] } tokio = { version = "1.49.0", features = ["macros", "net", "rt-multi-thread", "sync"] } diff --git a/src/models/cua_s1/native/README.md b/src/models/cua_s1/native/README.md deleted file mode 100644 index 1910b0dd..00000000 --- a/src/models/cua_s1/native/README.md +++ /dev/null @@ -1,38 +0,0 @@ -# Cua-S1 4B 0.2 native text worker - -A `/v1/systemone` worker for the `text` adapter in Rust. It handles requests like the reference worker (same validation, status codes, prompt token ids and answer format) and runs the Qwen3.5-4B forward pass on the CUDA kernels in [`src/backends/cuda/qwen3_5/`](../../../backends/cuda/qwen3_5/), with no Python and no PyTorch. Setup and launch are in [`recipe/cua_s1/native.md`](../../../../recipe/cua_s1/native.md). - -| File | Contents | -| --- | --- | -| `src/pyjson.rs` | JSON parsing and `json.dumps` output as the Python worker has them, since structured `state`, `instructions` and `criteria` values reach the prompt as `json.dumps` text. | -| `src/contract.rs` | The request mapping, prompt text, confidence and answers of `../text/contract.py`. | -| `src/server.rs` | The HTTP worker (`GET /health`, `POST /v1/systemone`) with the reference worker's limits and status codes. | -| `src/engine.rs` | Tokenization, the letter rows of the output projection, and scoring. | -| `src/model.rs` | The Qwen3.5 text model: weights, buffers, the layer loop, CUDA graphs and GEMM tuning. | -| `src/cuda.rs` | Loading `libqwen3_5_cuda.so` and the calls into it. | -| `tests/kernels.rs` | GPU checks of the attention and Gated DeltaNet kernels and of the GEMM plans. | - -## How a question is answered - -1. The request is parsed and mapped as in the reference worker, and each question's prompt is tokenized with the `tokenizer.json` exported next to the merged weights (the worker checks its SHA-256 against `cua_s1_export.json`). -2. One forward pass runs over the prompt: bfloat16 weights, with the `text` adapter merged into them by `recipe/cua_s1/export_text_merged.py`. The operations follow the Transformers implementation and round to bfloat16 where it does, except inside attention and the Gated DeltaNet prefill, which keep some intermediate results in bfloat16 as FlashAttention and flash-linear-attention do. -3. The final-norm hidden state at the last position is multiplied by the 26 letter rows of the output projection (float32 with float64 accumulation), and a softmax over the question's letters gives the option probabilities. - -## CUDA library, graphs and GEMM plans - -`libqwen3_5_cuda.so` is built by `src/backends/cuda/qwen3_5/build.sh` and loaded when the worker starts, so building the crate needs no CUDA toolkit. - -Prompts up to `--graph-max-tokens` (2048) run as a CUDA graph captured for their exact length on first use, and the 128 most recently used lengths keep theirs. A graph queues the same kernels with the same GEMM algorithms as an eager pass, so both give bitwise identical results. Longer prompts run eagerly. - -Projections that share an input run as one GEMM. GEMM algorithms are tuned at startup for a set of prompt lengths by timing cuBLASLt's candidates; `--gemm-search` times far more configurations for prompts up to `--graph-max-tokens`. `--gemm-plans` keeps the choices in a file, so later starts reuse them and give the same results. The file records the GPU, the cuBLASLt version, the workspace size and `--graph-max-tokens`; a file that does not match is refused, and removing it tunes again. - -## Differences from the reference worker - -- The probabilities are not bitwise identical: the adapter is merged, the kernels differ, and the GEMM algorithms depend on the prompt length and the GPU. They are held to the tolerance in [`../README.md`](../README.md#validation). -- Error messages quote names as Python's `repr` does, except that non-printable characters outside ASCII, such as U+00A0 or U+200B, are written as they are instead of escaped. -- The HTTP stacks differ outside the request body (Starlette and uvicorn there, axum and hyper here): trailing slashes, percent-encoded paths, `HEAD /health`, FastAPI's `/docs`, and which malformed HTTP requests are refused. -- `GET /health` also reports `"mode": "native"`, and its `dtype` is always `bfloat16`. - -## Tests - -`cargo test -p omni-cua-s1-native` runs the request-handling tests. The kernel checks in `tests/kernels.rs` need a GPU and run with `-- --ignored` and `CUA_S1_CUDA_LIB` pointing to a built `libqwen3_5_cuda.so`. diff --git a/src/models/cua_s1/native/src/contract.rs b/src/models/cua_s1/native/src/contract.rs index a117f6e1..0ca10f29 100644 --- a/src/models/cua_s1/native/src/contract.rs +++ b/src/models/cua_s1/native/src/contract.rs @@ -1,25 +1,18 @@ -//! Request mapping, prompt construction and answers for Cua-S1 4B 0.2, ported from -//! `src/models/cua_s1/text/contract.py`, so the two workers build the same prompts and -//! reject the same requests with the same status codes. +//! Request mapping, prompts and answers for Cua-S1 4B 0.2, as in +//! `src/models/cua_s1/README.md` and the reference worker's `text/contract.py`. -use std::fmt::Write as _; +use serde_json::{Map, Value, json}; -use crate::pyjson::{self, PyStr, Value, dumps, float_repr, repr, repr_str, write_json_str}; +use crate::json::{dumps, parse, quote}; pub const MODEL_NAME: &str = "cua-s1-4b-0.2"; -pub const ADAPTER_REPO: &str = "cua-ai/cua-s1-4b-0.2"; -pub const ADAPTER_REVISION: &str = "16818868b0cc7813808aae4e87b417657046ab79"; -pub const BASE_REPO: &str = "Qwen/Qwen3.5-4B"; -pub const BASE_REVISION: &str = "851bf6e806efd8d0a36b00ddf55e13ccb7b8cd0a"; - +pub const MODEL_ID: &str = "cua-ai/cua-s1-4b-0.2@16818868b0cc7813808aae4e87b417657046ab79:text"; pub const LETTERS: &str = "ABCDEFGHIJKLMNOPQRSTUVWXYZ"; -pub const MAX_OPTIONS: usize = 26; +pub const MAX_QUESTIONS: usize = 64; -// The system message, the user message layout and the fixed values below are copied -// from trycua/cua at 0e75660ce4c2edda519e0c795fa3ad98abf4e76f: -// `libs/cua-s1/python/src/cua_s1/four_b.py` (SYSTEM_PROMPT, build_prompt, -// _describe_option) and `libs/cua-driver/examples/jev-use/python/decision_models.py` -// (S1DecisionModel.score). +// The system message and the user message layout are copied from trycua/cua at +// 0e75660ce4c2edda519e0c795fa3ad98abf4e76f (`libs/cua-s1/python/src/cua_s1/four_b.py` +// and `libs/cua-driver/examples/jev-use/python/decision_models.py`). // // MIT License // @@ -47,39 +40,19 @@ current state of a screen and a fixed, closed list of candidate \ (element, action) options, each given a single letter. Choose exactly \ one option: the single best next action to take. Answer with ONLY that \ option's letter -- no words, no punctuation, no explanation."; -pub const APP: &str = "Cua Driver"; -pub const TASK_FAMILY: &str = "closed-candidate decision"; -pub const ROLE: &str = "Decision"; -pub const ACTION: &str = "select"; -/// A request the worker rejects, with the HTTP status to return. -#[derive(Debug, Clone, PartialEq)] pub struct RequestError { pub status: u16, pub message: String, } -impl RequestError { - pub fn new(status: u16, message: impl Into) -> Self { - Self { - status, - message: message.into(), - } - } - - fn unprocessable(message: impl Into) -> Self { - Self::new(422, message) +fn error(status: u16, message: impl Into) -> RequestError { + RequestError { + status, + message: message.into(), } } -impl From for RequestError { - fn from(e: pyjson::JsonError) -> Self { - RequestError::new(400, e.message()) - } -} - -/// One `choice` question mapped onto the prompt fields. -#[derive(Debug, Clone, PartialEq)] pub struct Question { pub name: String, pub goal: String, @@ -87,374 +60,193 @@ pub struct Question { pub labels: Vec, } -#[derive(Debug, Clone, PartialEq)] -pub struct Request { - pub state: String, - pub questions: Vec, -} - -pub fn parse_body(raw: &[u8]) -> Result, RequestError> { - Ok(pyjson::parse(raw)?) -} - -fn text(s: &PyStr) -> &str { - s.as_str().expect("parse rejects lone surrogates") +pub fn parse_body(raw: &[u8]) -> Result, RequestError> { + parse(raw).map_err(|message| error(400, message)) } -/// `state` or `instructions` as prompt text: a string as is, anything else as -/// `json.dumps(value, ensure_ascii=False)`. -fn as_text(value: &Value) -> String { +/// A string as is; an object or array as Python's `json.dumps` writes it. +fn text(value: &Value, place: &str) -> Result { match value { - Value::Str(s) => text(s).to_string(), - other => dumps(other), - } -} - -/// An option label escaped the way upstream's chooser does: -/// `json.dumps(value, ensure_ascii=False)[1:-1]`. -fn escape_label(value: &str) -> String { - let mut out = String::new(); - write_json_str(value, &mut out); - out[1..out.len() - 1].to_string() -} - -fn check_json_value(value: &Value, place: &str, allow_null: bool) -> Result<(), RequestError> { - match value { - Value::Null if !allow_null => Err(RequestError::unprocessable(format!( - "{place} must not be null" - ))), - Value::Bool(_) | Value::Int(_) | Value::Float(_) => Err(RequestError::unprocessable( + Value::String(s) => Ok(s.clone()), + Value::Object(_) | Value::Array(_) => Ok(dumps(value)), + _ => Err(error( + 422, format!("{place} must be a string, an object or an array"), )), - _ => Ok(()), } } -fn get<'a>(pairs: &'a [(PyStr, Value)], key: &str) -> Option<&'a Value> { - pairs.iter().find(|(k, _)| k == key).map(|(_, v)| v) -} - -/// Validate a `/v1/systemone` body and map it onto prompt fields. -pub fn map_request(body: &[(PyStr, Value)], max_questions: usize) -> Result { - match get(body, "model") { - Some(Value::Str(s)) if s == MODEL_NAME => {} - _ => { - return Err(RequestError::unprocessable(format!( - "'model' must be {}", - repr_str(&PyStr::new(MODEL_NAME)) - ))); - } +pub fn map_request(body: &Map) -> Result<(String, Vec), RequestError> { + if body.get("model").and_then(Value::as_str) != Some(MODEL_NAME) { + return Err(error(422, format!("'model' must be '{MODEL_NAME}'"))); } - - let state_value = - get(body, "state").ok_or_else(|| RequestError::unprocessable("'state' is required"))?; - check_json_value(state_value, "'state'", false)?; - let empty = match state_value { - Value::Str(s) => s == "", - Value::Object(p) => p.is_empty(), - Value::Array(a) => a.is_empty(), - _ => false, - }; - if empty { - return Err(RequestError::unprocessable("'state' must not be empty")); + let state = body.get("state").unwrap_or(&Value::Null); + if [json!(""), json!({}), json!([])].contains(state) { + return Err(error(422, "'state' must not be empty")); } - let state = as_text(state_value); - - let questions = match get(body, "questions") { + let state = text(state, "'state'")?; + let questions = match body.get("questions") { Some(Value::Object(q)) if !q.is_empty() => q, - _ => { - return Err(RequestError::unprocessable( - "'questions' must be a non-empty object", - )); - } + _ => return Err(error(422, "'questions' must be a non-empty object")), }; - if questions.len() > max_questions { - return Err(RequestError::new( - 413, - format!("too many questions ({} > {max_questions})", questions.len()), - )); + if questions.len() > MAX_QUESTIONS { + return Err(error(413, format!("more than {MAX_QUESTIONS} questions"))); } - - // Every question type is checked before the per-question checks, so a `score` - // or `noul` question anywhere rejects the whole request with that reason. - for (name, question) in questions { - let Value::Object(q) = question else { - return Err(RequestError::unprocessable(format!( - "question {} must be an object", - repr_str(name) - ))); + let mut mapped = Vec::with_capacity(questions.len()); + for (name, q) in questions { + let place = format!("question {}", quote(name)); + let Value::Object(q) = q else { + return Err(error(422, format!("{place} must be an object"))); }; - let missing = Value::Null; - let kind = get(q, "type").unwrap_or(&missing); - match kind { - Value::Str(s) if s == "score" || s == "noul" => { - return Err(RequestError::unprocessable(format!( - "question {}: type {} is not supported; Cua-S1 4B 0.2 answers 'choice' questions only", - repr_str(name), - repr(kind) - ))); - } - Value::Str(s) if s == "choice" => {} - _ => { - return Err(RequestError::unprocessable(format!( - "question {}: unknown type {}", - repr_str(name), - repr(kind) - ))); + match q.get("type").unwrap_or(&Value::Null) { + Value::String(t) if t == "choice" => {} + Value::String(t) if t == "score" || t == "noul" => { + return Err(error(422, format!("{place}: type '{t}' is not supported"))); } + other => return Err(error(422, format!("{place}: unknown type {other}"))), } - } - - let mut mapped = Vec::with_capacity(questions.len()); - for (name, question) in questions { - let Value::Object(q) = question else { - unreachable!("checked above") - }; - let place = format!("question {}", repr_str(name)); - let instructions = get(q, "instructions").ok_or_else(|| { - RequestError::unprocessable(format!("{place}: 'instructions' is required")) - })?; - check_json_value(instructions, &format!("{place}: 'instructions'"), true)?; - let goal = match instructions { - Value::Null => String::new(), - other => as_text(other), + let goal = match q.get("instructions") { + None => return Err(error(422, format!("{place}: 'instructions' is required"))), + Some(Value::Null) => String::new(), + Some(value) => text(value, &place)?, }; - - let criteria = match get(q, "criteria") { - Some(Value::Object(c)) => c, + let criteria = match q.get("criteria") { + Some(Value::Object(c)) if (1..=LETTERS.len()).contains(&c.len()) => c, _ => { - return Err(RequestError::unprocessable(format!( - "{place}: 'criteria' must be an object" - ))); + let message = format!("{place}: 'criteria' must be an object with 1 to 26 options"); + return Err(error(422, message)); } }; - if criteria.is_empty() { - return Err(RequestError::unprocessable(format!( - "{place}: 'criteria' must have at least one option" - ))); - } - if criteria.len() > MAX_OPTIONS { - return Err(RequestError::unprocessable(format!( - "{place}: {} options; at most {MAX_OPTIONS} are supported", - criteria.len() - ))); - } - let mut keys = Vec::with_capacity(criteria.len()); let mut labels = Vec::with_capacity(criteria.len()); for (key, value) in criteria { - check_json_value(value, &format!("{place}: option {}", repr_str(key)), true)?; let label = match value { - Value::Null => text(key).to_string(), - other => as_text(other), + Value::Null => key.clone(), + value => text(value, &format!("{place}: {}", quote(key)))?, }; - keys.push(text(key).to_string()); - labels.push(escape_label(&label)); + // escaped as `json.dumps(label, ensure_ascii=False)[1:-1]`, like upstream's chooser + let quoted = quote(&label); + labels.push(quoted[1..quoted.len() - 1].to_string()); } + let keys = criteria.keys().cloned().collect(); mapped.push(Question { - name: text(name).to_string(), + name: name.clone(), goal, keys, labels, }); } - Ok(Request { - state, - questions: mapped, - }) + Ok((state, mapped)) } -/// The user message for one question, matching upstream `build_prompt` (text). -pub fn user_message(state: &str, question: &Question) -> String { - let mut user = String::new(); - if !question.goal.is_empty() { - write!(user, "Goal: {}\n\n", question.goal).unwrap(); - } - write!(user, "App: {APP}\nTask family: {TASK_FAMILY}\n\n").unwrap(); - write!(user, "Accessibility tree:\n{state}\n\n").unwrap(); - user.push_str("Options:\n"); - for (i, (letter, label)) in LETTERS.chars().zip(&question.labels).enumerate() { - if i > 0 { - user.push('\n'); - } - write!(user, "{letter}. {ROLE} \"{label}\" -> {ACTION}").unwrap(); - } - user.push_str("\n\nAnswer with a single letter."); - user -} - -/// The prompt text the Qwen3.5 chat template renders for the system and user messages -/// with `add_generation_prompt=True` (thinking left on). The template trims message -/// content, which changes nothing here: the system prompt is fixed, and the user -/// message starts with "Goal: " or "App: " and ends with "letter.". +/// The prompt the Qwen3.5 chat template renders for the system and user messages with +/// `add_generation_prompt=True`. The template trims message content, which changes +/// nothing here: the user message starts with "Goal: " or "App: " and ends with "letter.". pub fn chat_text(state: &str, question: &Question) -> String { + let goal = match question.goal.as_str() { + "" => String::new(), + goal => format!("Goal: {goal}\n\n"), + }; + let options: Vec = LETTERS + .chars() + .zip(&question.labels) + .map(|(letter, label)| format!("{letter}. Decision \"{label}\" -> select")) + .collect(); format!( - "<|im_start|>system\n{SYSTEM_PROMPT}<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n\n", - user_message(state, question) + "<|im_start|>system\n{SYSTEM_PROMPT}<|im_end|>\n<|im_start|>user\n{goal}App: Cua Driver\n\ + Task family: closed-candidate decision\n\nAccessibility tree:\n{state}\n\nOptions:\n{}\n\n\ + Answer with a single letter.<|im_end|>\n<|im_start|>assistant\n\n", + options.join("\n") ) } -/// CPython 3.12's `sum()` over floats (Neumaier compensated summation). -fn py_sum(values: impl IntoIterator) -> f64 { - let (mut total, mut c) = (0.0f64, 0.0f64); - for x in values { - let t = total + x; - if total.abs() >= x.abs() { - c += (total - t) + x; - } else { - c += (x - t) + total; - } - total = t; - } - if c != 0.0 && c.is_finite() { - total += c; - } - total -} - -/// Normalized entropy, `1 - H(p) / ln(n)`, as the LAYA worker reports it. -pub fn confidence(probabilities: &[f64]) -> f64 { - let n = probabilities.len(); - if n < 2 { - return 1.0; - } - let entropy = -py_sum(probabilities.iter().map(|&p| p * p.clamp(1e-12, 1.0).ln())); - (1.0 - entropy / (n as f64).ln()).clamp(0.0, 1.0) -} - -/// One choice answer, already serialized the way the Python worker's response is -/// (`json.dumps(..., ensure_ascii=False, separators=(",", ":"))`). Ties go to the -/// earliest option. -pub fn answer_json(question: &Question, probabilities: &[f32]) -> Result { +/// The Jev choice answer; ties go to the earliest option. `confidence` is the +/// normalized entropy `1 - H(p) / ln(n)`, as the LAYA worker reports it. +pub fn answer(question: &Question, probabilities: &[f32]) -> Value { let p: Vec = probabilities.iter().map(|&x| x as f64).collect(); - if p.len() != question.keys.len() || !p.iter().all(|x| x.is_finite() && (0.0..=1.0).contains(x)) - { - return Err(format!("model returned invalid probabilities: {p:?}")); - } - let total = py_sum(p.iter().copied()); - // math.isclose(total, 1.0, abs_tol=1e-5) - if (total - 1.0).abs() > f64::max(1e-9 * total.abs().max(1.0), 1e-5) { - return Err(format!("model probabilities do not sum to one: {p:?}")); - } - let mut best = 0; - for (i, &x) in p.iter().enumerate() { - if x > p[best] { - best = i; - } - } - let mut out = String::from("{\"type\":\"choice\",\"choice\":"); - write_json_str(&question.keys[best], &mut out); - out.push_str(",\"probabilities\":{"); - for (i, (key, x)) in question.keys.iter().zip(&p).enumerate() { - if i > 0 { - out.push(','); - } - write_json_str(key, &mut out); - out.push(':'); - out.push_str(&float_repr(*x)); - } - out.push_str("},\"confidence\":"); - out.push_str(&float_repr(confidence(&p))); - out.push('}'); - Ok(out) -} - -pub fn model_identity(revision: &str) -> String { - format!("{ADAPTER_REPO}@{revision}:text") -} - -/// `{"detail": message}` as the Python worker serializes it. -pub fn detail_json(message: &str) -> String { - let mut out = String::from("{\"detail\":"); - write_json_str(message, &mut out); - out.push('}'); - out + let best = (0..p.len()).fold(0, |b, i| if p[i] > p[b] { i } else { b }); + let entropy: f64 = -p + .iter() + .filter(|&&x| x > 0.0) + .map(|&x| x * x.ln()) + .sum::(); + let n = p.len() as f64; + let probs: Map = question + .keys + .iter() + .cloned() + .zip(p.iter().map(|&x| json!(x))) + .collect(); + json!({ + "type": "choice", + "choice": question.keys[best], + "probabilities": probs, + "confidence": if p.len() > 1 { (1.0 - entropy / n.ln()).max(0.0) } else { 1.0 }, + }) } -pub const WARMUP_BODY: &str = r#"{"model": "cua-s1-4b-0.2", "state": "Dialog: 'Update installed.' Button: OK", "questions": {"warmup": {"type": "choice", "instructions": "Close the dialog.", "criteria": {"ok": "Click OK", "wait": "Wait"}}}}"#; - #[cfg(test)] mod tests { use super::*; - fn map(body: &str) -> Result { - map_request(&parse_body(body.as_bytes())?, 64) + fn map(body: &str) -> Result<(String, Vec), RequestError> { + map_request(&parse_body(body.as_bytes())?) } - fn detail(body: &str) -> (u16, String) { - let e = map(body).unwrap_err(); - (e.status, e.message) - } - - const OK: &str = r#"{"model": "cua-s1-4b-0.2", "state": "S", "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "A", "b": null}}}}"#; - #[test] fn maps_a_request() { - let r = map(OK).unwrap(); - assert_eq!(r.state, "S"); - assert_eq!(r.questions[0].keys, ["a", "b"]); - assert_eq!(r.questions[0].labels, ["A", "b"]); - let text = chat_text(&r.state, &r.questions[0]); + let body = r#"{"model": "cua-s1-4b-0.2", "state": {"a": [1.5, "é"]}, "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "Say \"hi\"\n", "b": null}}}}"#; + let (state, questions) = map(body).ok().unwrap(); + assert_eq!(state, r#"{"a": [1.5, "é"]}"#); + assert_eq!(questions[0].keys, ["a", "b"]); + assert_eq!(questions[0].labels, [r#"Say \"hi\"\n"#, "b"]); + let text = chat_text(&state, &questions[0]); + assert!(text.contains("<|im_start|>user\nGoal: go\n\nApp: Cua Driver\n")); assert!(text.ends_with( - "Options:\nA. Decision \"A\" -> select\nB. Decision \"b\" -> select\n\nAnswer with a single letter.<|im_end|>\n<|im_start|>assistant\n\n" + "Options:\nA. Decision \"Say \\\"hi\\\"\\n\" -> select\nB. Decision \"b\" -> select\n\n\ + Answer with a single letter.<|im_end|>\n<|im_start|>assistant\n\n" )); - assert!(text.contains("<|im_start|>user\nGoal: go\n\nApp: Cua Driver\n")); } #[test] - fn errors_match_python() { - let cases: &[(&str, u16, &str)] = &[ - (r#"{"state": "S"}"#, 422, "'model' must be 'cua-s1-4b-0.2'"), - (r#"{"model": "cua-s1-4b-0.2"}"#, 422, "'state' is required"), - ( - r#"{"model": "cua-s1-4b-0.2", "state": null}"#, - 422, - "'state' must not be null", - ), - ( - r#"{"model": "cua-s1-4b-0.2", "state": 1.5}"#, - 422, - "'state' must be a string, an object or an array", - ), - ( - r#"{"model": "cua-s1-4b-0.2", "state": {}}"#, - 422, - "'state' must not be empty", - ), - ( - r#"{"model": "cua-s1-4b-0.2", "state": "S", "questions": []}"#, - 422, - "'questions' must be a non-empty object", - ), + fn errors() { + let q = |question: &str| { + format!( + r#"{{"model": "cua-s1-4b-0.2", "state": "S", "questions": {{"q": {question}}}}}"# + ) + }; + let cases = [ + (r#"{"state": "S"}"#.to_string(), 422), ( - r#"{"model": "cua-s1-4b-0.2", "state": "S", "questions": {"q": {"type": "choice"}, "r": {"type": "noul"}}}"#, + r#"{"model": "cua-s1-4b-0.2", "state": {}}"#.to_string(), 422, - "question 'r': type 'noul' is not supported; Cua-S1 4B 0.2 answers 'choice' questions only", ), ( - r#"{"model": "cua-s1-4b-0.2", "state": "S", "questions": {"it's": {"type": [1, {"a": null}]}}}"#, + r#"{"model": "cua-s1-4b-0.2", "state": 1.5}"#.to_string(), 422, - "question \"it's\": unknown type [1, {'a': None}]", ), + (q(r#"{"type": "noul"}"#), 422), + (q(r#"{"type": "choice"}"#), 422), ( - r#"{"model": "cua-s1-4b-0.2", "state": "S", "questions": {"q": {"type": "choice"}}}"#, + q(r#"{"type": "choice", "instructions": null, "criteria": {"a": true}}"#), 422, - "question 'q': 'instructions' is required", ), ( - r#"{"model": "cua-s1-4b-0.2", "state": "S", "questions": {"q": {"type": "choice", "instructions": null, "criteria": {"a": true}}}}"#, + q(r#"{"type": "choice", "instructions": null, "criteria": {}}"#), 422, - "question 'q': option 'a' must be a string, an object or an array", ), + (r#"{"a": 1, "a": 2}"#.to_string(), 400), ]; - for (body, status, message) in cases { - assert_eq!(detail(body), (*status, message.to_string()), "{body}"); + for (body, status) in cases { + assert_eq!(map(&body).err().map(|e| e.status), Some(status), "{body}"); } } #[test] - fn confidence_matches_python() { - // values from the Python worker - let p = [0.00247262348420918f64, 0.9975274205207825]; - assert_eq!(float_repr(confidence(&p)), "0.9750249548256825"); + fn answer_and_confidence() { + let (_, questions) = map(r#"{"model": "cua-s1-4b-0.2", "state": "S", "questions": {"q": {"type": "choice", "instructions": null, "criteria": {"a": "A", "b": "B"}}}}"#).ok().unwrap(); + let tie = answer(&questions[0], &[0.5, 0.5]); + assert_eq!(tie["choice"], "a"); + assert!(tie["confidence"].as_f64().unwrap().abs() < 1e-12); + assert_eq!(answer(&questions[0], &[0.25, 0.75])["choice"], "b"); } } diff --git a/src/models/cua_s1/native/src/cuda.rs b/src/models/cua_s1/native/src/cuda.rs index c33c7af9..a15fef2f 100644 --- a/src/models/cua_s1/native/src/cuda.rs +++ b/src/models/cua_s1/native/src/cuda.rs @@ -21,18 +21,6 @@ pub struct Stream(*mut c_void); // from one thread at a time. unsafe impl Send for Stream {} -/// A tuned GEMM algorithm for one shape (`Cs1GemmPlan` in ops.h). -#[repr(C)] -#[derive(Clone, Copy, Default)] -pub struct GemmPlan { - pub m: i32, - pub n: i32, - pub k: i32, - pub ldy: i32, - pub cublaslt_version: u64, - pub algo: [u64; 8], -} - macro_rules! api { ($($name:ident($($arg:ident: $ty:ty),* $(,)?) $(-> $ret:ty)?;)*) => { /// The functions of the library, as declared in ops.h. @@ -63,17 +51,12 @@ api! { cs1_abi_version() -> u32; cs1_error_string(code: c_int) -> *const c_char; cs1_set_device(device: c_int) -> c_int; - cs1_device_info(name: *mut c_char, cap: usize, compute_capability: *mut c_int, sms: *mut c_int) -> c_int; cs1_malloc(ptr: *mut *mut c_void, bytes: usize) -> c_int; cs1_free(ptr: *mut c_void) -> c_int; cs1_stream_create(stream: *mut Stream) -> c_int; cs1_stream_sync(stream: Stream) -> c_int; cs1_upload(dst: *mut c_void, src: *const c_void, bytes: usize, stream: Stream) -> c_int; cs1_download(dst: *mut c_void, src: *const c_void, bytes: usize, stream: Stream) -> c_int; - cs1_graph_begin(stream: Stream) -> c_int; - cs1_graph_end(stream: Stream, exec: *mut *mut c_void) -> c_int; - cs1_graph_launch(exec: *mut c_void, stream: Stream) -> c_int; - cs1_graph_destroy(exec: *mut c_void) -> c_int; cs1_embed(ids: *const i32, table: *const c_void, out: *mut c_void, t: c_int, d: c_int, stream: Stream) -> c_int; cs1_rms_norm( x: *const c_void, w: *const c_void, out: *mut c_void, rows: c_int, d: c_int, eps: f32, stream: Stream, @@ -108,22 +91,10 @@ api! { q: *const c_void, k: *const c_void, v: *const c_void, ldv: c_int, out: *mut c_void, t: c_int, hq: c_int, hk: c_int, dh: c_int, scale: f32, stream: Stream, ) -> c_int; - cs1_attention_simple( - q: *const c_void, k: *const c_void, v: *const c_void, ldv: c_int, out: *mut c_void, t: c_int, hq: c_int, - hk: c_int, dh: c_int, scale: f32, stream: Stream, - ) -> c_int; cs1_sigmoid_gate(x: *mut c_void, gate: *const c_void, n: usize, stream: Stream) -> c_int; cs1_silu_mul(gate_up: *const c_void, ld: c_int, out: *mut c_void, t: c_int, i: c_int, stream: Stream) -> c_int; cs1_gemm_create(workspace_bytes: usize) -> *mut c_void; cs1_gemm_destroy(gemm: *mut c_void); - cs1_gemm_tune( - gemm: *mut c_void, x: *const c_void, w: *const c_void, y: *mut c_void, m: c_int, n: c_int, k: c_int, - ldy: c_int, exhaustive: c_int, stream: Stream, - ) -> c_int; - cs1_gemm_tune_done(gemm: *mut c_void); - cs1_gemm_export(gemm: *mut c_void, out: *mut GemmPlan, cap: usize) -> usize; - cs1_gemm_import(gemm: *mut c_void, plans: *const GemmPlan, n: usize) -> c_int; - cs1_gemm_version() -> usize; cs1_gemm( gemm: *mut c_void, x: *const c_void, w: *const c_void, y: *mut c_void, m: c_int, n: c_int, k: c_int, ldy: c_int, stream: Stream, @@ -165,21 +136,6 @@ pub fn load(path: &Path) -> Result<&'static Api> { Ok(API.get_or_init(|| api)) } -/// Name, compute capability (major * 10 + minor) and SM count of the current device. -pub fn device_info() -> Result<(String, i32, i32)> { - let mut name = [0 as c_char; 256]; - let (mut cc, mut sms) = (0, 0); - // SAFETY: `name` has room for 256 bytes, of which the library writes a terminated - // string; the two integers are valid for writes. - check( - unsafe { (api().cs1_device_info)(name.as_mut_ptr(), name.len(), &mut cc, &mut sms) }, - "reading the device properties", - )?; - // SAFETY: terminated by the library. - let name = unsafe { CStr::from_ptr(name.as_ptr()) }; - Ok((name.to_string_lossy().into_owned(), cc, sms)) -} - /// The loaded library; `load` must have succeeded before. pub fn api() -> &'static Api { API.get().expect("the CUDA library is not loaded") @@ -281,52 +237,3 @@ pub unsafe fn download(dst: &mut [u8], src: *const c_void, stream: Stream) -> Re "copy to host", ) } - -/// An instantiated CUDA graph, destroyed on drop. -pub struct Graph { - exec: *mut c_void, -} - -// SAFETY: the executable graph is only launched by its owner, one launch at a time. -unsafe impl Send for Graph {} - -impl Graph { - /// Capture the work `record` queues on `stream` (nothing runs) and instantiate it. - pub fn capture(stream: Stream, record: impl FnOnce() -> Result<()>) -> Result { - // SAFETY: plain runtime calls on a stream from new_stream; the capture is - // always ended, also when `record` fails. - unsafe { - check((api().cs1_graph_begin)(stream), "cudaStreamBeginCapture")?; - let recorded = record(); - let mut exec = std::ptr::null_mut(); - let ended = check( - (api().cs1_graph_end)(stream, &mut exec), - "capturing a CUDA graph", - ); - match recorded.and(ended) { - Ok(()) => Ok(Graph { exec }), - Err(e) => { - if !exec.is_null() { - (api().cs1_graph_destroy)(exec); - } - Err(e) - } - } - } - } - - pub fn launch(&self, stream: Stream) -> Result<()> { - // SAFETY: an instantiated graph whose buffers outlive it (see Model). - check( - unsafe { (api().cs1_graph_launch)(self.exec, stream) }, - "cudaGraphLaunch", - ) - } -} - -impl Drop for Graph { - fn drop(&mut self) { - // SAFETY: instantiated by capture and not destroyed before. - unsafe { (api().cs1_graph_destroy)(self.exec) }; - } -} diff --git a/src/models/cua_s1/native/src/engine.rs b/src/models/cua_s1/native/src/engine.rs index d2b10814..859156da 100644 --- a/src/models/cua_s1/native/src/engine.rs +++ b/src/models/cua_s1/native/src/engine.rs @@ -1,217 +1,60 @@ -//! The model side: prompt tokenization, and one prefill-only forward pass per -//! question through the native Qwen3.5 model, scored with the 26 letter rows of -//! the output projection. +//! Tokenization and scoring: one prefill-only forward pass per question through the +//! native Qwen3.5 model, scored with the 26 letter rows of the output projection. use std::path::Path; use std::sync::{Arc, Mutex}; -use std::time::Instant; -use anyhow::{Context, Result, bail, ensure}; +use anyhow::{Context, Result, ensure}; use serde_json::Value as Json; -use sha2::Digest; use tokenizers::Tokenizer; -use crate::contract::{self, LETTERS, Question}; +use crate::contract::{LETTERS, Question, chat_text}; use crate::model::Model; -use crate::model::Options; -/// What `cua_s1_export.json` records about a merged checkpoint. -#[derive(Debug, Clone)] -pub struct Provenance { - pub base_revision: String, - pub adapter_revision: String, -} - -/// Read and check the export record written next to the merged weights. -pub fn provenance(dir: &Path) -> Result { - let path = dir.join("cua_s1_export.json"); - let info: Json = serde_json::from_str(&std::fs::read_to_string(&path).with_context(|| { - format!( - "{} is missing; export the merged checkpoint first", - path.display() - ) - })?)?; - ensure!( - info["format"] == "cua-s1-text-merged/1", - "{}: unknown format {}", - path.display(), - info["format"] - ); - ensure!( - info["base"]["repo"] == contract::BASE_REPO, - "base is not {}", - contract::BASE_REPO - ); - ensure!( - info["adapter"]["repo"] == contract::ADAPTER_REPO && info["adapter"]["subfolder"] == "text", - "adapter is not the `text` adapter of {}", - contract::ADAPTER_REPO - ); - // Transformers 5.17 tokenizes Qwen3.5 with the rule saved in this file, not the - // one in the base repo's tokenizer.json, so the file must be the exported one. - let want = info["tokenizer"]["sha256"] - .as_str() - .context("cua_s1_export.json does not record the tokenizer's sha256")?; - let got = format!( - "{:x}", - sha2::Sha256::digest(std::fs::read(dir.join("tokenizer.json"))?) - ); - ensure!( - got == want, - "tokenizer.json (sha256 {got}) is not the one exported with the weights ({want})" - ); - let rev = |v: &Json| v.as_str().map(str::to_string).context("revision missing"); - Ok(Provenance { - base_revision: rev(&info["base"]["revision"])?, - adapter_revision: rev(&info["adapter"]["revision"])?, - }) -} - -/// Chat text and token ids for a question. -pub struct Prompter { - tokenizer: Tokenizer, - pub letter_ids: Vec, -} - -impl Prompter { - fn load(dir: &Path) -> Result { - let tokenizer = Tokenizer::from_file(dir.join("tokenizer.json")) - .map_err(|e| anyhow::anyhow!("tokenizer.json: {e}"))?; - let mut letter_ids = Vec::with_capacity(LETTERS.len()); - for letter in LETTERS.chars() { - let enc = tokenizer - .encode(letter.to_string(), false) - .map_err(|e| anyhow::anyhow!(e))?; - ensure!( - enc.get_ids().len() == 1, - "letter {letter} is not a single token" - ); - letter_ids.push(enc.get_ids()[0]); - } - Ok(Self { - tokenizer, - letter_ids, - }) - } - - pub fn encode(&self, state: &str, question: &Question) -> Result> { - let text = contract::chat_text(state, question); - let enc = self - .tokenizer - .encode(text, false) - .map_err(|e| anyhow::anyhow!(e))?; - Ok(enc.get_ids().to_vec()) - } -} - -/// Where the output projection can live in a Qwen3.5 text checkpoint; with tied -/// weights (as in Qwen3.5-4B) only the embedding is stored. -const HEAD_NAMES: &[&str] = &[ - "lm_head.weight", - "language_model.lm_head.weight", - "model.embed_tokens.weight", - "model.language_model.embed_tokens.weight", - "language_model.model.embed_tokens.weight", -]; - -/// The letter rows of the output projection, as float32, read straight from the -/// safetensors files. -fn letter_rows(dir: &Path, letter_ids: &[u32]) -> Result<(Vec, usize)> { - let index_path = dir.join("model.safetensors.index.json"); - let (file, name) = if index_path.exists() { - let index: Json = serde_json::from_str(&std::fs::read_to_string(&index_path)?)?; - let map = &index["weight_map"]; - let name = HEAD_NAMES - .iter() - .find(|n| map.get(**n).is_some()) - .with_context(|| { - format!( - "no output projection or embedding in {}", - index_path.display() - ) - })?; - let file = map[*name] - .as_str() - .context("weight_map entry is not a file name")?; - (dir.join(file), Some(name.to_string())) - } else { - (dir.join("model.safetensors"), None) - }; - let file = std::fs::File::open(&file)?; - // SAFETY: the checkpoint is not modified while the worker runs. - let mmap = unsafe { memmap2::Mmap::map(&file)? }; - let st = safetensors::SafeTensors::deserialize(&mmap)?; - let name = match name { - Some(n) => n, - None => { - let names = st.names(); - HEAD_NAMES - .iter() - .find(|n| names.iter().any(|m| m == *n)) - .context("no output projection or embedding in model.safetensors")? - .to_string() - } - }; - let view = st.tensor(&name)?; - ensure!( - view.dtype() == safetensors::Dtype::BF16 && view.shape().len() == 2, - "{name}: expected a 2-D bfloat16 tensor, got {:?} {:?}", - view.dtype(), - view.shape() - ); - let (vocab, hidden) = (view.shape()[0], view.shape()[1]); - let data = view.data(); - let mut rows = Vec::with_capacity(letter_ids.len() * hidden); - for &id in letter_ids { - let id = id as usize; - ensure!(id < vocab, "letter id {id} outside the vocabulary"); - let row = &data[id * hidden * 2..(id + 1) * hidden * 2]; - let (pairs, _) = row.as_chunks::<2>(); - rows.extend(pairs.iter().map(|&b| half::bf16::from_le_bytes(b).to_f32())); - } - Ok((rows, hidden)) -} +/// Tied with the output projection in Qwen3.5-4B. +const EMBEDDING: &str = "model.language_model.embed_tokens.weight"; pub struct Engine { - pub prompter: Prompter, + tokenizer: Tokenizer, model: Arc>, + /// the letter rows of the output projection, as float32 letters: Vec, - hidden: usize, - pub load_seconds: f64, - pub device: String, - pub provenance: Provenance, } impl Engine { - /// Check the export record and `tokenizer.json` (see `provenance`), load the CUDA - /// library and the model, and prepare CUDA graphs (see `Options`). - pub async fn load(dir: &Path, opts: &Options) -> Result { - let started = Instant::now(); - let provenance = provenance(dir)?; - let prompter = Prompter::load(dir)?; - let (letters, hidden) = letter_rows(dir, &prompter.letter_ids)?; - let dir = dir.to_path_buf(); - let opts = opts.clone(); - let model = tokio::task::spawn_blocking(move || Model::load(&dir, &opts)).await??; + pub async fn load(dir: &Path, library: &Path) -> Result { + // Written by export_text_merged.py; without it `dir` may hold the base model alone. ensure!( - model.cfg.hidden == hidden, - "hidden size {} does not match the head ({hidden})", - model.cfg.hidden + dir.join("cua_s1_export.json").exists(), + "{} is not a merged text checkpoint; see recipe/cua_s1/native.md", + dir.display() ); + let tokenizer = + Tokenizer::from_file(dir.join("tokenizer.json")).map_err(anyhow::Error::msg)?; + let ids = LETTERS + .chars() + .map(|c| { + tokenizer + .token_to_id(&c.to_string()) + .context("letter token") + }) + .collect::>>()?; + let (d, lib) = (dir.to_path_buf(), library.to_path_buf()); + let model = tokio::task::spawn_blocking(move || Model::load(&d, &lib)).await??; + let letters = letter_rows(dir, &ids, model.cfg.hidden)?; Ok(Self { - prompter, + tokenizer, model: Arc::new(Mutex::new(model)), letters, - hidden, - load_seconds: started.elapsed().as_secs_f64(), - device: "cuda".to_string(), - provenance, }) } - /// The longest prompt that runs as a CUDA graph (0: none). - pub fn graph_max_tokens(&self) -> usize { - self.model.lock().map(|m| m.graph_max_tokens()).unwrap_or(0) + pub fn encode(&self, state: &str, question: &Question) -> Result> { + let enc = self + .tokenizer + .encode(chat_text(state, question), false) + .map_err(anyhow::Error::msg)?; + Ok(enc.get_ids().to_vec()) } /// Option probabilities for one prompt: the final-norm hidden state at the last @@ -220,33 +63,57 @@ impl Engine { pub async fn score(&self, ids: Vec, n_options: usize) -> Result> { let model = self.model.clone(); let last = tokio::task::spawn_blocking(move || { - let mut model = model + model .lock() - .map_err(|_| anyhow::anyhow!("model lock poisoned"))?; - model.forward(&ids) + .map_err(|_| anyhow::anyhow!("poisoned"))? + .forward(&ids) }) .await??; - if last.len() != self.hidden { - bail!( - "hidden size {} does not match the head ({})", - last.len(), - self.hidden - ); - } - let logits: Vec = self + let logits: Vec = self .letters - .chunks_exact(self.hidden) + .chunks_exact(last.len()) .take(n_options) .map(|w| { w.iter() .zip(&last) .map(|(&a, &b)| a as f64 * b as f64) - .sum::() as f32 + .sum::() as f32 as f64 }) .collect(); - let max = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max) as f64; - let exps: Vec = logits.iter().map(|&l| (l as f64 - max).exp()).collect(); + let max = logits.iter().copied().fold(f64::NEG_INFINITY, f64::max); + let exps: Vec = logits.iter().map(|&l| (l - max).exp()).collect(); let total: f64 = exps.iter().sum(); - Ok(exps.iter().map(|e| (e / total) as f32).collect()) + let probs: Vec = exps.iter().map(|e| (e / total) as f32).collect(); + ensure!( + probs.iter().all(|p| p.is_finite()), + "non-finite probabilities" + ); + Ok(probs) + } +} + +/// The letter rows of the bfloat16 embedding, read from the safetensors files. +fn letter_rows(dir: &Path, ids: &[u32], hidden: usize) -> Result> { + let index: Json = serde_json::from_str(&std::fs::read_to_string( + dir.join("model.safetensors.index.json"), + )?)?; + let file = index["weight_map"][EMBEDDING].as_str().context(EMBEDDING)?; + let file = std::fs::File::open(dir.join(file))?; + // SAFETY: the checkpoint is not modified while the worker runs. + let mmap = unsafe { memmap2::Mmap::map(&file)? }; + let tensors = safetensors::SafeTensors::deserialize(&mmap)?; + let view = tensors.tensor(EMBEDDING)?; + ensure!( + view.dtype() == safetensors::Dtype::BF16 && view.shape()[1] == hidden, + "{EMBEDDING}: {:?} {:?}", + view.dtype(), + view.shape() + ); + let mut rows = Vec::with_capacity(ids.len() * hidden); + for &id in ids { + let row = &view.data()[id as usize * hidden * 2..(id as usize + 1) * hidden * 2]; + let (pairs, _) = row.as_chunks::<2>(); + rows.extend(pairs.iter().map(|&b| half::bf16::from_le_bytes(b).to_f32())); } + Ok(rows) } diff --git a/src/models/cua_s1/native/src/json.rs b/src/models/cua_s1/native/src/json.rs new file mode 100644 index 00000000..a0e2ed5a --- /dev/null +++ b/src/models/cua_s1/native/src/json.rs @@ -0,0 +1,245 @@ +//! Request JSON, on serde_json. +//! +//! - [`parse`] rejects what the contract rejects with a 400: invalid JSON or UTF-8, +//! `NaN`/`Infinity`, numbers out of range, lone surrogates and nesting deeper than +//! serde_json's limit (all serde_json errors), and keys repeated in an object. +//! - [`dumps`] writes `json.dumps(value, ensure_ascii=False)`: separators `, ` and +//! `: `, key order kept, floats as Python's `repr`. + +use std::fmt::Write as _; +use std::io; + +use serde::Serialize; +use serde::de::{self, Deserializer, MapAccess, SeqAccess, Visitor}; +use serde_json::{Map, Number, Value}; + +/// Decode a request body into its top-level object; the error is the 400 message. +pub fn parse(raw: &[u8]) -> Result, String> { + let mut de = serde_json::Deserializer::from_slice(raw); + let value = de + .deserialize_any(NoDuplicates) + .and_then(|v| de.end().map(|()| v)) + .map_err(|e| format!("request body is not valid JSON: {e}"))?; + match value { + Value::Object(map) => Ok(map), + _ => Err("request body must be a JSON object".into()), + } +} + +/// Builds a `Value` like serde_json does, but fails on a repeated key. +struct NoDuplicates; + +impl<'de> de::Deserialize<'de> for Wrapped { + fn deserialize>(d: D) -> Result { + d.deserialize_any(NoDuplicates).map(Wrapped) + } +} + +struct Wrapped(Value); + +impl<'de> Visitor<'de> for NoDuplicates { + type Value = Value; + + fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { + f.write_str("a JSON value") + } + fn visit_unit(self) -> Result { + Ok(Value::Null) + } + fn visit_bool(self, b: bool) -> Result { + Ok(Value::Bool(b)) + } + fn visit_i64(self, n: i64) -> Result { + Ok(Value::Number(n.into())) + } + fn visit_u64(self, n: u64) -> Result { + Ok(Value::Number(n.into())) + } + fn visit_f64(self, x: f64) -> Result { + Number::from_f64(x) + .map(Value::Number) + .ok_or_else(|| E::custom("number out of range")) + } + fn visit_str(self, s: &str) -> Result { + Ok(Value::String(s.to_owned())) + } + fn visit_string(self, s: String) -> Result { + Ok(Value::String(s)) + } + fn visit_seq>(self, mut seq: A) -> Result { + let mut items = Vec::new(); + while let Some(Wrapped(v)) = seq.next_element()? { + items.push(v); + } + Ok(Value::Array(items)) + } + fn visit_map>(self, mut map: A) -> Result { + let mut obj = Map::new(); + while let Some(key) = map.next_key::()? { + let Wrapped(v) = map.next_value()?; + if obj.contains_key(&key) { + return Err(de::Error::custom(format_args!( + "duplicate key {}", + quote(&key) + ))); + } + obj.insert(key, v); + } + Ok(Value::Object(obj)) + } +} + +/// A string as a JSON literal, which is also how error messages quote names. +pub fn quote(s: &str) -> String { + serde_json::to_string(s).expect("a string serializes") +} + +/// `json.dumps(value, ensure_ascii=False)`. +pub fn dumps(value: &Value) -> String { + let mut out = Vec::new(); + let mut ser = serde_json::Serializer::with_formatter(&mut out, PyFormatter); + value.serialize(&mut ser).expect("a Value serializes"); + String::from_utf8(out).expect("serde_json writes UTF-8") +} + +/// serde_json's compact output with Python's separators and float format; its string +/// escaping (`"`, `\\` and control characters, `\u00XX` in lowercase) is Python's. +struct PyFormatter; + +impl serde_json::ser::Formatter for PyFormatter { + fn begin_array_value( + &mut self, + w: &mut W, + first: bool, + ) -> io::Result<()> { + if first { Ok(()) } else { w.write_all(b", ") } + } + fn begin_object_key( + &mut self, + w: &mut W, + first: bool, + ) -> io::Result<()> { + if first { Ok(()) } else { w.write_all(b", ") } + } + fn begin_object_value(&mut self, w: &mut W) -> io::Result<()> { + w.write_all(b": ") + } + fn write_f64(&mut self, w: &mut W, x: f64) -> io::Result<()> { + w.write_all(float_repr(x).as_bytes()) + } +} + +/// Python's `repr(float)`: the shortest digits that round-trip, in fixed notation +/// for exponents from -5 to 15 and scientific notation otherwise. (Python breaks the +/// rare exact ties between two shortest candidates to even; this does not.) +pub fn float_repr(x: f64) -> String { + if x == 0.0 { + return if x.is_sign_negative() { "-0.0" } else { "0.0" }.into(); + } + let sci = format!("{x:e}"); + let (mantissa, exp) = sci.split_once('e').expect("{:e} has an exponent"); + let exp: i32 = exp.parse().expect("integer exponent"); + let (neg, mantissa) = match mantissa.strip_prefix('-') { + Some(m) => (true, m), + None => (false, mantissa), + }; + let digits: String = mantissa.chars().filter(|c| *c != '.').collect(); + let mut out = String::new(); + if neg { + out.push('-'); + } + let decpt = exp + 1; + if decpt <= -4 || decpt > 16 { + out.push_str(&digits[..1]); + if digits.len() > 1 { + out.push('.'); + out.push_str(&digits[1..]); + } + let sign = if exp < 0 { '-' } else { '+' }; + write!(out, "e{sign}{:02}", exp.unsigned_abs()).unwrap(); + } else if decpt <= 0 { + out.push_str("0."); + out.extend(std::iter::repeat_n('0', (-decpt) as usize)); + out.push_str(&digits); + } else if decpt as usize >= digits.len() { + out.push_str(&digits); + out.extend(std::iter::repeat_n('0', decpt as usize - digits.len())); + out.push_str(".0"); + } else { + out.push_str(&digits[..decpt as usize]); + out.push('.'); + out.push_str(&digits[decpt as usize..]); + } + out +} + +#[cfg(test)] +mod tests { + use super::*; + + fn err(body: &str) -> String { + parse(body.as_bytes()).unwrap_err() + } + + #[test] + fn float_repr_matches_python_examples() { + let cases = [ + (1.0, "1.0"), + (1e16, "1e+16"), + (1e15, "1000000000000000.0"), + (1e-5, "1e-05"), + (1e-4, "0.0001"), + (-0.0, "-0.0"), + (3.14e-07, "3.14e-07"), + (5e-324, "5e-324"), + (1.7976931348623157e308, "1.7976931348623157e+308"), + (5.960464477539063e-08, "5.960464477539063e-08"), + ]; + for (x, want) in cases { + assert_eq!(float_repr(x), want, "{x:e}"); + } + } + + #[test] + fn dumps_matches_python() { + let v = Value::Object(parse(br#"{"a": [1.0, 1e16, 1e-5, 0.0001, -0.0, 123456789012345678, 3.14e-07, true, null], "b": {}}"#).unwrap()); + assert_eq!( + dumps(&v), + r#"{"a": [1.0, 1e+16, 1e-05, 0.0001, -0.0, 123456789012345678, 3.14e-07, true, null], "b": {}}"# + ); + let s = Value::String("\u{0}\u{1f}\u{7f}\u{2028}\"\\/\t\u{8}\u{c}é😀".into()); + assert_eq!( + dumps(&s), + "\"\\u0000\\u001f\u{7f}\u{2028}\\\"\\\\/\\t\\b\\fé😀\"" + ); + } + + #[test] + fn rejects_what_the_contract_rejects() { + assert_eq!(err("[]"), "request body must be a JSON object"); + for body in [ + r#"{"a": NaN}"#, + r#"{"a": 1e400}"#, + r#"{"a": "\ud800x"}"#, + r#"{"a": 1, "b": 2, "a": 3}"#, + r#"{"a": [1,]}"#, + r#"{} x"#, + "\u{feff}{}", + ] { + assert!( + err(body).starts_with("request body is not valid JSON"), + "{body}" + ); + } + assert!(err(r#"{"a": 1, "a": 2}"#).contains("duplicate key \"a\"")); + assert!( + err(&format!( + "{{\"a\": {}1{}}}", + "[".repeat(200), + "]".repeat(200) + )) + .contains("recursion limit") + ); + assert!(parse(b"{\"a\": \"\xff\"}").is_err()); + } +} diff --git a/src/models/cua_s1/native/src/lib.rs b/src/models/cua_s1/native/src/lib.rs index 3be5ae17..0d3aa5fd 100644 --- a/src/models/cua_s1/native/src/lib.rs +++ b/src/models/cua_s1/native/src/lib.rs @@ -5,6 +5,5 @@ pub mod contract; pub mod cuda; pub mod engine; +pub mod json; pub mod model; -pub mod pyjson; -pub mod server; diff --git a/src/models/cua_s1/native/src/main.rs b/src/models/cua_s1/native/src/main.rs index 3ade9523..121a9e3d 100644 --- a/src/models/cua_s1/native/src/main.rs +++ b/src/models/cua_s1/native/src/main.rs @@ -1,125 +1,113 @@ -//! Cua-S1 4B 0.2 (`text` adapter) `/v1/systemone` worker on native CUDA kernels. +//! Cua-S1 4B 0.2 (`text` adapter) `/v1/systemone` worker on the native CUDA kernels. //! -//! omni-cua-s1-native --model [--port 8000] +//! CUA_S1_MODEL= omni-cua-s1-native //! -//! See recipe/cua_s1/native.md for building the CUDA library and exporting the merged -//! checkpoint. +//! `CUA_S1_CUDA_LIB` (default: next to this executable), `CUA_S1_HOST` and `CUA_S1_PORT` +//! are optional; see recipe/cua_s1/native.md. -use std::path::PathBuf; use std::sync::Arc; -use std::time::Instant; -use anyhow::{Result, ensure}; -use clap::Parser; -use clap::builder::RangedU64ValueParser; +use anyhow::{Context, Result, ensure}; +use axum::body::Bytes; +use axum::extract::rejection::BytesRejection; +use axum::extract::{DefaultBodyLimit, State}; +use axum::http::StatusCode; +use axum::response::{IntoResponse, Response}; +use axum::routing::{get, post}; +use axum::{Json, Router}; +use serde_json::{Value, json}; -use omni_cua_s1_native::contract; +use omni_cua_s1_native::contract::{self, MODEL_ID}; use omni_cua_s1_native::cuda; use omni_cua_s1_native::engine::Engine; -use omni_cua_s1_native::model::Options; -use omni_cua_s1_native::server::{self, App, Limits}; +use omni_cua_s1_native::json::quote; -#[derive(Parser)] -#[command(about = "Cua-S1 4B 0.2 text worker on native CUDA kernels")] -struct Args { - /// Merged text checkpoint: Qwen/Qwen3.5-4B with the `text` adapter merged, plus - /// the cua_s1_export.json that recipe/cua_s1/export_text_merged.py writes. - #[arg(long, env = "CUA_S1_MODEL")] - model: PathBuf, - /// libqwen3_5_cuda.so, built by src/backends/cuda/qwen3_5/build.sh [default: next - /// to this executable]. - #[arg(long, env = "CUA_S1_CUDA_LIB")] - cuda_lib: Option, - #[arg(long, env = "CUA_S1_HOST", default_value = "127.0.0.1")] - host: String, - #[arg(long, env = "CUA_S1_PORT", default_value_t = 8000)] - port: u16, - #[arg(long, env = "CUA_S1_MAX_BODY_BYTES", default_value_t = 4 << 20)] - max_body_bytes: usize, - #[arg(long, env = "CUA_S1_MAX_QUESTIONS", default_value_t = 64)] - max_questions: usize, - /// Per question; 0 disables the check. - #[arg(long, env = "CUA_S1_MAX_PROMPT_TOKENS", default_value_t = 16384)] - max_prompt_tokens: usize, - /// Prompts up to this many tokens run as a CUDA graph captured for their length - /// on first use; longer ones run eagerly. 0 runs everything eagerly, with - /// cuBLASLt's first-choice GEMM algorithms and no tuning. - #[arg(long, env = "CUA_S1_GRAPH_MAX_TOKENS", default_value_t = 2048)] - graph_max_tokens: usize, - /// How many prompt lengths keep their captured graph. - #[arg(long, env = "CUA_S1_GRAPH_CACHE", default_value_t = 128, - value_parser = RangedU64ValueParser::::new().range(1..))] - graph_cache: usize, - /// GEMM algorithm choices: read from this file if it exists, else tuned at - /// startup and written to it, so later starts make the same choices. A file tuned - /// on another GPU or cuBLASLt version, or for another --graph-max-tokens, is - /// refused. - #[arg(long, env = "CUA_S1_GEMM_PLANS")] - gemm_plans: Option, - /// When tuning, time far more cuBLASLt configurations for prompts up to - /// --graph-max-tokens (each algorithm with its tiles, stage counts, swizzles and - /// several split-K factors) instead of the heuristic's shortlist. Takes about a - /// minute; use it with --gemm-plans so that it runs once. - #[arg(long, env = "CUA_S1_GEMM_SEARCH")] - gemm_search: bool, +const MAX_BODY_BYTES: usize = 4 << 20; +const MAX_PROMPT_TOKENS: usize = 16384; +const WARMUP: &[u8] = br#"{"model": "cua-s1-4b-0.2", "state": "Dialog: Update installed.", "questions": {"q": {"type": "choice", "instructions": "Close it.", "criteria": {"ok": "OK", "wait": "Wait"}}}}"#; + +fn reply(status: u16, body: Value) -> Response { + (StatusCode::from_u16(status).unwrap(), Json(body)).into_response() } -impl Args { - fn options(&self) -> Result { - ensure!( - self.graph_max_tokens > 0 || (self.gemm_plans.is_none() && !self.gemm_search), - "--gemm-plans and --gemm-search need --graph-max-tokens above 0" - ); - Ok(Options { - library: match &self.cuda_lib { - Some(path) => path.clone(), - None => cuda::default_library()?, - }, - graph_max_tokens: self.graph_max_tokens, - graph_cache: self.graph_cache, - gemm_plans: self.gemm_plans.clone(), - gemm_search: self.gemm_search, - }) +/// Every prompt is tokenized and checked against the limit before any forward pass. +async fn decide(engine: &Engine, raw: &[u8]) -> Response { + let (state, questions) = match contract::parse_body(raw).and_then(|b| contract::map_request(&b)) + { + Ok(request) => request, + Err(e) => return reply(e.status, json!({"detail": e.message})), + }; + let failed = |e: anyhow::Error| { + eprintln!("inference failed: {e:#}"); + reply(500, json!({"detail": "inference failed"})) + }; + let mut prompts = Vec::with_capacity(questions.len()); + for q in &questions { + let ids = match engine.encode(&state, q) { + Ok(ids) => ids, + Err(e) => return failed(e), + }; + if ids.len() > MAX_PROMPT_TOKENS { + let message = format!( + "question {}: {} prompt tokens, over {MAX_PROMPT_TOKENS}", + quote(&q.name), + ids.len() + ); + return reply(413, json!({"detail": message})); + } + prompts.push(ids); + } + let tokens: usize = prompts.iter().map(Vec::len).sum(); + let mut answers = serde_json::Map::new(); + for (q, ids) in questions.iter().zip(prompts) { + match engine.score(ids, q.keys.len()).await { + Ok(probs) => answers.insert(q.name.clone(), contract::answer(q, &probs)), + Err(e) => return failed(e), + }; + } + reply( + 200, + json!({"model": MODEL_ID, "answers": answers, "usage": {"input_tokens": tokens, "output_tokens": 0}}), + ) +} + +async fn systemone( + State(engine): State>, + body: Result, +) -> Response { + match body { + Ok(raw) => decide(&engine, &raw).await, + Err(e) => reply(e.status().as_u16(), json!({"detail": e.body_text()})), } } #[tokio::main] async fn main() -> Result<()> { - let args = Args::parse(); - let engine = Engine::load(&args.model, &args.options()?).await?; - if engine.provenance.adapter_revision != contract::ADAPTER_REVISION { - eprintln!( - "warning: adapter revision {} is not the pinned {}", - engine.provenance.adapter_revision, - contract::ADAPTER_REVISION - ); - } - if engine.provenance.base_revision != contract::BASE_REVISION { - eprintln!( - "warning: base revision {} is not the pinned {}", - engine.provenance.base_revision, - contract::BASE_REVISION - ); - } - println!( - "loaded in {:.1} s on {} (bfloat16, graphs up to {} tokens)", - engine.load_seconds, - engine.device, - engine.graph_max_tokens() - ); - let api_key = std::env::var_os("CUA_S1_API_KEY").map(|k| k.into_encoded_bytes()); - let limits = Limits { - max_body_bytes: args.max_body_bytes, - max_questions: args.max_questions, - max_prompt_tokens: args.max_prompt_tokens, + let model = std::env::var_os("CUA_S1_MODEL").context("set CUA_S1_MODEL")?; + let library = match std::env::var_os("CUA_S1_CUDA_LIB") { + Some(path) => path.into(), + None => cuda::default_library()?, }; - let revision = engine.provenance.adapter_revision.clone(); - let app = Arc::new(App::new(engine, limits, api_key, &revision)); - let started = Instant::now(); - server::warmup(&app).await?; - println!("warmed up in {:.1} s", started.elapsed().as_secs_f64()); - let listener = tokio::net::TcpListener::bind((args.host.as_str(), args.port)).await?; - println!("listening on {}:{}", args.host, args.port); - axum::serve(listener, server::router(app)).await?; + let engine = Arc::new(Engine::load(model.as_ref(), &library).await?); + // one decision before listening, so the first request does not pay for first-call setup + ensure!( + decide(&engine, WARMUP).await.status() == StatusCode::OK, + "warmup failed" + ); + let host = std::env::var("CUA_S1_HOST").unwrap_or_else(|_| "127.0.0.1".into()); + let port: u16 = std::env::var("CUA_S1_PORT") + .map_or(Ok(8000), |p| p.parse()) + .context("CUA_S1_PORT")?; + let app = Router::new() + .route( + "/health", + get(|| async { Json(json!({"status": "ready", "model": MODEL_ID})) }), + ) + .route("/v1/systemone", post(systemone)) + .layer(DefaultBodyLimit::max(MAX_BODY_BYTES)) + .with_state(engine); + let listener = tokio::net::TcpListener::bind((host.as_str(), port)).await?; + println!("listening on {host}:{port}"); + axum::serve(listener, app).await?; Ok(()) } diff --git a/src/models/cua_s1/native/src/model.rs b/src/models/cua_s1/native/src/model.rs index 9321c6e8..2b040779 100644 --- a/src/models/cua_s1/native/src/model.rs +++ b/src/models/cua_s1/native/src/model.rs @@ -1,17 +1,7 @@ //! The Qwen3.5 text model (the language model of Qwen/Qwen3.5-4B), prefill only: one //! forward pass over a prompt, returning the final-norm hidden state of the last -//! position. The layer loop, buffers, CUDA graphs and kernel choice live here; the -//! operations are the CUDA kernels in `src/backends/cuda/qwen3_5`. -//! -//! Prompts up to `graph_max_tokens` run as a CUDA graph captured for their exact -//! length on first use (the most recent `graph_cache` lengths are kept); longer ones -//! run eagerly. A graph queues the same kernels with the same GEMM algorithms as the -//! eager pass of that length, so both give bitwise identical results. -//! -//! GEMM algorithms are tuned at startup for the lengths in TUNE_ROWS; other lengths -//! borrow a nearby tuned one (see gemm.cu). The choices can be saved to a file and -//! reused, so that restarts do not change them; the file records the GPU, the -//! cuBLASLt version and the tuned lengths, and one that does not match is refused. +//! position. The layer loop and buffers live here; the operations are the CUDA +//! kernels in `src/backends/cuda/qwen3_5`. //! //! The order of operations follows `modeling_qwen3_5.py`, and so do the points where //! it rounds to bfloat16, except inside attention and the Gated DeltaNet prefill (see @@ -19,64 +9,20 @@ //! token, so the multimodal rotary sections all get the same position and the //! rotary embedding is the plain one. -use std::collections::{HashMap, VecDeque}; +use std::collections::HashMap; use std::ffi::c_void; -use std::path::{Path, PathBuf}; +use std::path::Path; use anyhow::{Context, Result, bail, ensure}; use serde_json::Value as Json; -use crate::cuda::{self, DeviceBuffer, Graph, Stream, check}; +use crate::cuda::{self, DeviceBuffer, Stream, check}; const ALIGN: usize = 256; const BF16: usize = 2; const F32: usize = 4; const GEMM_WORKSPACE: usize = 32 << 20; -/// What the GEMM plans depend on besides the shapes: the GPU, cuBLASLt, the -/// workspace, and the longest prompt tuned in the graph range. -fn plan_setup(graph_max_tokens: usize) -> Result { - let (gpu, compute_capability, sms) = cuda::device_info()?; - // SAFETY: takes no arguments. - let cublaslt = unsafe { (cuda::api().cs1_gemm_version)() }; - Ok(serde_json::json!({ - "gpu": gpu, - "compute_capability": compute_capability, - "sms": sms, - "cublaslt": cublaslt, - "workspace_bytes": GEMM_WORKSPACE, - "graph_max_tokens": graph_max_tokens, - })) -} - -/// Prompt lengths the GEMM algorithms are tuned for, at most twice apart, so that a -/// length up to the last one borrows a tuned algorithm for at most twice its length -/// and longer ones borrow the last one's. Past these, cuBLASLt's first choice for long -/// prompts is an older, half-rate tensor-core kernel on sm_89. -pub const TUNE_ROWS: &[usize] = &[ - 64, 96, 128, 160, 192, 224, 256, 320, 384, 448, 512, 640, 768, 1024, 1536, 2048, 4096, 8192, - 16384, -]; - -/// Startup options of the model. -#[derive(Debug, Clone)] -pub struct Options { - /// libqwen3_5_cuda.so (see src/backends/cuda/qwen3_5/build.sh). - pub library: PathBuf, - /// Longest prompt that runs as a CUDA graph; 0 runs everything eagerly and skips - /// GEMM tuning. - pub graph_max_tokens: usize, - /// How many prompt lengths keep their captured graph. - pub graph_cache: usize, - /// Tuned GEMM algorithms: read from this file if it exists, else tuned and - /// written to it. - pub gemm_plans: Option, - /// Tune the GEMMs of graph-length prompts over far more cuBLASLt configurations - /// than the heuristic's shortlist (see gemm.cu; about a minute). Longer prompts - /// keep the shortlist, whose choices did better in whole forward passes. - pub gemm_search: bool, -} - #[derive(Debug, Clone)] pub struct Config { pub hidden: usize, @@ -566,14 +512,8 @@ pub struct Model { layers: Vec, stream: Stream, gemm: *mut c_void, - /// Graphs by prompt length, least recently used first in `graph_lru`; they point - /// into `graph_scratch`, which is never reallocated. - graphs: HashMap, - graph_lru: VecDeque, - graph_scratch: Option, - graph_cache: usize, - /// For prompts longer than `graph_scratch` holds; grows as needed. - eager_scratch: Option, + /// Buffers for the longest prompt so far; grows as needed. + scratch: Option, } // SAFETY: the raw pointers are device addresses and a cuBLASLt handle owned by the @@ -582,17 +522,16 @@ unsafe impl Send for Model {} impl Drop for Model { fn drop(&mut self) { - self.graphs.clear(); // SAFETY: created by cs1_gemm_create and not destroyed before. unsafe { (cuda::api().cs1_gemm_destroy)(self.gemm) }; } } impl Model { - /// Load the weights, then tune the GEMMs (or read their plans) for graph use. - pub fn load(dir: &Path, opts: &Options) -> Result { + /// Load the CUDA library and the weights. + pub fn load(dir: &Path, library: &Path) -> Result { let cfg = Config::load(dir)?; - cuda::load(&opts.library)?; + cuda::load(library)?; cuda::set_device(0)?; let stream = cuda::new_stream()?; let weights = Weights::load(dir, stream)?; @@ -653,7 +592,7 @@ impl Model { // SAFETY: plain allocation; checked for null below. let gemm = unsafe { (cuda::api().cs1_gemm_create)(GEMM_WORKSPACE) }; ensure!(!gemm.is_null(), "cuBLASLt setup failed"); - let mut model = Self { + let model = Self { cfg, _weights: weights, embed, @@ -661,193 +600,11 @@ impl Model { layers, stream, gemm, - graphs: HashMap::new(), - graph_lru: VecDeque::new(), - graph_scratch: None, - graph_cache: opts.graph_cache.max(1), - eager_scratch: None, + scratch: None, }; - if opts.graph_max_tokens > 0 { - model.prepare_graphs( - opts.graph_max_tokens, - opts.gemm_plans.as_deref(), - opts.gemm_search, - )?; - } Ok(model) } - /// The longest prompt that runs as a graph (0: none). - pub fn graph_max_tokens(&self) -> usize { - self.graph_scratch.as_ref().map_or(0, |s| s.cap) - } - - /// Allocate the graph buffers, then read or tune the GEMM plans. - fn prepare_graphs(&mut self, max: usize, plans: Option<&Path>, search: bool) -> Result<()> { - let s = Scratch::new(&self.cfg, max, self.stream)?; - // one eager pass over the longest prompt fills every buffer and sets up the kernels - let zeros = vec![0u8; max * 4]; - // SAFETY: the ids buffer holds `max` int32 values. - unsafe { cuda::upload(s.at(s.ids), &zeros, self.stream)? }; - self.run(&s, max)?; - cuda::synchronize(self.stream)?; - let setup = plan_setup(max)?; - let loaded = match plans { - Some(path) if path.exists() => { - let n = self.import_plans(path, &setup).with_context(|| { - format!( - "GEMM plans in {} not used; remove the file, or pass another \ - --gemm-plans, to tune again", - path.display() - ) - })?; - eprintln!("GEMM plans: {n} read from {}", path.display()); - true - } - _ => false, - }; - if !loaded { - let mut rows: Vec = TUNE_ROWS.iter().copied().filter(|&m| m < max).collect(); - rows.push(max); - for m in rows { - self.tune(&s, m, search)?; - } - // longer prompts run eagerly; tune those lengths in a temporary buffer - let long: Vec = TUNE_ROWS.iter().copied().filter(|&m| m > max).collect(); - if let Some(&cap) = long.last() { - let big = Scratch::new(&self.cfg, cap, self.stream)?; - for m in long { - self.tune(&big, m, false)?; - } - } - // SAFETY: frees only the tuning buffers. - unsafe { (cuda::api().cs1_gemm_tune_done)(self.gemm) }; - if let Some(path) = plans { - let n = self.export_plans(path, &setup, search)?; - eprintln!("GEMM plans: {n} written to {}", path.display()); - } - } - self.graph_scratch = Some(s); - Ok(()) - } - - fn export_plans(&self, path: &Path, setup: &Json, search: bool) -> Result { - // SAFETY: a null buffer with capacity 0 only counts. - let n = unsafe { (cuda::api().cs1_gemm_export)(self.gemm, std::ptr::null_mut(), 0) }; - let mut plans = vec![cuda::GemmPlan::default(); n]; - // SAFETY: `plans` has room for n records. - unsafe { (cuda::api().cs1_gemm_export)(self.gemm, plans.as_mut_ptr(), n) }; - let records: Vec = plans - .iter() - .map(|p| { - serde_json::json!({ - "m": p.m, "n": p.n, "k": p.k, "ldy": p.ldy, - "algo": p.algo.iter().map(|w| format!("{w:016x}")).collect::>(), - }) - }) - .collect(); - let doc = serde_json::json!({ - "format": "cua-s1-native-gemm-plans/1", - "setup": setup, - "search": search, - "plans": records, - }); - std::fs::write(path, serde_json::to_string_pretty(&doc)? + "\n") - .with_context(|| format!("{}", path.display()))?; - Ok(n) - } - - fn import_plans(&self, path: &Path, setup: &Json) -> Result { - let doc: Json = serde_json::from_str(&std::fs::read_to_string(path)?)?; - ensure!( - doc["format"] == "cua-s1-native-gemm-plans/1", - "unknown format {}", - doc["format"] - ); - ensure!( - doc["setup"].is_object(), - "the file does not record the GPU and cuBLASLt version it was tuned for" - ); - ensure!( - doc["setup"] == *setup, - "tuned for {}, this run is {setup}", - doc["setup"] - ); - let int = |v: &Json| v.as_i64().map(|x| x as i32).context("bad plan field"); - let plans = doc["plans"] - .as_array() - .context("no plans")? - .iter() - .map(|p| { - let words = p["algo"].as_array().context("bad algo")?; - ensure!(words.len() == 8, "bad algo"); - let mut algo = [0u64; 8]; - for (dst, w) in algo.iter_mut().zip(words) { - *dst = u64::from_str_radix(w.as_str().context("bad algo")?, 16)?; - } - Ok(cuda::GemmPlan { - m: int(&p["m"])?, - n: int(&p["n"])?, - k: int(&p["k"])?, - ldy: int(&p["ldy"])?, - // the file's setup, checked above, records the version - cublaslt_version: setup["cublaslt"].as_u64().context("no cuBLASLt version")?, - algo, - }) - }) - .collect::>>()?; - // SAFETY: `plans` holds plans.len() records. - check( - unsafe { (cuda::api().cs1_gemm_import)(self.gemm, plans.as_ptr(), plans.len()) }, - "importing GEMM plans", - )?; - Ok(plans.len()) - } - - /// Pick the GEMM algorithms for `m` rows, timing the first layer of each kind. - fn tune(&self, s: &Scratch, m: usize, exhaustive: bool) -> Result<()> { - let mut jobs: Vec<(usize, &Tensor, usize)> = Vec::new(); - let layer = &self.layers[0]; - jobs.push((s.x, &layer.gate_up, s.gate_up)); - jobs.push((s.act, &layer.down, s.delta)); - if let Some(la) = self.layers.iter().find_map(|l| match &l.mixer { - Mixer::Linear(la) => Some(la), - Mixer::Full(_) => None, - }) { - jobs.push((s.x, &la.in_proj, s.gdn_in)); - jobs.push((s.ln, &la.out, s.delta)); - } - if let Some(fa) = self.layers.iter().find_map(|l| match &l.mixer { - Mixer::Full(fa) => Some(fa), - Mixer::Linear(_) => None, - }) { - jobs.push((s.x, &fa.qkv, s.attn_in)); - jobs.push((s.ao, &fa.o, s.delta)); - } - for (x, w, y) in jobs { - let (n, k) = (w.shape[0] as i32, w.shape[1] as i32); - // SAFETY: x and y are scratch buffers sized for `cap` >= m rows of this shape. - check( - unsafe { - (cuda::api().cs1_gemm_tune)( - self.gemm, - s.at(x), - w.ptr, - s.at(y), - m as i32, - n, - k, - n, - exhaustive.into(), - self.stream, - ) - }, - "gemm tuning", - )?; - } - Ok(()) - } - fn gemm(&self, s: &Scratch, x: usize, w: &Tensor, y: usize, m: usize) -> Result<()> { let (n, k) = (w.shape[0] as i32, w.shape[1] as i32); // SAFETY: x and y are scratch buffers sized for m rows of w's shape. @@ -870,7 +627,6 @@ impl Model { } /// The final-norm hidden state at the last position, as float32. - /// Runs from the graph for this length when graphs are on and it fits, else eagerly. pub fn forward(&mut self, ids: &[u32]) -> Result> { let t = ids.len(); ensure!(t > 0, "empty prompt"); @@ -880,36 +636,19 @@ impl Model { "token id outside the vocabulary" ); cuda::set_device(0)?; - let in_graph_scratch = self.graph_scratch.as_ref().is_some_and(|s| t <= s.cap); - if !in_graph_scratch && self.eager_scratch.as_ref().is_none_or(|s| t > s.cap) { - self.eager_scratch = None; - self.eager_scratch = Some(Scratch::new( + if self.scratch.as_ref().is_none_or(|s| t > s.cap) { + self.scratch = None; + self.scratch = Some(Scratch::new( &self.cfg, t.next_multiple_of(1024), self.stream, )?); } - let s = if in_graph_scratch { - self.graph_scratch.as_ref() - } else { - self.eager_scratch.as_ref() - } - .unwrap(); + let s = self.scratch.as_ref().unwrap(); let ids32: Vec = ids.iter().flat_map(|&i| (i as i32).to_le_bytes()).collect(); // SAFETY: the ids buffer holds at least t int32 values. unsafe { cuda::upload(s.at(s.ids), &ids32, self.stream)? }; - if in_graph_scratch { - self.graph_for(t)?; - self.graphs[&t].launch(self.stream)?; - } else { - self.run(s, t)?; - } - let s = if in_graph_scratch { - self.graph_scratch.as_ref() - } else { - self.eager_scratch.as_ref() - } - .unwrap(); + self.run(s, t)?; let mut last = vec![0u8; h * BF16]; // SAFETY: x holds at least t rows of the hidden size. unsafe { cuda::download(&mut last, s.at(s.x + (t - 1) * h * BF16), self.stream)? }; @@ -920,28 +659,8 @@ impl Model { .collect()) } - /// Make sure a graph for `t` tokens exists (capturing it if needed) and mark it - /// as the most recently used, dropping the least recently used beyond the cache. - fn graph_for(&mut self, t: usize) -> Result<()> { - if self.graphs.contains_key(&t) { - self.graph_lru.retain(|&x| x != t); - } else { - while self.graphs.len() >= self.graph_cache { - let Some(old) = self.graph_lru.pop_front() else { - break; - }; - self.graphs.remove(&old); - } - let s = self.graph_scratch.as_ref().context("no graph buffers")?; - let graph = Graph::capture(self.stream, || self.run(s, t))?; - self.graphs.insert(t, graph); - } - self.graph_lru.push_back(t); - Ok(()) - } - - /// Queue one forward pass over the first `t` ids in `s` (nothing else is queued, - /// so it can be captured). The final-norm hidden states end up in `s.x`. + /// Queue one forward pass over the first `t` ids in `s`. The final-norm hidden + /// states end up in `s.x`. fn run(&self, s: &Scratch, t: usize) -> Result<()> { let cfg = &self.cfg; let st = self.stream; diff --git a/src/models/cua_s1/native/src/pyjson.rs b/src/models/cua_s1/native/src/pyjson.rs deleted file mode 100644 index 7c60e1d5..00000000 --- a/src/models/cua_s1/native/src/pyjson.rs +++ /dev/null @@ -1,858 +0,0 @@ -//! JSON the way the Python worker sees it. -//! -//! - [`parse`] follows CPython 3.12's `json.loads` together with the checks that -//! `contract.parse_body` adds: duplicate keys, `NaN`/`Infinity`, numbers that are -//! out of range for a float, lone surrogates and very deep nesting are all rejected, -//! and errors surface in the same order as in Python. -//! - [`dumps`] follows `json.dumps(value, ensure_ascii=False)`. -//! - [`repr`] follows Python's `repr`, which the error messages quote, except that -//! characters outside ASCII are written as they are unless they are lone surrogates. - -use std::collections::HashSet; -use std::fmt::Write as _; - -/// Nesting limits of the Python worker, measured on CPython 3.12.13 under uvicorn. -/// Both come from interpreter recursion limits, so they depend on the call stack and -/// are not documented constants. -/// -/// Parsing fails when a container would open more than this many levels deep (the -/// `json` C scanner's recursion check). -pub const MAX_PARSE_DEPTH: usize = 9990; - -/// After parsing, `parse_body` walks the value with a recursive Python function -/// (`_check_unicode`), one call per value, containers and scalars alike. A walk that -/// needs more nested calls than this fails. -pub const MAX_CHECK_CALLS: usize = 969; - -/// `sys.int_info.default_max_str_digits`: longer integers fail to convert in Python. -const MAX_INT_DIGITS: usize = 4300; - -/// A string as Python holds it: a sequence of code points that may include lone -/// surrogates, stored as generalized UTF-8. -#[derive(Clone, Debug, Default, PartialEq, Eq, Hash)] -pub struct PyStr(Vec); - -impl PyStr { - pub fn new(s: &str) -> Self { - PyStr(s.as_bytes().to_vec()) - } - - fn push(&mut self, cp: u32) { - let b = &mut self.0; - match cp { - 0..=0x7f => b.push(cp as u8), - 0x80..=0x7ff => { - b.push(0xc0 | (cp >> 6) as u8); - b.push(0x80 | (cp & 0x3f) as u8); - } - 0x800..=0xffff => { - b.push(0xe0 | (cp >> 12) as u8); - b.push(0x80 | ((cp >> 6) & 0x3f) as u8); - b.push(0x80 | (cp & 0x3f) as u8); - } - _ => { - b.push(0xf0 | (cp >> 18) as u8); - b.push(0x80 | ((cp >> 12) & 0x3f) as u8); - b.push(0x80 | ((cp >> 6) & 0x3f) as u8); - b.push(0x80 | (cp & 0x3f) as u8); - } - } - } - - /// The text, or None if it holds a lone surrogate. - pub fn as_str(&self) -> Option<&str> { - std::str::from_utf8(&self.0).ok() - } - - pub fn code_points(&self) -> impl Iterator + '_ { - let b = &self.0; - let mut i = 0; - std::iter::from_fn(move || { - let lead = *b.get(i)? as u32; - let (len, init) = match lead { - 0..=0x7f => (1, lead), - 0xc0..=0xdf => (2, lead & 0x1f), - 0xe0..=0xef => (3, lead & 0x0f), - _ => (4, lead & 0x07), - }; - let cp = b[i + 1..i + len] - .iter() - .fold(init, |acc, &c| (acc << 6) | (c as u32 & 0x3f)); - i += len; - Some(cp) - }) - } -} - -impl PartialEq for PyStr { - fn eq(&self, other: &str) -> bool { - self.0 == other.as_bytes() - } -} - -#[derive(Clone, Debug, PartialEq)] -pub enum Value { - Null, - Bool(bool), - /// An integer as Python would print it (JSON integers never have leading zeros, - /// so this is the source text, with `-0` read as `0`). - Int(String), - Float(f64), - Str(PyStr), - Array(Vec), - /// Key order is kept; keys are unique once parsing succeeds. - Object(Vec<(PyStr, Value)>), -} - -/// Values can nest thousands of levels deep before the depth check rejects them, so -/// they are dropped without recursion. -impl Drop for Value { - fn drop(&mut self) { - let mut stack: Vec = Vec::new(); - take_children(self, &mut stack); - while let Some(mut v) = stack.pop() { - take_children(&mut v, &mut stack); - } - } -} - -fn take_children(value: &mut Value, stack: &mut Vec) { - match value { - Value::Array(items) => stack.append(items), - Value::Object(pairs) => stack.extend(pairs.drain(..).map(|(_, v)| v)), - _ => {} - } -} - -/// Why a body was rejected as JSON; every case is a 400. -#[derive(Debug, PartialEq)] -pub enum JsonError { - Utf8, - Syntax, - Depth, - Duplicate(PyStr), - Constant(&'static str), - Range(String), - NotObject, -} - -impl JsonError { - /// The `detail` the Python worker returns for this error. - pub fn message(&self) -> String { - match self { - JsonError::Utf8 => "request body must be valid UTF-8 text".into(), - JsonError::Syntax => "request body must be valid JSON".into(), - JsonError::Depth => "request body is nested too deeply".into(), - JsonError::Duplicate(key) => { - format!("duplicate key {} in a JSON object", repr_str(key)) - } - JsonError::Constant(name) => format!("{name} is not valid JSON"), - JsonError::Range(text) => format!("number {text} is out of range"), - JsonError::NotObject => "request body must be a JSON object".into(), - } - } -} - -/// Decode a request body into its top-level object. -pub fn parse(raw: &[u8]) -> Result, JsonError> { - let text = std::str::from_utf8(raw).map_err(|_| JsonError::Utf8)?; - let mut p = Parser { - s: text.as_bytes(), - i: 0, - }; - p.ws(); - let mut value = p.value()?; - p.ws(); - if p.i != p.s.len() { - return Err(JsonError::Syntax); - } - check(&value, 1)?; - match &mut value { - Value::Object(pairs) => Ok(std::mem::take(pairs)), - _ => Err(JsonError::NotObject), - } -} - -/// `_check_unicode` from `contract.py`: visit values in order (an object's key before -/// its value) and fail at the first lone surrogate or the first call past -/// [`MAX_CHECK_CALLS`], whichever comes first. -fn check(value: &Value, calls: usize) -> Result<(), JsonError> { - if calls > MAX_CHECK_CALLS { - return Err(JsonError::Depth); - } - match value { - Value::Str(s) if s.as_str().is_none() => Err(JsonError::Utf8), - Value::Array(items) => items.iter().try_for_each(|v| check(v, calls + 1)), - Value::Object(pairs) => pairs.iter().try_for_each(|(k, v)| { - if k.as_str().is_none() { - return Err(JsonError::Utf8); - } - check(v, calls + 1) - }), - _ => Ok(()), - } -} - -/// An array or object that has been opened but not closed yet. -enum Open { - Array(Vec), - /// the members so far, and the key whose value is being parsed - Object(Vec<(PyStr, Value)>, PyStr), -} - -struct Parser<'a> { - s: &'a [u8], - i: usize, -} - -impl Parser<'_> { - fn peek(&self) -> Option { - self.s.get(self.i).copied() - } - - fn rest_starts_with(&self, lit: &[u8]) -> bool { - self.s[self.i..].starts_with(lit) - } - - fn ws(&mut self) { - while let Some(b' ' | b'\t' | b'\n' | b'\r') = self.peek() { - self.i += 1; - } - } - - /// One JSON value. Containers are parsed with an explicit stack, so deep nesting - /// cannot overflow the thread's stack before [`MAX_PARSE_DEPTH`] stops it. - fn value(&mut self) -> Result { - let mut open: Vec = Vec::new(); - loop { - // the start of a value: open a container or read a scalar - let mut done = match self.peek() { - Some(b'{') => { - if open.len() + 1 > MAX_PARSE_DEPTH { - return Err(JsonError::Depth); - } - self.i += 1; - self.ws(); - if self.peek() == Some(b'}') { - self.i += 1; - Value::Object(Vec::new()) - } else { - let key = self.key()?; - open.push(Open::Object(Vec::new(), key)); - continue; - } - } - Some(b'[') => { - if open.len() + 1 > MAX_PARSE_DEPTH { - return Err(JsonError::Depth); - } - self.i += 1; - self.ws(); - if self.peek() == Some(b']') { - self.i += 1; - Value::Array(Vec::new()) - } else { - open.push(Open::Array(Vec::new())); - continue; - } - } - _ => self.scalar()?, - }; - // hand the finished value to its container, closing containers that end - loop { - match open.last_mut() { - None => return Ok(done), - Some(Open::Array(items)) => { - items.push(done); - self.ws(); - match self.peek() { - Some(b',') => { - self.i += 1; - self.ws(); - break; - } - Some(b']') => { - self.i += 1; - let Some(Open::Array(items)) = open.pop() else { - unreachable!() - }; - done = Value::Array(items); - } - _ => return Err(JsonError::Syntax), - } - } - Some(Open::Object(pairs, key)) => { - pairs.push((std::mem::take(key), done)); - self.ws(); - match self.peek() { - Some(b',') => { - self.i += 1; - self.ws(); - *key = self.key()?; - break; - } - Some(b'}') => { - self.i += 1; - let Some(Open::Object(pairs, _)) = open.pop() else { - unreachable!() - }; - check_duplicates(&pairs)?; - done = Value::Object(pairs); - } - _ => return Err(JsonError::Syntax), - } - } - } - } - } - } - - /// A member name and its colon. - fn key(&mut self) -> Result { - if self.peek() != Some(b'"') { - return Err(JsonError::Syntax); - } - self.i += 1; - let key = self.string()?; - self.ws(); - if self.peek() != Some(b':') { - return Err(JsonError::Syntax); - } - self.i += 1; - self.ws(); - Ok(key) - } - - fn scalar(&mut self) -> Result { - match self.peek() { - Some(b'"') => { - self.i += 1; - Ok(Value::Str(self.string()?)) - } - Some(b'n') if self.rest_starts_with(b"null") => { - self.i += 4; - Ok(Value::Null) - } - Some(b't') if self.rest_starts_with(b"true") => { - self.i += 4; - Ok(Value::Bool(true)) - } - Some(b'f') if self.rest_starts_with(b"false") => { - self.i += 5; - Ok(Value::Bool(false)) - } - Some(b'N') if self.rest_starts_with(b"NaN") => Err(JsonError::Constant("NaN")), - Some(b'I') if self.rest_starts_with(b"Infinity") => { - Err(JsonError::Constant("Infinity")) - } - Some(b'-') if self.rest_starts_with(b"-Infinity") => { - Err(JsonError::Constant("-Infinity")) - } - _ => self.number(), - } - } - - fn digits(&mut self) { - while let Some(b'0'..=b'9') = self.peek() { - self.i += 1; - } - } - - fn number(&mut self) -> Result { - let start = self.i; - if self.peek() == Some(b'-') { - self.i += 1; - } - match self.peek() { - Some(b'0') => self.i += 1, - Some(b'1'..=b'9') => self.digits(), - _ => return Err(JsonError::Syntax), - } - let mut is_float = false; - if self.peek() == Some(b'.') && matches!(self.s.get(self.i + 1), Some(b'0'..=b'9')) { - self.i += 1; - self.digits(); - is_float = true; - } - if let Some(b'e' | b'E') = self.peek() { - let e_start = self.i; - self.i += 1; - if let Some(b'+' | b'-') = self.peek() { - self.i += 1; - } - let digits_start = self.i; - self.digits(); - if self.i > digits_start { - is_float = true; - } else { - // not an exponent after all; what follows is left for the caller - self.i = e_start; - } - } - let text = std::str::from_utf8(&self.s[start..self.i]).expect("ASCII"); - if is_float { - let v: f64 = text.parse().map_err(|_| JsonError::Syntax)?; - if !v.is_finite() { - return Err(JsonError::Range(text.to_string())); - } - Ok(Value::Float(v)) - } else { - let digits = text.strip_prefix('-').unwrap_or(text); - if digits.len() > MAX_INT_DIGITS { - return Err(JsonError::Syntax); - } - Ok(Value::Int(if digits == "0" { - "0".into() - } else { - text.into() - })) - } - } - - fn hex4(&self, at: usize) -> Option { - let h = self.s.get(at..at + 4)?; - let mut v = 0; - for &c in h { - v = v * 16 + (c as char).to_digit(16)?; - } - Some(v) - } - - /// The rest of a string whose opening quote has been read. - fn string(&mut self) -> Result { - let mut out = PyStr::default(); - loop { - match self.peek().ok_or(JsonError::Syntax)? { - b'"' => { - self.i += 1; - return Ok(out); - } - b'\\' => { - let esc = *self.s.get(self.i + 1).ok_or(JsonError::Syntax)?; - self.i += 2; - let cp = match esc { - b'"' => '"' as u32, - b'\\' => '\\' as u32, - b'/' => '/' as u32, - b'b' => 0x08, - b'f' => 0x0c, - b'n' => '\n' as u32, - b'r' => '\r' as u32, - b't' => '\t' as u32, - b'u' => { - let mut c = self.hex4(self.i).ok_or(JsonError::Syntax)?; - self.i += 4; - // a high surrogate joins a directly following low one - if (0xd800..=0xdbff).contains(&c) - && self.rest_starts_with(b"\\u") - && let Some(c2) = self.hex4(self.i + 2) - && (0xdc00..=0xdfff).contains(&c2) - { - c = 0x10000 + ((c - 0xd800) << 10) + (c2 - 0xdc00); - self.i += 6; - } - c - } - _ => return Err(JsonError::Syntax), - }; - out.push(cp); - } - 0x00..=0x1f => return Err(JsonError::Syntax), - _ => { - let start = self.i; - while let Some(c) = self.peek() { - if c == b'"' || c == b'\\' || c < 0x20 { - break; - } - self.i += 1; - } - out.0.extend_from_slice(&self.s[start..self.i]); - } - } - } - } -} - -/// Python's object_pairs_hook runs when an object closes and reports the first key -/// that repeats an earlier one. -fn check_duplicates(pairs: &[(PyStr, Value)]) -> Result<(), JsonError> { - let mut seen = HashSet::with_capacity(pairs.len()); - for (key, _) in pairs { - if !seen.insert(key) { - return Err(JsonError::Duplicate(key.clone())); - } - } - Ok(()) -} - -/// Python's `repr(float)`: the shortest digits that round-trip, in fixed notation -/// for exponents from -5 to 15 and scientific notation otherwise. -pub fn float_repr(x: f64) -> String { - if x == 0.0 { - return if x.is_sign_negative() { "-0.0" } else { "0.0" }.into(); - } - let sci = format!("{x:e}"); - let (mantissa, exp) = sci.split_once('e').expect("{:e} has an exponent"); - let exp: i32 = exp.parse().expect("integer exponent"); - let (neg, mantissa) = match mantissa.strip_prefix('-') { - Some(m) => (true, m), - None => (false, mantissa), - }; - let digits: String = mantissa.chars().filter(|c| *c != '.').collect(); - let (digits, exp) = break_tie_to_even(x, digits, exp); - let mut out = String::new(); - if neg { - out.push('-'); - } - let decpt = exp + 1; - if decpt <= -4 || decpt > 16 { - out.push_str(&digits[..1]); - if digits.len() > 1 { - out.push('.'); - out.push_str(&digits[1..]); - } - let sign = if exp < 0 { '-' } else { '+' }; - write!(out, "e{sign}{:02}", exp.unsigned_abs()).unwrap(); - } else if decpt <= 0 { - out.push_str("0."); - out.extend(std::iter::repeat_n('0', (-decpt) as usize)); - out.push_str(&digits); - } else if decpt as usize >= digits.len() { - out.push_str(&digits); - out.extend(std::iter::repeat_n('0', decpt as usize - digits.len())); - out.push_str(".0"); - } else { - out.push_str(&digits[..decpt as usize]); - out.push('.'); - out.push_str(&digits[decpt as usize..]); - } - out -} - -/// When `x` lies exactly halfway between two shortest round-tripping decimals, Python -/// (David Gay's dtoa) takes the one with an even last digit, while Rust's shortest -/// formatting may take the other. Ties need an exact decimal expansion only one -/// digit longer than the shortest, which takes 16 or more significant digits. -fn break_tie_to_even(x: f64, digits: String, exp: i32) -> (String, i32) { - let n = digits.len(); - if n < 15 { - return (digits, exp); - } - // every finite double has at most 767 significant decimal digits - let exact = format!("{:.800e}", x.abs()); - let (mantissa, e) = exact.split_once('e').expect("{:e} has an exponent"); - let e: i32 = e.parse().expect("integer exponent"); - let full: String = mantissa.chars().filter(|c| *c != '.').collect(); - let full = full.trim_end_matches('0'); - if full.len() != n + 1 || !full.ends_with('5') { - return (digits, exp); - } - let floor = &full[..n]; - let last = floor.as_bytes()[n - 1] - b'0'; - let even = if last.is_multiple_of(2) { - floor.to_string() - } else if last == 9 { - // the even choice would carry, which exact ties (values in [2^50, 2^51) - // ending in .25 or .75) never need - return (digits, exp); - } else { - // the even choice is floor + 1 in the last digit - let mut up = floor.as_bytes().to_vec(); - up[n - 1] += 1; - String::from_utf8(up).expect("ASCII digits") - }; - // At a power of two the double below is half as far away, so one of the two - // candidates may read back as a different double; then it is not a tie. - let back: f64 = format!("{}.{}e{}", &even[..1], &even[1..], e) - .parse() - .expect("decimal digits"); - if back.to_bits() != x.abs().to_bits() { - return (digits, exp); - } - (even, e) -} - -/// A JSON string literal as `json.dumps(s, ensure_ascii=False)` writes it. -pub fn write_json_str(s: &str, out: &mut String) { - out.push('"'); - for c in s.chars() { - match c { - '"' => out.push_str("\\\""), - '\\' => out.push_str("\\\\"), - '\n' => out.push_str("\\n"), - '\r' => out.push_str("\\r"), - '\t' => out.push_str("\\t"), - '\u{8}' => out.push_str("\\b"), - '\u{c}' => out.push_str("\\f"), - c if (c as u32) < 0x20 => write!(out, "\\u{:04x}", c as u32).unwrap(), - c => out.push(c), - } - } - out.push('"'); -} - -/// `json.dumps(value, ensure_ascii=False)`. The value must not hold lone surrogates -/// (`parse` rejects them). -pub fn dumps(value: &Value) -> String { - let mut out = String::new(); - write_value(value, &mut out); - out -} - -fn write_value(value: &Value, out: &mut String) { - match value { - Value::Null => out.push_str("null"), - Value::Bool(b) => out.push_str(if *b { "true" } else { "false" }), - Value::Int(text) => out.push_str(text), - Value::Float(x) => out.push_str(&float_repr(*x)), - Value::Str(s) => write_json_str(s.as_str().expect("checked UTF-8"), out), - Value::Array(items) => { - out.push('['); - for (i, item) in items.iter().enumerate() { - if i > 0 { - out.push_str(", "); - } - write_value(item, out); - } - out.push(']'); - } - Value::Object(pairs) => { - out.push('{'); - for (i, (key, item)) in pairs.iter().enumerate() { - if i > 0 { - out.push_str(", "); - } - write_json_str(key.as_str().expect("checked UTF-8"), out); - out.push_str(": "); - write_value(item, out); - } - out.push('}'); - } - } -} - -/// Python's `repr(str)`, except that characters outside ASCII are kept as they are -/// unless they are lone surrogates. -pub fn repr_str(s: &PyStr) -> String { - let cps: Vec = s.code_points().collect(); - let squote = cps.contains(&('\'' as u32)); - let dquote = cps.contains(&('"' as u32)); - let quote = if squote && !dquote { '"' } else { '\'' }; - let mut out = String::with_capacity(cps.len() + 2); - out.push(quote); - for cp in cps { - match cp { - _ if cp == quote as u32 || cp == '\\' as u32 => { - out.push('\\'); - out.push(char::from_u32(cp).unwrap()); - } - 0x09 => out.push_str("\\t"), - 0x0a => out.push_str("\\n"), - 0x0d => out.push_str("\\r"), - 0..=0x1f | 0x7f => write!(out, "\\x{cp:02x}").unwrap(), - _ => match char::from_u32(cp) { - Some(c) => out.push(c), - // a lone surrogate - None => write!(out, "\\u{cp:04x}").unwrap(), - }, - } - } - out.push(quote); - out -} - -/// Python's `repr` of a decoded JSON value. -pub fn repr(value: &Value) -> String { - let mut out = String::new(); - write_repr(value, &mut out); - out -} - -fn write_repr(value: &Value, out: &mut String) { - match value { - Value::Null => out.push_str("None"), - Value::Bool(b) => out.push_str(if *b { "True" } else { "False" }), - Value::Int(text) => out.push_str(text), - Value::Float(x) => out.push_str(&float_repr(*x)), - Value::Str(s) => out.push_str(&repr_str(s)), - Value::Array(items) => { - out.push('['); - for (i, item) in items.iter().enumerate() { - if i > 0 { - out.push_str(", "); - } - write_repr(item, out); - } - out.push(']'); - } - Value::Object(pairs) => { - out.push('{'); - for (i, (key, item)) in pairs.iter().enumerate() { - if i > 0 { - out.push_str(", "); - } - out.push_str(&repr_str(key)); - out.push_str(": "); - write_repr(item, out); - } - out.push('}'); - } - } -} - -#[cfg(test)] -mod tests { - use super::*; - - fn err(body: &str) -> String { - parse(body.as_bytes()).unwrap_err().message() - } - - #[test] - fn float_repr_matches_python_examples() { - let cases = [ - (1.0, "1.0"), - (1e16, "1e+16"), - (1e15, "1000000000000000.0"), - (1e-5, "1e-05"), - (1e-4, "0.0001"), - (-0.0, "-0.0"), - (1e22, "1e+22"), - (3.14e-07, "3.14e-07"), - (0.1, "0.1"), - (5e-324, "5e-324"), - (1.7976931348623157e308, "1.7976931348623157e+308"), - // 2^-24: a digit string halfway to the next decimal, but the lower one - // reads back as the double below - (5.960464477539063e-08, "5.960464477539063e-08"), - (123456789.0, "123456789.0"), - (0.0024726232513785362, "0.0024726232513785362"), - // exactly halfway between two 17-digit candidates: Python takes the even one - (f64::from_bits(0xc31d5a973d792fa1), "-2065594985630696.2"), - ]; - for (x, want) in cases { - assert_eq!(float_repr(x), want, "{x:e}"); - } - } - - #[test] - fn dumps_matches_python() { - let v = parse(br#"{"a": [1.0, 1e16, 1e-5, 0.0001, -0.0, 1e22, 123456789012345678, 3.14e-07, -0, true, null]}"#).unwrap(); - let v = Value::Object(v); - assert_eq!( - dumps(&v), - r#"{"a": [1.0, 1e+16, 1e-05, 0.0001, -0.0, 1e+22, 123456789012345678, 3.14e-07, 0, true, null]}"# - ); - let s = Value::Str(PyStr::new("\u{0}\u{1f}\u{7f}\u{2028}\"\\/\t\u{8}\u{c}é😀")); - assert_eq!( - dumps(&s), - "\"\\u0000\\u001f\u{7f}\u{2028}\\\"\\\\/\\t\\b\\fé😀\"" - ); - } - - #[test] - fn repr_str_quotes_and_escapes() { - let s = PyStr::new("a\u{7f}\u{a0}😀é"); - assert_eq!(repr_str(&s), "'a\\x7f\u{a0}😀é'"); - assert_eq!(repr_str(&PyStr::new("it's")), "\"it's\""); - assert_eq!(repr_str(&PyStr::new("both'\"")), "'both\\'\"'"); - let v = Value::Object(parse(br#"{"k": [1, 2.5, null, true, "x"], "e": {}}"#).unwrap()); - assert_eq!(repr(&v), "{'k': [1, 2.5, None, True, 'x'], 'e': {}}"); - } - - #[test] - fn surrogates() { - let v = parse(b"{\"a\": \"\\ud83d\\ude00\"}").unwrap(); - assert_eq!(v[0].1, Value::Str(PyStr::new("😀"))); - assert_eq!( - err(r#"{"a": "\ud800x"}"#), - "request body must be valid UTF-8 text" - ); - assert_eq!( - err(r#"{"a": "\ude00\ud83d"}"#), - "request body must be valid UTF-8 text" - ); - // duplicate-key errors come first, and quote the surrogate - assert_eq!( - err(r#"{"\ud800": 1, "\ud800": 2}"#), - "duplicate key '\\ud800' in a JSON object" - ); - } - - #[test] - fn errors() { - assert_eq!(err("[]"), "request body must be a JSON object"); - assert_eq!(err("\u{feff}{}"), "request body must be valid JSON"); - assert_eq!(err(r#"{"a": NaN}"#), "NaN is not valid JSON"); - assert_eq!(err(r#"{"a": -Infinity}"#), "-Infinity is not valid JSON"); - assert_eq!(err(r#"{"a": 1e400}"#), "number 1e400 is out of range"); - assert_eq!(err(r#"{"a": 1e400x"#), "number 1e400 is out of range"); - assert_eq!( - err(r#"{"a": 1, "b": 2, "a": 3}"#), - "duplicate key 'a' in a JSON object" - ); - assert_eq!(err(r#"{"a": [1,]}"#), "request body must be valid JSON"); - assert_eq!(err("{\"a\": \"\u{1}\"}"), "request body must be valid JSON"); - assert_eq!(err(r#"{"a": 01}"#), "request body must be valid JSON"); - assert_eq!(err(r#"{"a": 1.}"#), "request body must be valid JSON"); - assert_eq!(err(r#"{"a": "\x"}"#), "request body must be valid JSON"); - assert_eq!(err(r#"{} x"#), "request body must be valid JSON"); - assert!(parse(br#"{"a": 1e-400}"#).is_ok()); - let long = format!("{{\"a\": {}}}", "1".repeat(4300)); - assert!(parse(long.as_bytes()).is_ok()); - let long = format!("{{\"a\": -{}}}", "1".repeat(4301)); - assert_eq!(err(&long), "request body must be valid JSON"); - assert_eq!(parse(b"{\"a\": \"\xff\"}").unwrap_err(), JsonError::Utf8); - } - - /// The Python worker's behaviour, measured over HTTP: `state` nested `d` lists deep - /// inside the top-level object. - #[test] - fn depth() { - let body = |d: usize, leaf: &str, tail: &str| { - format!( - "{{\"model\": \"m\", \"state\": {}{leaf}{}{tail}}}", - "[".repeat(d), - "]".repeat(d) - ) - }; - let deep = "request body is nested too deeply"; - // a scalar leaf takes one more call of the check than an empty container - assert!(parse(body(967, "\"x\"", "").as_bytes()).is_ok()); - assert_eq!(err(&body(968, "\"x\"", "")), deep); - assert!(parse(body(968, "", "").as_bytes()).is_ok()); - assert_eq!(err(&body(969, "", "")), deep); - assert!(parse(body(967, "{}", "").as_bytes()).is_ok()); - assert_eq!(err(&body(968, "{}", "")), deep); - // up to the parse limit, later parse errors win over depth - assert_eq!( - err(&(body(9989, "1", "") + "x")), - "request body must be valid JSON" - ); - assert_eq!(err(&(body(9990, "1", "") + "x")), deep); - assert_eq!( - err(&body(5000, "1", ", \"state\": 1")), - "duplicate key 'state' in a JSON object" - ); - // after parsing, the first problem in walk order wins - let deep_state = format!("{}1{}", "[".repeat(1500), "]".repeat(1500)); - assert_eq!( - err(&format!( - "{{\"model\": \"\\ud800\", \"state\": {deep_state}}}" - )), - "request body must be valid UTF-8 text" - ); - assert_eq!( - err(&format!("{{\"state\": {deep_state}, \"z\": \"\\ud800\"}}")), - deep - ); - assert_eq!( - err(&format!("{{\"a\": {deep_state}, \"\\ud800\": 1}}")), - deep - ); - // far past the limit: fails cleanly, without overflowing the stack - assert_eq!(err(&"[".repeat(1_000_000)), deep); - let wide = format!("{{\"a\": {}1{}}}", "[".repeat(9989), "]".repeat(9989)); - assert_eq!(err(&wide), deep); - } -} diff --git a/src/models/cua_s1/native/src/server.rs b/src/models/cua_s1/native/src/server.rs deleted file mode 100644 index 8dc87e12..00000000 --- a/src/models/cua_s1/native/src/server.rs +++ /dev/null @@ -1,227 +0,0 @@ -//! HTTP routes, matching the Python worker: `GET /health` and `POST /v1/systemone`, -//! one decision at a time, with the same status codes and response format. - -use std::sync::Arc; - -use axum::Router; -use axum::body::Body; -use axum::extract::State; -use axum::http::{HeaderMap, StatusCode, header}; -use axum::response::{IntoResponse, Response}; -use axum::routing::{get, post}; -use http_body_util::BodyExt; - -use crate::contract::{self, Request, RequestError, detail_json, map_request, parse_body}; -use crate::engine::Engine; -use crate::pyjson::{PyStr, repr_str, write_json_str}; - -pub struct Limits { - pub max_body_bytes: usize, - pub max_questions: usize, - /// per question; 0 disables the check - pub max_prompt_tokens: usize, -} - -pub struct App { - pub engine: Engine, - pub limits: Limits, - /// `Bearer ` as raw bytes, when a key is set - expected_auth: Option>, - pub identity: String, - /// held for the whole decision, so forward passes run one at a time - turn: tokio::sync::Mutex<()>, -} - -impl App { - /// `api_key` is the raw value of `CUA_S1_API_KEY`; empty means no key. - pub fn new(engine: Engine, limits: Limits, api_key: Option>, revision: &str) -> Self { - let expected_auth = api_key - .filter(|k| !k.is_empty()) - .map(|k| [b"Bearer ".as_slice(), &k].concat()); - Self { - engine, - limits, - expected_auth, - identity: contract::model_identity(revision), - turn: tokio::sync::Mutex::new(()), - } - } -} - -fn json_response(status: StatusCode, body: String) -> Response { - (status, [(header::CONTENT_TYPE, "application/json")], body).into_response() -} - -fn error(status: u16, message: &str) -> Response { - json_response( - StatusCode::from_u16(status).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), - detail_json(message), - ) -} - -enum DecideError { - Request(RequestError), - Internal(anyhow::Error), -} - -/// Score each question and build the response body. Every prompt is tokenized and -/// checked against the prompt limit before any forward pass runs. -async fn decide(app: &App, request: &Request) -> Result { - let limit = app.limits.max_prompt_tokens; - let mut encoded = Vec::with_capacity(request.questions.len()); - for question in &request.questions { - let ids = app - .engine - .prompter - .encode(&request.state, question) - .map_err(DecideError::Internal)?; - if limit > 0 && ids.len() > limit { - return Err(DecideError::Request(RequestError::new( - 413, - format!( - "question {}: prompt is {} tokens, over the {limit}-token limit", - repr_str(&PyStr::new(&question.name)), - ids.len() - ), - ))); - } - encoded.push(ids); - } - let mut answers = String::from("{"); - let mut prompt_tokens = 0; - for (i, (question, ids)) in request.questions.iter().zip(encoded).enumerate() { - prompt_tokens += ids.len(); - let probs = app - .engine - .score(ids, question.keys.len()) - .await - .map_err(DecideError::Internal)?; - let answer = contract::answer_json(question, &probs) - .map_err(|e| DecideError::Internal(anyhow::anyhow!(e)))?; - if i > 0 { - answers.push(','); - } - write_json_str(&question.name, &mut answers); - answers.push(':'); - answers.push_str(&answer); - } - answers.push('}'); - let mut out = String::from("{\"model\":"); - write_json_str(&app.identity, &mut out); - out.push_str(",\"answers\":"); - out.push_str(&answers); - out.push_str(&format!( - ",\"usage\":{{\"input_tokens\":{prompt_tokens},\"output_tokens\":0}}}}" - )); - Ok(out) -} - -/// The Python worker reads the header as Latin-1 text (Starlette) and encodes it back -/// as UTF-8 before `hmac.compare_digest`; the same bytes are compared here, so both -/// workers accept and reject the same headers. -fn authorized(app: &App, headers: &HeaderMap) -> bool { - let Some(expected) = &app.expected_auth else { - return true; - }; - let raw = headers - .get(header::AUTHORIZATION) - .map(|v| v.as_bytes()) - .unwrap_or(b""); - let mut supplied = Vec::with_capacity(raw.len()); - for &b in raw { - if b < 0x80 { - supplied.push(b); - } else { - supplied.extend_from_slice(&[0xc0 | (b >> 6), 0x80 | (b & 0x3f)]); - } - } - supplied.len() == expected.len() - && supplied - .iter() - .zip(expected) - .fold(0u8, |acc, (a, b)| acc | (a ^ b)) - == 0 -} - -async fn health(State(app): State>) -> Response { - let mut out = String::from("{\"status\":\"ready\",\"modality\":\"text\",\"model\":"); - write_json_str(&app.identity, &mut out); - out.push_str(",\"device\":"); - write_json_str(&app.engine.device, &mut out); - out.push_str(",\"dtype\":\"bfloat16\",\"mode\":\"native\"}"); - json_response(StatusCode::OK, out) -} - -async fn systemone(State(app): State>, headers: HeaderMap, body: Body) -> Response { - if !authorized(&app, &headers) { - return error(401, "invalid or missing bearer token"); - } - let max = app.limits.max_body_bytes; - if let Some(len) = headers - .get(header::CONTENT_LENGTH) - .and_then(|v| v.to_str().ok()) - && !len.is_empty() - && len.bytes().all(|b| b.is_ascii_digit()) - && len.parse::().map_or(true, |n| n > max as u128) - { - return error(413, "request body too large"); - } - let mut raw = Vec::new(); - let mut body = body; - while let Some(frame) = body.frame().await { - let Ok(frame) = frame else { - return error(400, "request body could not be read"); - }; - if let Some(chunk) = frame.data_ref() { - raw.extend_from_slice(chunk); - if raw.len() > max { - return error(413, "request body too large"); - } - } - } - let request = match parse_body(&raw).and_then(|b| map_request(&b, app.limits.max_questions)) { - Ok(r) => r, - Err(e) => return error(e.status, &e.message), - }; - let _turn = app.turn.lock().await; - match decide(&app, &request).await { - Ok(body) => json_response(StatusCode::OK, body), - Err(DecideError::Request(e)) => error(e.status, &e.message), - Err(DecideError::Internal(e)) => { - eprintln!("inference failed: {e:#}"); - error(500, "inference failed") - } - } -} - -async fn not_found() -> Response { - error(404, "Not Found") -} - -async fn method_not_allowed() -> Response { - error(405, "Method Not Allowed") -} - -pub fn router(app: Arc) -> Router { - Router::new() - .route("/health", get(health)) - .route("/v1/systemone", post(systemone)) - .fallback(not_found) - .method_not_allowed_fallback(method_not_allowed) - .with_state(app) -} - -/// One decision through the whole request path, before the server listens. -pub async fn warmup(app: &App) -> anyhow::Result<()> { - let request = map_request( - &parse_body(contract::WARMUP_BODY.as_bytes()).map_err(|e| anyhow::anyhow!(e.message))?, - 64, - ) - .map_err(|e| anyhow::anyhow!(e.message))?; - let _turn = app.turn.lock().await; - match decide(app, &request).await { - Ok(_) => Ok(()), - Err(DecideError::Request(e)) => anyhow::bail!(e.message), - Err(DecideError::Internal(e)) => Err(e), - } -} diff --git a/src/models/cua_s1/native/tests/kernels.rs b/src/models/cua_s1/native/tests/kernels.rs index 6b373232..cd044c83 100644 --- a/src/models/cua_s1/native/tests/kernels.rs +++ b/src/models/cua_s1/native/tests/kernels.rs @@ -1,5 +1,4 @@ -//! GPU checks of the attention and Gated DeltaNet kernels on random inputs, and of the -//! GEMM plan import. They need +//! GPU checks of the attention and Gated DeltaNet kernels on random inputs. They need //! a GPU and CUA_S1_CUDA_LIB pointing at libqwen3_5_cuda.so, so they only run when //! asked for: //! @@ -9,7 +8,7 @@ use std::path::PathBuf; use half::bf16; -use omni_cua_s1_native::cuda::{self, DeviceBuffer, GemmPlan, Stream, api, check}; +use omni_cua_s1_native::cuda::{self, DeviceBuffer, Stream, api, check}; fn setup() -> Stream { let lib = std::env::var_os("CUA_S1_CUDA_LIB") @@ -62,7 +61,7 @@ fn from_device(buf: &DeviceBuffer, n: usize, st: Stream) -> Vec { #[test] #[ignore = "needs a GPU and CUA_S1_CUDA_LIB"] -fn flash_attention_matches_float32_kernel() { +fn flash_attention_matches_float64_reference() { let st = setup(); let (hq, hk, dh) = (16usize, 4usize, 256usize); for (t, amp) in [ @@ -73,68 +72,70 @@ fn flash_attention_matches_float32_kernel() { (700, 8.0), (2048, 0.5), ] { - let q = to_device(&random(t * hq * dh, 1, amp), st); - let k = to_device(&random(t * hk * dh, 2, amp), st); + let (qh, kh) = (random(t * hq * dh, 1, amp), random(t * hk * dh, 2, amp)); // v is read in place from the q|k|v projection output, rows of 10240 as in the model let (ldv, v_at) = (10240usize, (hq * 2 + hk) * dh); - let qkv = to_device(&random(t * ldv, 3, 1.0), st); - let v = qkv.at(v_at * 2); - let flash = DeviceBuffer::new(t * hq * dh * 2).unwrap(); - let simple = DeviceBuffer::new(t * hq * dh * 2).unwrap(); - let (ti, hqi, hki, dhi, ldv) = (t as i32, hq as i32, hk as i32, dh as i32, ldv as i32); + let qkvh = random(t * ldv, 3, 1.0); + let (q, k, qkv) = (to_device(&qh, st), to_device(&kh, st), to_device(&qkvh, st)); + let out = DeviceBuffer::new(t * hq * dh * 2).unwrap(); // SAFETY: every buffer holds t rows of the given widths. - unsafe { - check( - (api().cs1_attention)( - q.at(0), - k.at(0), - v, - ldv, - flash.at(0), - ti, - hqi, - hki, - dhi, - 0.0625, - st, - ), - "flash", + let code = unsafe { + (api().cs1_attention)( + q.at(0), + k.at(0), + qkv.at(v_at * 2), + ldv as i32, + out.at(0), + t as i32, + hq as i32, + hk as i32, + dh as i32, + 0.0625, + st, ) - .unwrap(); - check( - (api().cs1_attention_simple)( - q.at(0), - k.at(0), - v, - ldv, - simple.at(0), - ti, - hqi, - hki, - dhi, - 0.0625, - st, - ), - "simple", - ) - .unwrap(); - } - let a = from_device(&flash, t * hq * dh, st); - let b = from_device(&simple, t * hq * dh, st); - // per (token, head): the largest difference over the largest magnitude - let mut worst = 0f32; - for (ra, rb) in a.chunks_exact(dh).zip(b.chunks_exact(dh)) { - let d = ra - .iter() - .zip(rb) - .map(|(x, y)| (x - y).abs()) - .fold(0f32, f32::max); - let m = rb.iter().map(|y| y.abs()).fold(1e-3f32, f32::max); - assert!( - ra.iter().all(|x| x.is_finite()), - "non-finite output at t = {t}" - ); - worst = worst.max(d / m); + }; + check(code, "attention").unwrap(); + let got = from_device(&out, t * hq * dh, st); + assert!( + got.iter().all(|x| x.is_finite()), + "non-finite output at t = {t}" + ); + // about 64 query rows per length, each against causal attention in float64; + // per (row, head): the largest difference over the largest magnitude + let mut worst = 0f64; + for i in (0..t).step_by(t.div_ceil(64)).chain([t - 1]) { + for h in 0..hq { + let g = h / (hq / hk); + let qi = &qh[(i * hq + h) * dh..][..dh]; + let s: Vec = (0..=i) + .map(|j| { + let kj = &kh[(j * hk + g) * dh..][..dh]; + qi.iter() + .zip(kj) + .map(|(a, b)| a.to_f64() * b.to_f64()) + .sum::() + * 0.0625 + }) + .collect(); + let m = s.iter().copied().fold(f64::NEG_INFINITY, f64::max); + let w: Vec = s.iter().map(|x| (x - m).exp()).collect(); + let z: f64 = w.iter().sum(); + let want: Vec = (0..dh) + .map(|d| { + (0..=i) + .map(|j| w[j] * qkvh[j * ldv + v_at + g * dh + d].to_f64()) + .sum::() + / z + }) + .collect(); + let row = &got[(i * hq + h) * dh..][..dh]; + let diff = row + .iter() + .zip(&want) + .map(|(&x, y)| (x as f64 - y).abs()) + .fold(0f64, f64::max); + worst = worst.max(diff / want.iter().map(|y| y.abs()).fold(1e-3, f64::max)); + } } eprintln!("attention t = {t}, amplitude {amp}: largest relative difference {worst:.2e}"); assert!(worst < 1.6e-2, "t = {t}: {worst}"); @@ -261,46 +262,3 @@ fn gated_delta_rule_matches_recurrent_reference() { assert!(worst <= 2e-2 * scale, "t = {t}: {worst} vs scale {scale}"); } } - -/// A tuned GEMM plan moves to another cuBLASLt handle through export and import, and a -/// plan that says it was tuned with another cuBLASLt version is refused without -/// changing anything. -#[test] -#[ignore = "needs a GPU and CUA_S1_CUDA_LIB"] -fn gemm_plans_import_only_for_their_cublaslt_version() { - let st = setup(); - let (m, n, k) = (64i32, 256i32, 512i32); - let x = to_device(&random((m * k) as usize, 21, 1.0), st); - let w = to_device(&random((n * k) as usize, 22, 1.0), st); - let y = DeviceBuffer::new((m * n * 2) as usize).unwrap(); - // SAFETY: the buffers hold m x k, n x k and m x n values; the handles are destroyed - // at the end and not used after. - unsafe { - let api = api(); - let (a, b) = ( - (api.cs1_gemm_create)(32 << 20), - (api.cs1_gemm_create)(32 << 20), - ); - assert!(!a.is_null() && !b.is_null()); - check( - (api.cs1_gemm_tune)(a, x.at(0), w.at(0), y.at(0), m, n, k, n, 0, st), - "tune", - ) - .unwrap(); - (api.cs1_gemm_tune_done)(a); - let count = (api.cs1_gemm_export)(a, std::ptr::null_mut(), 0); - assert_eq!(count, 1); - let mut plans = vec![GemmPlan::default(); count]; - (api.cs1_gemm_export)(a, plans.as_mut_ptr(), count); - assert_eq!(plans[0].cublaslt_version, (api.cs1_gemm_version)() as u64); - - let mut other = plans.clone(); - other[0].cublaslt_version += 1; - assert_ne!((api.cs1_gemm_import)(b, other.as_ptr(), 1), 0); - assert_eq!((api.cs1_gemm_export)(b, std::ptr::null_mut(), 0), 0); - check((api.cs1_gemm_import)(b, plans.as_ptr(), 1), "import").unwrap(); - assert_eq!((api.cs1_gemm_export)(b, std::ptr::null_mut(), 0), 1); - (api.cs1_gemm_destroy)(a); - (api.cs1_gemm_destroy)(b); - } -}