diff --git a/Cargo.lock b/Cargo.lock index 11b8ba2..ba292dd 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -452,6 +452,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b1f9e3d69d39e4862ffed03ed071a76f9a13ba1d9109d355b0f0aa6b15e393c4" dependencies = [ "futures-core", + "futures-sink", ] [[package]] @@ -974,6 +975,20 @@ dependencies = [ "minimal-lexical", ] +[[package]] +name = "omni-clm" +version = "0.1.0" +dependencies = [ + "anyhow", + "base64 0.22.1", + "memmap2", + "reqwest", + "safetensors 0.6.2", + "serde", + "serde_json", + "sha2", +] + [[package]] name = "omni-cua-s1-native" version = "0.1.0" @@ -1346,6 +1361,7 @@ checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147" dependencies = [ "base64 0.22.1", "bytes", + "futures-channel", "futures-core", "futures-util", "http", diff --git a/Cargo.toml b/Cargo.toml index d5abbe6..242d51c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,3 +1,3 @@ [workspace] -members = ["src/frontend", "src/runtime", "src/models/cua_s1/native", "src/models/qwen3_5/native", "src/models/open_jev/native", "src/models/laya", "src/backends/cuda"] +members = ["src/frontend", "src/runtime", "src/models/clm", "src/models/cua_s1/native", "src/models/qwen3_5/native", "src/models/open_jev/native", "src/models/laya", "src/backends/cuda"] resolver = "3" diff --git a/docs/supported-models.md b/docs/supported-models.md index f9a8539..db8ade2 100644 --- a/docs/supported-models.md +++ b/docs/supported-models.md @@ -15,16 +15,17 @@ Models that are being added are also tracked in issues labeled [new model](https | Cua-S1 4B 0.2, `text` adapter | [Native Rust worker](../recipe/cua_s1/native.md) on the [Qwen3.5 CUDA kernels](../src/backends/cuda/qwen3_5/README.md) | Not supported | Validated on compute capability 8.9 ([#19](https://github.com/ThinkFlowLab/system1-omni/pull/19), [#52](https://github.com/ThinkFlowLab/system1-omni/pull/52)) | Not supported | Compute capability 8.0 or newer, the CUDA toolkit to build, weights merged with `export_text_merged.py` | | Cua-S1 4B 0.2, `multimodal` adapter | Reference worker on Transformers and PEFT, [`src/frontend/cua_s1.py`](../src/frontend/cua_s1.py); no recipe yet | Not supported | Validated ([#17](https://github.com/ThinkFlowLab/system1-omni/pull/17), [#18](https://github.com/ThinkFlowLab/system1-omni/pull/18)) | Not supported | The state is one PNG or JPEG image; upstream's `weights.lock.json` next to the base weights | | Open-Jev-27B-v1.1 | [Native Rust/CUDA worker](../recipe/open_jev/native.md) on the shared Qwen3.5/3.8 executor | Not supported | Validated on H200 (sm_90) for the [74 single-candidate workload](../recipe/open_jev/validation.md) | Not supported | Compute capability 8.0 or newer, CUDA toolkit to build, exported merged weights and trained head | -| CLM-v0.1-8B | [External `clm-serve` recipe](../recipe/clm/README.md) with a CPU stub embeddings server | **Stub-encoder contract checks only** ([#23](https://github.com/ThinkFlowLab/system1-omni/pull/23)); not real Qwen3-8B decisions | Real encoder unverified by the merged recipe | Unverified | Python, upstream CLM and head checkpoint; a real encoder requires a separate embeddings server | +| CLM-v0.1-8B | [Native Rust engine](../src/models/clm/README.md), `clm-run`, in front of a Qwen3-8B `/v1/embeddings` server; [CLM's own `clm-serve` recipe](../recipe/clm/README.md) is the stub-encoder path | Engine-only agreement **4.5e-06** on identical vectors ([#29](https://github.com/ThinkFlowLab/system1-omni/pull/29), [validation](../recipe/clm/native/VALIDATION.md)); contract checks documented ([#23](https://github.com/ThinkFlowLab/system1-omni/pull/23)) | **Validated on compute capability 8.9** ([#29](https://github.com/ThinkFlowLab/system1-omni/pull/29), [validation](../recipe/clm/native/VALIDATION.md)) | Unverified | Python, upstream CLM and head checkpoint for the recipe; the native engine needs a separate `/v1/embeddings` server and Rust | - **Validated:** covered by the recipe on `main` or by the checks in the linked merged pull request. - **Unverified:** the worker accepts this device, but no recipe or merged pull request covers it. - **Planned:** not implemented yet; the linked issue tracks it. The Cua-S1 workers answer `choice` questions only. -LAYA's English worker and Open-Jev support `choice`, `score`, and `noul` text questions. -CLM's merged recipe exercises these answer shapes with stub embeddings; it does -not validate decision quality. MPS validation above is for a Python/PyTorch +LAYA's English worker, Open-Jev and CLM support `choice`, `score`, and `noul` text questions. +CLM's CUDA row is the path through a real Qwen3-8B encoder; its engine-only figure comes +from the same comparison run against the deterministic stub, where both sides get identical +vectors. Neither number validates decision quality. MPS validation above is for a Python/PyTorch worker, not a native Metal backend. The [architecture contracts](architecture.md) describe the native target. diff --git a/recipe/clm/README.md b/recipe/clm/README.md index f0f7726..39b16dc 100644 --- a/recipe/clm/README.md +++ b/recipe/clm/README.md @@ -90,3 +90,7 @@ GPU=0 PORT=8090 UTIL=0.35 ./serve_qwen3_8b.sh # from the CLM checkout; needs ``` Everything downstream is unchanged, which is the property this recipe is meant to demonstrate. + +`recipe/clm/native/VALIDATION.md` records what that path measures when it is run against the +native engine instead of `clm-serve`: the engine alone on identical vectors, and the whole +path on a real Qwen3-8B. diff --git a/recipe/clm/native/VALIDATION.md b/recipe/clm/native/VALIDATION.md new file mode 100644 index 0000000..61b6437 --- /dev/null +++ b/recipe/clm/native/VALIDATION.md @@ -0,0 +1,121 @@ +# CLM CUDA validation + +`omni-clm` owns everything after the encoder: a frozen Qwen3-8B runs as its own process +behind an `/v1/embeddings` endpoint, and the engine owns the two projection heads, the +cosine score, the temperature and the typed answer. Agreement with the reference is +therefore two separate questions — whether the engine's own arithmetic matches, and whether +the whole path still agrees once a real encoder is in front of it — and they are measured +separately below, because one number cannot answer both. + +Both phases ran on one RTX 4090 on 2026-10-06. Every number below is the output of +`compare_with_reference.py`, which sends the same requests to `clm-run` and to CLM's own +`Engine` and `Schema` and compares the answers. + +## The environment + +| | | +| --- | --- | +| GPU | NVIDIA GeForce RTX 4090, 24 GB, compute capability 8.9 | +| Driver | 595.71.05 | +| CUDA | 13.0, V13.0.88 | +| Python | 3.12.3 | +| torch | 2.12.1+cu130 | +| transformers | 5.17.0, the encoder | +| contrastive-lm | 0.1.0, the reference | +| Rust | 1.98.1 | +| Engine commit | `a2019f9` | +| `clm-run` | sha256 `791859b51047edd79ad6f380bba9a78ea1fdddb5c4a8ad797472ac541251dc92` | + +## Phase A: the engine alone + +`recipe/clm/stub_embedder.py` derives its vectors from the text alone, and both sides ask +it, so the same text reaches both as the same vector. The encoder's numerical difference is +out of the comparison; what remains is parsing, the text the heads see, the projections, the +cosine, the temperature and the answer shape — and, because a difference in the text the two +sides send would produce a different vector, the text is in the comparison too. + +| case | kind | agreement | +| --- | --- | --- | +| `choice_two` | choice | `max\|dp\|=3.12e-07`, both pick `billing` | +| `choice_five` | choice | `max\|dp\|=4.49e-06`, both pick `c` | +| `score_three` | score | `max\|dp\|=1.13e-06`, `1.132559` vs `1.132558` | +| `noul_stmt` | noul | `0.271456` vs `0.271451` | +| `noul_numbers` | noul | `0.225941` vs `0.225943` | + +Worst case **4.49e-06**. Nothing here is passed a tolerance to hide behind: the answers are +compared at the 0.025 the script uses for the real encoder, and they land more than three +orders of magnitude inside it. + +## Phase B: end to end, on Qwen3-8B + +`recipe/clm/native/transformers_encoder.py` serves Qwen3-8B — the model the deployment +encodes with — behind the same endpoint, so vLLM is not needed to check the client. + +| case | kind | agreement | +| --- | --- | --- | +| `choice_two` | choice | `max\|dp\|=1.72e-05`, both pick `billing` | +| `choice_five` | choice | `max\|dp\|=1.97e-02`, both pick `a` | +| `score_three` | score | `max\|dp\|=1.08e-02`, `0.385166` vs `0.403241` | +| `noul_stmt` | noul | `0.631038` vs `0.646805` | +| `noul_numbers` | noul | `0.323871` vs `0.321280` | + +Worst case **1.97e-02**, against the script's tolerance of 0.025, and all five cases agree. + +## Why the two numbers differ, and why the tolerance is 0.025 + +CLM scores with `softmax(exp(logit_scale) * cos(...) / temperature)`, and the published +checkpoint's `logit_scale` is 4.6132, whose exponential the engine caps at 100. A difference +of `d` in a cosine therefore becomes `100 d` in the logit, and two forward passes that +disagree by ~1e-4 in the cosine — which is what two encoder implementations produce for +Qwen3-8B — disagree by ~1e-2 in the probabilities. + +Phase A is what makes that attribution a measurement rather than a story: with the encoder +taken out, the engine agrees to 4.49e-06, so the ~1e-2 in phase B is the encoder and not the +engine. A deployment serving the same Qwen3-8B weights would be on the other side of that +difference, which is a property of how this model scores, not a defect on either side. + +## Artifacts + +| file | sha256 | +| --- | --- | +| `CLM_v0.1-8B.pt` | `b2b4a8c9c2d39263eff78a351eb909a342ce9b3bf21a3f07c1d1bf15f1c4eda5` | +| `clm-export/oracle.json` | `7d058de4a691b6d4949aa51c53c7ec95a136ac4c14290580e2d81ba8f7192c16` | +| Qwen3-8B `config.json` | `f7c4eadfbbf522470667b797a3c89be2524832d2d599797248dc304fff447c30` | + +`clm-export/model.safetensors` is deliberately not pinned. The safetensors serializer +carries `__metadata__` through a hash map, so its key order — and with it the file hash — +differs from run to run: four exports of this same checkpoint gave four file hashes while +their tensor entries were byte-identical and every one of them wrote the same `oracle.json`. +`oracle.json` holds the FP32/FP16/BF16 hash of every tensor and is stable, byte-identical +across two machines, so that is what pins the conversion. + +## Reproduction + +Both phases take the same comparison; only the endpoint changes. + +```sh +# the checkpoint, once +python recipe/clm/native/export_weights.py CLM_v0.1-8B.pt clm-export + +# Phase A, no GPU and no 8B encoder +python recipe/clm/stub_embedder.py --port 8090 & + +# Phase B, instead of the stub, on a GPU +python recipe/clm/native/transformers_encoder.py --model /path/to/Qwen3-8B --port 8090 & + +python recipe/clm/native/compare_with_reference.py \ + --checkpoint clm-export --bin target/release/clm-run \ + --emb-url http://127.0.0.1:8090/v1/embeddings \ + --pt CLM_v0.1-8B.pt --temperature 1.0 +``` + +## What this does not establish + +- **Decision quality.** Five fixed requests show the port is faithful; they are not a + labelled set and say nothing about how often CLM is right. +- **vLLM.** The encoder here is Transformers. A deployment's pooling server is a different + implementation and sits on the other side of the ~1e-2 above; nothing here measures it. +- **Latency or throughput.** Nothing was timed. The engine's candidate-vector cache is + exercised but not measured. +- **Other hardware.** One RTX 4090, compute capability 8.9. The engine is Rust and has no + device-specific code; the CUDA part of this path is `transformers_encoder.py`. diff --git a/recipe/clm/native/compare_with_reference.py b/recipe/clm/native/compare_with_reference.py new file mode 100644 index 0000000..32a9652 --- /dev/null +++ b/recipe/clm/native/compare_with_reference.py @@ -0,0 +1,154 @@ +"""Compare omni-clm's decision with the CLM reference, over the same embeddings. + +Both sides call the same OpenAI-compatible `/v1/embeddings` endpoint, so the vectors are +identical and the only difference left is the code between the vectors and the answer. + + python recipe/clm/native/compare_with_reference.py \ + --checkpoint /path/to/clm-export \ + --bin /path/to/clm-run \ + --emb-url http://127.0.0.1:8090/v1/embeddings \ + --pt /path/to/CLM_v0.1-8B.pt + +## On the tolerance + +CLM scores with `exp(logit_scale) = 100`, so a difference of `d` in a cosine similarity +becomes `100 d` in the logit. Two encoder forward passes that disagree by ~1e-4 in the +cosine therefore disagree by ~1e-2 in the probabilities, and that is what two different +implementations of Qwen3-8B produce -- vLLM's bf16 pooling against Transformers' forward. +It is a property of how this model scores, not a defect on either side. + +The default tolerance reflects it. The engine's own arithmetic is checked separately and +much more tightly: with the vectors held fixed the two implementations agree to 3e-06. +""" + +from __future__ import annotations + +import argparse +import json +import subprocess +import sys +from pathlib import Path + +CASES = [ + ("choice_two", "choice", {"billing": "Charges and refunds", "technical": "Software problems"}), + ("choice_five", "choice", {"a": "alpha", "b": "beta", "c": "gamma", "d": "delta", "e": "epsilon"}), + ("score_three", "score", ["Not urgent", "Needs attention soon", "Needs attention immediately"]), + ("noul_stmt", "noul", None), + # A `noul` reads its two descriptions in `false`/`true` order, not in the order they + # were written, so a number here has to be found by the key it was written under. + # Counting positions gave `false: 2` / `true: 18446744073709551616` reversed, which + # is a different pair of candidate texts and so a different answer. + ("noul_numbers", "noul", {"true": 18446744073709551616, "false": 2}), +] + +STATE = "I was charged twice for order 4411 and want the second charge refunded." + +# exp(logit_scale) with the published checkpoint, and what it does to a cosine difference. +SCALE = 100.0 +TOLERANCE = 2.5e-2 + + +def request_for(name: str, kind: str, criteria) -> dict: + q = {"type": kind, "instructions": f"Question {name}"} + if criteria is not None: + q["criteria"] = criteria + return {"model": "clm-latest", "state": STATE, "questions": {name: q}} + + +def reference(emb_url: str, emb_model: str, pt: Path, req: dict, temperature: float) -> dict: + """CLM's own heads and schema, applied to the same embeddings the Rust side gets.""" + from clm.embedder import Embedder + from clm.engine import Engine + + embedder = Embedder(url=emb_url, model=emb_model) + engine = Engine(embedder=embedder, checkpoint=str(pt)) + return engine.answer(req["state"], req["questions"], "clm-latest", temperature) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("--checkpoint", type=Path, required=True, help="converted export dir") + parser.add_argument("--bin", type=Path, required=True, help="clm-run") + parser.add_argument("--emb-url", required=True) + parser.add_argument("--emb-model", default="qwen3-8b") + parser.add_argument("--pt", type=Path, required=True, help="CLM_v0.1-8B.pt") + parser.add_argument("--temperature", type=float, default=1.0) + parser.add_argument( + "--tolerance", + type=float, + default=TOLERANCE, + help=f"logits are scale*cos with scale={SCALE:g}, so two encoder implementations " + f"differ by about this much in the probabilities", + ) + args = parser.parse_args() + + requests_ = [(n, request_for(n, k, c)) for n, k, c in CASES] + payload = "".join(json.dumps(r) + "\n" for _, r in requests_) + proc = subprocess.run( + [ + str(args.bin), + str(args.checkpoint), + "--emb-url", + args.emb_url, + "--model", + args.emb_model, + # Both sides must run the same temperature, or a non-default one compares two + # different configurations and reports a port regression that is not there. + "--temperature", + str(args.temperature), + ], + input=payload, + capture_output=True, + text=True, + timeout=600, + ) + if proc.returncode != 0: + print(proc.stderr[-2000:], file=sys.stderr) + raise SystemExit(f"clm-run exited {proc.returncode}") + mine = [json.loads(line) for line in proc.stdout.splitlines() if line.strip()] + if len(mine) != len(requests_): + raise SystemExit(f"clm-run answered {len(mine)} of {len(requests_)}") + + failures = 0 + for (name, req), got in zip(requests_, mine): + want = reference(args.emb_url, args.emb_model, args.pt, req, args.temperature) + mine_p = got["answers"][name] + want_p = want["answers"][name] + kind = req["questions"][name]["type"] + + if kind == "noul": + got_v, want_v = mine_p["noul"], want_p["noul"] + ok = abs(got_v - want_v) <= args.tolerance + label, detail = "noul", f"{got_v:.6f} vs {want_v:.6f}" + else: + gk, wk = mine_p["probabilities"], want_p["probabilities"] + if set(gk) != set(wk): + print(f"FAIL {name}: keys differ {sorted(gk)} vs {sorted(wk)}") + failures += 1 + continue + worst = max(abs(gk[k] - wk[k]) for k in gk) + same_pick = ( + mine_p.get("choice") == want_p.get("choice") + if kind == "choice" + else abs(mine_p["score"] - want_p["score"]) <= args.tolerance + ) + ok = worst <= args.tolerance and same_pick + label = "choice" if kind == "choice" else "score" + detail = f"max|dp|={worst:.2e} " + ( + f"choice={mine_p['choice']}/{want_p['choice']}" + if kind == "choice" + else f"score={mine_p['score']:.6f}/{want_p['score']:.6f}" + ) + failures += not ok + print(f"{'PASS' if ok else 'FAIL'} {name:12} {label:7} {detail}") + + print( + f"\n{'all cases agree' if not failures else str(failures) + ' FAILED'} " + f"(tolerance {args.tolerance:g}; point --emb-url at recipe/clm/stub_embedder.py " + f"to measure the engine alone, on identical vectors)" + ) + raise SystemExit(1 if failures else 0) + + +if __name__ == "__main__": + main() diff --git a/recipe/clm/native/export_weights.py b/recipe/clm/native/export_weights.py new file mode 100644 index 0000000..ab1215f --- /dev/null +++ b/recipe/clm/native/export_weights.py @@ -0,0 +1,87 @@ +"""Export a CLM head checkpoint to safetensors, and record the conversion oracle. + +A CLM checkpoint is a ``torch.save`` dict (``state_head``/``action_head`` state dicts, +``logit_scale``, ``cfg``), so it is a pickle and no non-Python reader can open it. This +writes the tensors to safetensors with the head name as a prefix, keeps the scalar and +config entries in the safetensors metadata, and emits the FP32/FP16/BF16 hash of every +tensor so a reader can be checked without comparing floats directly. + + python recipe/clm/native/export_weights.py CLM_v0.1-8B.pt OUT_DIR + python recipe/clm/native/export_weights.py CLM_v0.1-8B.pt OUT_DIR --oracle oracle.json + +Writes ``model.safetensors`` and, unless ``--no-oracle``, ``oracle.json``. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +from pathlib import Path + +import torch +from safetensors.torch import save_file + +HEADS = ("state_head", "action_head") + + +def tensors(ckpt: dict) -> dict[str, torch.Tensor]: + """Every parameter, prefixed by its head, in a stable order.""" + out: dict[str, torch.Tensor] = {} + for head in HEADS: + state = ckpt.get(head) + if not isinstance(state, dict): + raise SystemExit(f"checkpoint has no {head!r} state dict") + for name, value in state.items(): + if not isinstance(value, torch.Tensor): + raise SystemExit(f"{head}.{name} is {type(value).__name__}, not a tensor") + out[f"{head}.{name}"] = value.detach().to(torch.float32).contiguous() + return out + + +def oracle(weights: dict[str, torch.Tensor]) -> list[dict]: + rows = [] + for name, x in weights.items(): + row: dict = {"name": name, "shape": list(x.shape), "source_dtype": str(x.dtype)} + for key, dtype in (("f32", torch.float32), ("f16", torch.float16), ("bf16", torch.bfloat16)): + y = x.to(torch.float32).to(dtype).contiguous() + row[key] = hashlib.sha256(y.view(torch.uint8).numpy().tobytes()).hexdigest() + rows.append(row) + return rows + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("checkpoint", type=Path, help="CLM_v0.1-8B.pt") + parser.add_argument("output", type=Path, help="directory to write into") + parser.add_argument("--oracle", type=Path, help="where to write the conversion oracle") + parser.add_argument("--no-oracle", action="store_true", help="skip the oracle") + args = parser.parse_args() + + torch.set_num_threads(4) + ckpt = torch.load(args.checkpoint, map_location="cpu", weights_only=False) + cfg = dict(ckpt.get("cfg") or {}) + weights = tensors(ckpt) + args.output.mkdir(parents=True, exist_ok=True) + + metadata = { + "format": "clm-heads", + "logit_scale": repr(float(ckpt["logit_scale"])), + "hidden_size": str(int(ckpt.get("hidden_size") or cfg.get("hidden_size"))), + "projection_dim": str(int(ckpt.get("projection_dim") or cfg.get("projection_dim"))), + "cfg": json.dumps(cfg, sort_keys=True), + } + path = args.output / "model.safetensors" + save_file(weights, str(path), metadata=metadata) + + params = sum(v.numel() for v in weights.values()) + print(f"WROTE {path} tensors={len(weights)} params={params}", flush=True) + + if not args.no_oracle: + where = args.oracle or (args.output / "oracle.json") + where.write_text(json.dumps(oracle(weights), indent=2) + "\n") + print(f"ORACLE {where} rows={len(weights)}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/recipe/clm/native/head_oracle.py b/recipe/clm/native/head_oracle.py new file mode 100644 index 0000000..38c865d --- /dev/null +++ b/recipe/clm/native/head_oracle.py @@ -0,0 +1,162 @@ +"""Reference decisions for the CLM heads, for the Rust port to be checked against. + +Reads the exported safetensors (not the .pt) so both sides load byte-identical weights, +computes a decision for each fixed case with plain NumPy, and writes the result as JSON. + + python recipe/clm/native/head_oracle.py EXPORT_DIR OUT.json + +The embeddings here are synthesised from the case name, so the numbers are meaningless +as model output -- the point is that two implementations of the same arithmetic agree. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +from pathlib import Path + +import numpy as np +from safetensors import safe_open + +HIDDEN = 4096 + + +def embedding(text: str, dim: int) -> np.ndarray: + """A stable unit vector per text, so the cases are reproducible without an encoder.""" + out: list[float] = [] + counter = 0 + while len(out) < dim: + digest = hashlib.sha256(f"{counter}:{text}".encode()).digest() + for i in range(0, len(digest) - 3, 4): + if len(out) == dim: + break + out.append(int.from_bytes(digest[i:i + 4], "big") / 2**31 - 1.0) + counter += 1 + x = np.asarray(out, dtype=np.float32) + return x / np.linalg.norm(x) + + +def load_head(f, prefix: str, width: int, proj: int, hidden: int, blocks: int) -> dict: + def get(name: str) -> np.ndarray: + return np.asarray(f.get_tensor(f"{prefix}.{name}"), dtype=np.float32) + + head = { + "inp_w": get("inp.weight"), "inp_b": get("inp.bias"), + "out_w": get("out.weight"), "out_b": get("out.bias"), + } + for i in range(blocks): + head[f"hidden{i}_w"] = get(f"hidden.{i}.weight") + head[f"hidden{i}_b"] = get(f"hidden.{i}.bias") + head[f"norm{i}_w"] = get(f"norms.{i}.weight") + head[f"norm{i}_b"] = get(f"norms.{i}.bias") + return head + + +def gelu(x: np.ndarray) -> np.ndarray: + # The erf form, which is what torch.nn.GELU() uses by default. + from math import sqrt + + return 0.5 * x * (1.0 + _erf(x / sqrt(2.0))) + + +def _erf(x: np.ndarray) -> np.ndarray: + """Vectorised erf via the same A&S 7.1.26 form the Rust side uses.""" + sign = np.sign(x) + x = np.abs(x) + t = 1.0 / (1.0 + 0.3275911 * x) + y = 1.0 - (((((1.0614054 * t - 1.45315203) * t) + 1.42141374) * t - 0.284496736) * t + + 0.2548296) * t * np.exp(-x * x) + return sign * y + + +def layernorm(x: np.ndarray, w: np.ndarray, b: np.ndarray, eps: float = 1e-5) -> np.ndarray: + mean = x.mean(axis=-1, keepdims=True) + var = x.var(axis=-1, keepdims=True) + return (x - mean) / np.sqrt(var + eps) * w + b + + +def project(head: dict, cfg: dict, x: np.ndarray) -> np.ndarray: + h = gelu(x @ head["inp_w"].T + head["inp_b"]) + blocks = cfg["depth"] - 2 + for i in range(blocks): + z = h @ head[f"hidden{i}_w"].T + head[f"hidden{i}_b"] + if cfg["layernorm"]: + z = layernorm(z, head[f"norm{i}_w"], head[f"norm{i}_b"]) + z = gelu(z) + h = h + z if cfg["residual"] else z + out = h @ head["out_w"].T + head["out_b"] + return out / np.linalg.norm(out, axis=-1, keepdims=True) + + +def softmax(v: np.ndarray) -> np.ndarray: + e = np.exp(v - v.max()) + return e / e.sum() + + +def confidence(probs: np.ndarray) -> float: + if probs.size < 2: + return 1.0 + j = int(np.argmax(probs)) + rest = np.delete(probs, j).mean() + return float(min(1.0, max(0.0, probs[j] - rest))) + + +CASES = [ + ("choice_binary", "choice", ["billing", "technical"]), + ("choice_five", "choice", ["a", "b", "c", "d", "e"]), + ("score_three", "score", ["0", "1", "2"]), + ("noul", "noul", ["false", "true"]), + ("choice_single", "choice", ["only"]), +] + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("export", type=Path, help="directory holding model.safetensors") + parser.add_argument("output", type=Path) + args = parser.parse_args() + + with safe_open(args.export / "model.safetensors", framework="np") as f: + md = f.metadata() + cfg = json.loads(md["cfg"]) + hidden = int(md["hidden_size"]) + proj = int(md["projection_dim"]) + width = cfg["width"] + blocks = cfg["depth"] - 2 + # The reference caps this: heads.py does exp(logit_scale).clamp(max=100.0), + # and the published checkpoint's 4.6132 exponentiates to 100.82, so the cap + # binds. Without it every probability is about 0.8 % off. + scale = min(float(np.exp(np.float32(md["logit_scale"]))), 100.0) + state = load_head(f, "state_head", width, proj, hidden, blocks) + action = load_head(f, "action_head", width, proj, hidden, blocks) + + rows = [] + for name, kind, keys in CASES: + state_vec = embedding(f"state::{name}", hidden) + cand_vecs = np.stack([embedding(f"cand::{name}::{k}", hidden) for k in keys]) + zs = project(state, cfg, state_vec[None, :])[0] + zc = project(action, cfg, cand_vecs) + cos = zc @ zs + temperature = 1.0 if name != "choice_five" else 2.5 + probs = softmax((np.float32(scale) * cos / np.float32(temperature)).astype(np.float32)) + row = { + "name": name, "kind": kind, "keys": keys, "temperature": temperature, + "probabilities": [float(p) for p in probs], + } + if kind == "choice": + row["choice"] = keys[int(np.argmax(probs))] + row["confidence"] = confidence(probs) + elif kind == "score": + row["score"] = float(sum(i * float(p) for i, p in enumerate(probs))) + row["confidence"] = confidence(probs) + else: + row["noul"] = float(probs[keys.index("true")]) + rows.append(row) + + args.output.write_text(json.dumps({"cases": rows}, indent=2) + "\n") + print(f"ORACLE {args.output} cases={len(rows)}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/recipe/clm/native/text_oracle.py b/recipe/clm/native/text_oracle.py new file mode 100644 index 0000000..11306bf --- /dev/null +++ b/recipe/clm/native/text_oracle.py @@ -0,0 +1,111 @@ +"""The text the CLM heads see, from the reference implementation itself. + +`omni-clm` reimplements `state_text`, `candidates` and `to_text`. A separator in the +wrong place does not fail loudly — it shifts every probability — so the Rust side is +checked byte-for-byte against this file, which is produced by importing `clm.schema` +rather than by transcribing it. + + python recipe/clm/native/text_oracle.py OUT.json +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +from clm.schema import build_pairs, candidates, state_text, to_text + +STATES = [ + "I was charged twice.", + {"body": "Charged twice", "order": 4411, "urgent": True}, + {"ticket": {"id": 7, "tags": ["a", "b"]}, "note": None}, + [{"k": 1}, {"k": 2}], + {"empty_obj": {}, "empty_arr": [], "n": 0.5}, + {"nested": {"deep": {"x": "y"}}}, + # `str(float)` switches to exponent form outside [1e-4, 1e16) and keeps a `.0` on an + # integral float, so these pin the number spelling the reference produces. + {"tiny": 1e-5, "smaller": 1e-7, "edge": 1e-4, "round": 1e15, "huge": 1e16, "neg": -1e-6}, + # 17 significant digits, where a parse that is not correctly rounded lands on + # the neighbouring double and renders differently. + {"seventeen": 7.8190461323667115, "inexact": 9007199254740993.0}, + # A JSON integer is an arbitrary-precision int in Python and keeps every digit; one + # larger than a u64 cannot survive a detour through a double. + {"big": 18446744073709551616, "huge": 340282366920938463463374607431768211456, + "negzero": -0, "negbig": -18446744073709551616}, +] + +QUESTIONS = [ + {"type": "choice", "instructions": "Which team?", + "criteria": {"billing": "Charges and refunds", "tech": "Software problems"}}, + {"type": "choice", "instructions": "Pick", "criteria": {"a": "", "b": None}}, + {"type": "score", "instructions": "How urgent?", "criteria": ["Not urgent", "Soon", "Now"]}, + {"type": "noul", "instructions": "Does the customer ask for a refund?", "criteria": None}, + {"type": "noul", "instructions": "Refund?", + "criteria": {"true": "Yes they do", "false": "No they do not"}}, + # An empty container is a description that renders to nothing, not a missing one, so + # these exercise the difference between "absent" and "renders empty". + {"type": "choice", "instructions": "Pick a bucket", + "criteria": {"empty_obj": {}, "empty_list": []}}, + {"type": "noul", "instructions": "Is it so?", "criteria": {"true": {}, "false": []}}, + # Numbers are spelled by `str(float)` in criteria too, not only in the state. + {"type": "score", "instructions": "How much?", "criteria": [1e-5, 0.5, 1e16]}, +] + +# Whole request bodies, as the text a caller sends them in. `QUESTIONS` above is built +# from Python values, so it cannot say `18446744073709551616` and mean an integer: it is +# the literal text that carries those digits. These go through the request path instead, +# which is the path a server takes and the only one the digits survive. +RAW_REQUESTS = [ + # A `noul` reads its two descriptions in `false`/`true` order, not in the order they + # were written, so a number has to be found by its key rather than by counting. + '{"state": {"n": 1}, "questions": {"q": {"type": "noul", "instructions": "Is it so?",' + ' "criteria": {"true": 1, "false": 2}}}}', + # A statement that is a number rather than a string. The default `noul` candidates are + # built from the statement, so it has to keep every digit there as well as in the + # state text. + '{"state": {}, "questions": {"q": {"type": "noul",' + ' "instructions": 18446744073709551616, "criteria": null}}}', + # The same digits in a state, in a `choice`'s descriptions and in a `score`'s levels, + # with the two questions reaching the same state text. + '{"state": {"big": 18446744073709551616, "seventeen": 7.8190461323667115},' + ' "questions": {' + '"a": {"type": "choice", "instructions": "Pick",' + ' "criteria": {"x": 340282366920938463463374607431768211456, "y": 7.8190461323667115}},' + ' "b": {"type": "score", "instructions": "How much?",' + ' "criteria": [1e-5, 18446744073709551616]}}}', +] + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("output", type=Path) + args = parser.parse_args() + + out = {"to_text": [to_text(s) for s in STATES], "cases": [], "raw_cases": []} + for state in STATES: + for q in QUESTIONS: + keys, texts = candidates(q) + out["cases"].append({ + "state": state, + "kind": q["type"], + "instructions": q.get("instructions") or "", + "state_text": state_text(state, q.get("instructions")), + "keys": keys, + "candidate_texts": texts, + }) + for line in RAW_REQUESTS: + body = json.loads(line) + pairs = build_pairs(body["state"], body["questions"]) + out["raw_cases"].append({ + "line": line, + "questions": {qid: {"state_text": s, "keys": k, "candidate_texts": t} + for qid, (s, k, t) in pairs.items()}, + }) + args.output.write_text(json.dumps(out, ensure_ascii=False, indent=1) + "\n") + print(f"TEXT_ORACLE {args.output} to_text={len(out['to_text'])} " + f"cases={len(out['cases'])} raw_cases={len(out['raw_cases'])}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/recipe/clm/native/transformers_encoder.py b/recipe/clm/native/transformers_encoder.py new file mode 100644 index 0000000..74fe85a --- /dev/null +++ b/recipe/clm/native/transformers_encoder.py @@ -0,0 +1,119 @@ +#!/usr/bin/env python3 +"""A `/v1/embeddings` endpoint backed by transformers, for verifying omni-clm. + +`vllm serve --runner pooling` is the deployment shape, but vLLM is a large install and +this is only needed to produce vectors for a parity check. Transformers loads the same +Qwen3-8B checkpoint and pools the last token, which is what CLM's heads were trained +against — `serve_qwen3_8b.sh` in the CLM repository uses vLLM with `--runner pooling` +precisely because it does the same thing. + + python recipe/clm/native/transformers_encoder.py --model /path/to/Qwen3-8B --port 8090 + +Not for production: one request at a time, no batching across callers. +""" + +from __future__ import annotations + +import argparse +import base64 +import json +import struct +import threading +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer + +MODEL = None +TOKENIZER = None +LOCK = threading.Lock() + + +def embed(texts: list[str]) -> list[list[float]]: + import torch + + with LOCK: + batch = TOKENIZER(texts, return_tensors="pt", padding=True, truncation=True, max_length=2048) + batch = {k: v.to(MODEL.device) for k, v in batch.items()} + with torch.no_grad(): + out = MODEL(**batch) + # Last-token pooling, as CLM's encoder does and its heads were trained on. + mask = batch["attention_mask"] + last = mask.sum(dim=1) - 1 + rows = torch.arange(mask.shape[0], device=mask.device) + hidden = out.last_hidden_state[rows, last] + return hidden.float().cpu().tolist() + + +class Handler(BaseHTTPRequestHandler): + def log_message(self, *args): + pass + + def _json(self, status: int, payload: dict) -> None: + body = json.dumps(payload).encode() + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def do_GET(self) -> None: # noqa: N802 + if self.path.startswith("/v1/models"): + self._json(200, {"object": "list", "data": [{"id": ARGS.model_name, "object": "model"}]}) + else: + self._json(404, {"error": "not found"}) + + def do_POST(self) -> None: # noqa: N802 + if not self.path.startswith("/v1/embeddings"): + self._json(404, {"error": "not found"}) + return + length = int(self.headers.get("Content-Length") or 0) + body = json.loads(self.rfile.read(length) or b"{}") + texts = body.get("input") or [] + if isinstance(texts, str): + texts = [texts] + try: + vectors = embed(texts) + except Exception as exc: # noqa: BLE001 + self._json(500, {"error": f"{type(exc).__name__}: {exc}"}) + return + data = [] + tokens = 0 + for index, vec in enumerate(vectors): + raw = struct.pack(f"<{len(vec)}f", *vec) + data.append({ + "object": "embedding", + "index": index, + "embedding": base64.b64encode(raw).decode(), + }) + tokens += len(TOKENIZER.tokenize(texts[index])) + self._json(200, { + "object": "list", + "data": data, + "model": body.get("model", ARGS.model_name), + "usage": {"prompt_tokens": tokens, "total_tokens": tokens}, + }) + + +def main() -> None: + global MODEL, TOKENIZER, ARGS + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("--model", required=True, help="path to Qwen3-8B") + parser.add_argument("--port", type=int, default=8090) + parser.add_argument("--model-name", default="qwen3-8b") + parser.add_argument("--dtype", default="bfloat16") + ARGS = parser.parse_args() + + import torch + from transformers import AutoModel, AutoTokenizer + + print(f"loading {ARGS.model} ({ARGS.dtype})", flush=True) + TOKENIZER = AutoTokenizer.from_pretrained(ARGS.model) + MODEL = AutoModel.from_pretrained(ARGS.model, dtype=getattr(torch, ARGS.dtype)) + MODEL = MODEL.to("cuda" if torch.cuda.is_available() else "cpu").eval() + print(f"ready on {MODEL.device}, hidden {MODEL.config.hidden_size}", flush=True) + + server = ThreadingHTTPServer(("127.0.0.1", ARGS.port), Handler) + print(f"listening on http://127.0.0.1:{ARGS.port}/v1/embeddings", flush=True) + server.serve_forever() + + +if __name__ == "__main__": + main() diff --git a/src/models/clm/Cargo.toml b/src/models/clm/Cargo.toml new file mode 100644 index 0000000..63b88ba --- /dev/null +++ b/src/models/clm/Cargo.toml @@ -0,0 +1,35 @@ +[package] +name = "omni-clm" +version = "0.1.0" +edition = "2024" +publish = false +description = "CLM: projection heads, typed decision scoring and the embeddings client" + +[dependencies] +anyhow = "1" +# The embeddings payload is base64 f32; reqwest already pulls this in, so it is free. +base64 = "0.22" +memmap2 = "0.9" +reqwest = { version = "0.12", default-features = false, features = ["blocking", "json", "rustls-tls"] } +safetensors = "0.6" +serde = { version = "1", features = ["derive"] } +# preserve_order keeps object key order, which to_text depends on: the heads are +# trained on the caller's field order, so re-sorting a state would change the text. +# float_roundtrip: a float literal must parse to the same double `json.loads` +# produces, or the heads see a different number. +# raw_value: a rendered field keeps the text it arrived as, so an integer literal +# above u64::MAX keeps every digit. It adds `RawValue` and does not change how +# `Number` behaves, unlike `arbitrary_precision`, which would reach every crate +# in the workspace. +serde_json = { version = "1", features = ["preserve_order", "float_roundtrip", "raw_value"] } +sha2 = "0.10" + +# Test bodies live under the repository-level tests/ tree; this is the explicit +# registration CONTRIBUTING asks for, matching src/models/laya/Cargo.toml. +[[test]] +name = "checkpoint" +path = "../../../tests/clm/checkpoint.rs" + +[[test]] +name = "text" +path = "../../../tests/clm/text.rs" diff --git a/src/models/clm/README.md b/src/models/clm/README.md new file mode 100644 index 0000000..4328883 --- /dev/null +++ b/src/models/clm/README.md @@ -0,0 +1,76 @@ +# CLM model engine + +CLM is the second model the project tracks ([#9](https://github.com/ThinkFlowLab/system1-omni/issues/9)), after LAYA. It decides differently in a way LAYA does not cover: **the engine does not compute embeddings.** A frozen `Qwen/Qwen3-8B` encoder runs as its own process behind an OpenAI-compatible `/v1/embeddings` endpoint, and the engine owns everything after it — two projection heads, the cosine score, and the typed answer. + +That split is the point of implementing it second. LAYA's engine owns one forward pass; this one owns a client to someone else's server, plus a candidate-vector cache that persists across requests. + +## What is here + +`omni-clm` reads a converted checkpoint and answers a request. It does not serve HTTP; the request path is a library plus the `clm-run` harness, and the frontend owns the socket. + +| module | responsibility | +| --- | --- | +| `config` | the head geometry, read from the safetensors metadata | +| `weights` | tensor inventory, shape checks, FP32 loading | +| `scoring` | projection, cosine score, softmax, and the three answer types | +| `embedding` | the `/v1/embeddings` client, and a hashing encoder for CPU-only checks | +| `serve` | request parsing, the text the heads see, and the answer shape | + +`clm-run CHECKPOINT_DIR --emb-url URL` reads one request object per line and writes one response per line, which is how `recipe/clm/native/compare_with_reference.py` drives it. + +## The checkpoint is converted first + +A CLM checkpoint is a `torch.save` dict, so it is a pickle and no non-Python reader can open it. `recipe/clm/native/export_weights.py` writes the tensors to safetensors with the head name as a prefix and keeps `cfg`, `hidden_size`, `projection_dim` and `logit_scale` in the metadata. + +```sh +python recipe/clm/native/export_weights.py CLM_v0.1-8B.pt OUT_DIR +``` + +`CLM_v0.1-8B.pt` is the published checkpoint from `Contrastive-LM/CLM-v0.1-8B`; it is 75 MB and holds the two heads, not the 8B encoder. The export is 16 tensors, 18.9 M parameters. + +## The decision, in one pass + +A decision is `softmax(exp(logit_scale) * cos(state_head(s), action_head(c)) / temperature)` over a question's candidates. Both heads are `inp → [LayerNorm →] hidden → out` with GELU, and both projections are L2-normalised before the dot product; the published checkpoint sets `layernorm: true` and `residual: false` with one hidden block. + +The three question types differ only after the distribution exists, which is why they share one scoring path: + +| type | answer | +| --- | --- | +| `choice` | the argmax key, with `confidence` and the full distribution | +| `noul` | the `true` entry of a two-candidate distribution | +| `score` | the expected level index, `sum(i * p_i)`, with `confidence` | + +`confidence` is the top probability minus the mean of the rest, clamped to `[0, 1]`, and `1.0` for a single candidate — the TypeSafe-style definition `schema.py` uses. + +## CPU checks + +The default tests need no checkpoint: + +```sh +cargo test -p omni-clm +``` + +To check the full checkpoint and the decision arithmetic, export it first and point `CLM_EXPORT` at the directory: + +```sh +python recipe/clm/native/export_weights.py /path/to/CLM_v0.1-8B.pt /tmp/clm-export +python recipe/clm/native/head_oracle.py /tmp/clm-export /tmp/clm-export/head-oracle.json +CLM_EXPORT=/tmp/clm-export cargo test -p omni-clm -- --ignored +``` + +Two tests run there: every tensor's FP32 conversion hash against the export oracle, and **five decisions checked against an independent NumPy implementation of the same arithmetic** (`head_oracle.py`) on synthesised embeddings, so the two sides need no encoder to disagree. The embeddings are hash-derived and carry no meaning as model output — the check is that two implementations of the same maths agree. + +The text the heads see is pinned the same way, against `clm.schema` itself: + +```sh +python recipe/clm/native/text_oracle.py /tmp/text-oracle.json +CLM_TEXT_ORACLE=/tmp/text-oracle.json cargo test -p omni-clm -- --ignored +``` + +`to_text`, `state_text` and `candidates` are compared byte for byte over six states and five questions. A separator in the wrong place does not fail loudly — it shifts every probability — so the oracle is produced by importing the reference rather than by transcribing it. + +The default CI job skips these because it does not download the checkpoint. + +## Checking a real decision + +With an encoder reachable, `recipe/clm/native/compare_with_reference.py` sends the same request to `clm-run` and to the reference engine and compares the answers case by case. Both sides call the same embeddings endpoint, so the vectors are identical and the only difference left is the code between the vectors and the answer. diff --git a/src/models/clm/src/bin/clm-run.rs b/src/models/clm/src/bin/clm-run.rs new file mode 100644 index 0000000..3e4998c --- /dev/null +++ b/src/models/clm/src/bin/clm-run.rs @@ -0,0 +1,105 @@ +//! `clm-run`: a decision for a JSON request, so the engine can be checked by hand and +//! against the reference. +//! +//! clm-run CHECKPOINT_DIR --emb-url http://127.0.0.1:8090/v1/embeddings [--model NAME] +//! [--temperature T] [--max-tokens N|none] +//! +//! Reads one request object per line on stdin and writes one response object per line on +//! stdout, which is the shape the reference's own `laya-run`-style harnesses use. +use anyhow::{Context, Result, bail}; +use omni_clm::{Engine, Heads, HttpEncoder, Request, Weights, serve}; +use std::io::{BufRead, Write}; +use std::path::PathBuf; +use std::time::{Duration, Instant}; + +fn main() -> Result<()> { + let args: Vec = std::env::args().skip(1).collect(); + let mut checkpoint: Option = None; + let mut emb_url = "http://127.0.0.1:8090/v1/embeddings".to_string(); + let mut model = "qwen3-8b".to_string(); + let mut temperature = 1.0f32; + let mut max_tokens: Option = Some(2048); + let mut i = 0; + while i < args.len() { + match args[i].as_str() { + "--emb-url" => { + emb_url = args.get(i + 1).context("--emb-url needs a value")?.clone(); + i += 2; + } + "--model" => { + model = args.get(i + 1).context("--model needs a value")?.clone(); + i += 2; + } + "--temperature" => { + temperature = args + .get(i + 1) + .context("--temperature needs a value")? + .parse() + .context("--temperature is not a number")?; + i += 2; + } + "--max-tokens" => { + max_tokens = match args + .get(i + 1) + .context("--max-tokens needs a value")? + .as_str() + { + "none" | "off" => None, + value => Some(value.parse().context("--max-tokens is not a number")?), + }; + i += 2; + } + other if !other.starts_with("--") => { + checkpoint = Some(PathBuf::from(other)); + i += 1; + } + other => bail!("unknown flag {other}"), + } + } + let checkpoint = checkpoint.context("usage: clm-run CHECKPOINT_DIR [--emb-url URL]")?; + + let weights = Weights::open(&checkpoint.join("model.safetensors")) + .with_context(|| format!("open {}", checkpoint.display()))?; + let heads = Heads::load(&weights)?; + eprintln!( + "clm-run: {} tensors, hidden {} -> projection {}, logit_scale {:.4}", + omni_clm::head_tensors(&heads.config.head)?.len(), + heads.config.head.hidden_size, + heads.config.head.projection_dim, + heads.config.logit_scale, + ); + let encoder = HttpEncoder::new(emb_url.clone(), model, Duration::from_secs(90))? + .with_max_tokens(max_tokens); + let engine = Engine::new(heads, encoder); + + let stdin = std::io::stdin(); + let mut stdout = std::io::stdout(); + for line in stdin.lock().lines() { + let line = line?; + let trimmed = line.trim(); + if trimmed.is_empty() { + continue; + } + // `parse_line`, not `parse`: the text is what keeps an integer literal's + // digits, and it is gone once the line has been parsed. + let request = Request::parse_line(trimmed)?; + let started = Instant::now(); + let decision = engine + .decide(&request, temperature) + .with_context(|| format!("decide on {} bytes of request", trimmed.len()))?; + let elapsed = started.elapsed(); + let response = serde_json::json!({ + "model": request.model.clone().unwrap_or_else(|| "clm-latest".into()), + "answers": serve::answers_json(&decision.answers), + "usage": { + "billing_units": decision.answers.len(), + "input_tokens": decision.encoder_tokens, + "output_tokens": 0, + }, + "elapsed_ms": elapsed.as_secs_f64() * 1000.0, + }); + writeln!(stdout, "{response}")?; + stdout.flush()?; + } + Ok(()) +} diff --git a/src/models/clm/src/config.rs b/src/models/clm/src/config.rs new file mode 100644 index 0000000..ded0580 --- /dev/null +++ b/src/models/clm/src/config.rs @@ -0,0 +1,109 @@ +//! CLM head configuration, read from the exported safetensors metadata. +//! +//! A CLM checkpoint is a `torch.save` dict, so `recipe/clm/native/export_weights.py` +//! converts it first; everything this crate reads is safetensors. The head geometry is +//! not a file of its own, so it travels in the safetensors metadata as `cfg`. +use anyhow::{Context, Result, ensure}; +use serde::Deserialize; + +/// The head shape the checkpoint was trained with. `depth` counts `inp`, the hidden +/// blocks and `out`, so `depth - 2` is the number of `hidden.N` / `norms.N` pairs. +#[derive(Debug, Clone, Deserialize, PartialEq, Eq)] +pub struct HeadConfig { + pub hidden_size: usize, + pub projection_dim: usize, + pub width: usize, + pub depth: usize, + pub activation: String, + pub layernorm: bool, + pub residual: bool, + #[serde(default)] + pub model: Option, +} + +impl HeadConfig { + /// Number of hidden blocks; `depth` includes the input and output projections. + pub fn hidden_blocks(&self) -> Result { + ensure!( + self.depth >= 2, + "depth {} cannot cover an input and an output projection", + self.depth + ); + Ok(self.depth - 2) + } +} + +/// The reference's cap on `exp(logit_scale)`. +pub const MAX_SCALE: f32 = 100.0; + +#[derive(Debug, Clone, PartialEq)] +pub struct Config { + pub head: HeadConfig, + /// `exp(logit_scale)` is the inverse InfoNCE temperature the heads were trained with. + pub logit_scale: f32, +} + +impl Config { + /// Build from the safetensors metadata written by the export tool. + pub fn from_metadata( + metadata: Option<&std::collections::HashMap>, + ) -> Result { + let metadata = metadata.context("the checkpoint carries no metadata")?; + ensure!( + metadata.get("format").map(String::as_str) == Some("clm-heads"), + "not a converted CLM head checkpoint: format is {:?}", + metadata.get("format") + ); + let cfg = metadata.get("cfg").context("metadata has no cfg")?; + // `hidden_size` and `projection_dim` are repeated at the top level, and `cfg` may + // omit them: the exporter takes them from the checkpoint's own entries. They are + // required fields of `HeadConfig`, so they have to be filled in *before* + // deserializing — a fallback applied afterwards never runs, because the parse + // fails on the missing field first. + let mut raw: serde_json::Map = + serde_json::from_str(cfg).with_context(|| format!("parse cfg {cfg}"))?; + for key in ["hidden_size", "projection_dim"] { + if !raw.contains_key(key) + && let Some(value) = metadata.get(key).and_then(|v| v.parse::().ok()) + { + raw.insert(key.to_string(), value.into()); + } + } + let head: HeadConfig = serde_json::from_value(serde_json::Value::Object(raw)) + .with_context(|| format!("parse cfg {cfg}"))?; + let logit_scale: f32 = metadata + .get("logit_scale") + .context("metadata has no logit_scale")? + .parse() + .context("logit_scale is not a number")?; + + ensure!(head.hidden_size > 0, "hidden_size must be positive"); + ensure!(head.width > 0, "width must be positive"); + ensure!(head.projection_dim > 0, "projection_dim must be positive"); + ensure!( + head.activation == "gelu" || head.activation == "relu" || head.activation == "silu", + "unsupported activation {:?}", + head.activation + ); + // The upstream head applies LayerNorm before every hidden block and adds the + // residual after it; the published 0.1 checkpoint sets both to true/false + // respectively, and the export keeps them so the two paths stay distinguishable. + head.hidden_blocks()?; + ensure!( + logit_scale.is_finite(), + "logit_scale {logit_scale} is not finite" + ); + Ok(Self { head, logit_scale }) + } + + /// `exp(logit_scale)`, capped as the reference caps it. + /// + /// `heads.py` computes `exp(logit_scale).clamp(max=100.0)`, and the published + /// checkpoint's `logit_scale` is 4.6132, whose exponential is 100.82 — so the cap + /// binds and the effective scale is 100, not 100.82. Without it every probability is + /// off by about 0.8 %, which is what the end-to-end comparison against the reference + /// caught; no CPU-side oracle can, because both sides of those share this constant. + pub fn scale(&self) -> f32 { + self.logit_scale.exp().min(MAX_SCALE) + } +} diff --git a/src/models/clm/src/embedding.rs b/src/models/clm/src/embedding.rs new file mode 100644 index 0000000..094244a --- /dev/null +++ b/src/models/clm/src/embedding.rs @@ -0,0 +1,197 @@ +//! The embeddings client: the engine's half of the split. +//! +//! CLM does not compute embeddings. A frozen Qwen3-8B encoder runs as its own process +//! behind an OpenAI-compatible `/v1/embeddings` endpoint, and this module is the client +//! to it. Everything the engine owns happens after the vectors come back. +//! +//! The wire shape follows `src/clm/embedder.py` in the CLM reference: a POST of +//! `{model, input, encoding_format, truncate_prompt_tokens}`, a base64 `f32` payload per +//! input in `index` order, and an `l2` normalisation applied on receipt — the encoder is +//! asked for raw vectors and the client normalises, so a server that already normalises +//! is harmless. +//! +//! `truncate_prompt_tokens` is not optional in practice: the reference sends its +//! `max_tokens` (2048 by default, matching the documented `--max-model-len 2048`), and a +//! text longer than that is truncated there but rejected by a server asked to embed it +//! whole. +use anyhow::{Context, Result, bail, ensure}; +use base64::Engine as _; +use std::time::Duration; + +use crate::scoring::normalize; + +/// Sends texts to the encoder and returns one row per text. +/// +/// A trait rather than a struct so the engine can be driven without an encoder; the +/// tests use [`HashingEncoder`], and a deployment uses [`HttpEncoder`]. +pub trait Encoder: Send + Sync { + /// One `hidden_size`-wide, L2-normalised vector per input, in the order given, and + /// the tokens the encoder reported spending on them. + fn embed(&self, texts: &[String]) -> Result<(Vec>, u64)>; +} + +/// A deterministic encoder with no server behind it. +/// +/// The vector is derived from a SHA-256 of the text filled to `dim` and then normalised, +/// which is exactly what `recipe/clm/native/head_oracle.py` does — so a decision taken +/// against this encoder can be compared with the Python oracle, and the whole engine can +/// be exercised on a machine with no GPU and no weights. The values carry no meaning as +/// model output; the point is that both implementations see the same numbers. +pub struct HashingEncoder { + dim: usize, +} + +impl HashingEncoder { + pub fn new(dim: usize) -> Self { + Self { dim } + } + + /// The vector for one text, as both this encoder and the oracle compute it. + pub fn vector(text: &str, dim: usize) -> Vec { + use sha2::{Digest, Sha256}; + + let mut out: Vec = Vec::with_capacity(dim); + let mut counter = 0u32; + while out.len() < dim { + let digest = Sha256::digest(format!("{counter}:{text}").as_bytes()); + for chunk in digest.as_chunks::<4>().0 { + if out.len() == dim { + break; + } + out.push(u32::from_be_bytes(*chunk) as f64 as f32 / 2f64.powi(31) as f32 - 1.0); + } + counter += 1; + } + normalize(&mut out); + out + } +} + +impl Encoder for HashingEncoder { + /// No server, so nothing was spent. + fn embed(&self, texts: &[String]) -> Result<(Vec>, u64)> { + Ok((texts.iter().map(|t| Self::vector(t, self.dim)).collect(), 0)) + } +} + +/// An OpenAI-compatible `/v1/embeddings` endpoint, which is what `vllm serve --runner +/// pooling` exposes. +pub struct HttpEncoder { + url: String, + model: String, + client: reqwest::blocking::Client, + /// Batching, as `embedder.py` does: a request carries at most this many inputs. + batch: usize, + /// `embedder.py`'s `max_tokens`, sent as `truncate_prompt_tokens`. `None` sends no + /// limit at all, which is what a reference built with `max_tokens=None` does. + max_tokens: Option, +} + +impl HttpEncoder { + pub fn new( + url: impl Into, + model: impl Into, + timeout: Duration, + ) -> Result { + let client = reqwest::blocking::Client::builder() + .timeout(timeout) + .build() + .context("build the embeddings HTTP client")?; + Ok(Self { + url: url.into(), + model: model.into(), + client, + batch: 512, + // The reference's default, and the deployment's `--max-model-len`. + max_tokens: Some(2048), + }) + } + + /// Override the truncation limit sent as `truncate_prompt_tokens`. + pub fn with_max_tokens(mut self, max_tokens: Option) -> Self { + self.max_tokens = max_tokens; + self + } + + /// The request body, which is `embedder.py`'s own. + fn body(&self, texts: &[String]) -> serde_json::Value { + let mut body = serde_json::json!({ + "model": self.model, + "input": texts, + "encoding_format": "base64", + }); + if let Some(max_tokens) = self.max_tokens { + body["truncate_prompt_tokens"] = max_tokens.into(); + } + body + } + + /// One request: its rows in `index` order, and the tokens the endpoint reported. + fn fetch(&self, texts: &[String]) -> Result<(Vec>, u64)> { + let body = self.body(texts); + let response = self + .client + .post(&self.url) + .json(&body) + .send() + .with_context(|| format!("embeddings request to {}", self.url))?; + let status = response.status(); + let payload: serde_json::Value = response + .json() + .with_context(|| format!("embeddings response from {} was not JSON", self.url))?; + if !status.is_success() { + bail!("embeddings endpoint returned {status}: {payload}"); + } + + let data = payload["data"] + .as_array() + .context("embeddings response has no data array")?; + let mut out: Vec>> = vec![None; texts.len()]; + for row in data { + let index = row["index"].as_u64().context("a data row has no index")? as usize; + ensure!(index < out.len(), "data index {index} is out of range"); + let encoded = row["embedding"] + .as_str() + .context("embedding is not a base64 string")?; + let raw = base64::engine::general_purpose::STANDARD + .decode(encoded) + .context("embedding is not valid base64")?; + ensure!( + raw.len() % 4 == 0, + "embedding payload is {} bytes, not a whole number of f32", + raw.len() + ); + let mut row: Vec = raw + .as_chunks::<4>() + .0 + .iter() + .map(|b| f32::from_le_bytes(*b)) + .collect(); + normalize(&mut row); + out[index] = Some(row); + } + let tokens = payload["usage"]["prompt_tokens"].as_u64().unwrap_or(0); + out.into_iter() + .enumerate() + .map(|(i, v)| v.with_context(|| format!("no embedding for input {i}"))) + .collect::>>() + .map(|v| (v, tokens)) + } +} + +impl Encoder for HttpEncoder { + fn embed(&self, texts: &[String]) -> Result<(Vec>, u64)> { + let mut out = Vec::with_capacity(texts.len()); + let mut tokens = 0; + for chunk in texts.chunks(self.batch) { + let (rows, spent) = self.fetch(chunk)?; + out.extend(rows); + tokens += spent; + } + Ok((out, tokens)) + } +} + +#[cfg(test)] +#[path = "../../../../tests/clm/embedding.rs"] +mod tests; diff --git a/src/models/clm/src/lib.rs b/src/models/clm/src/lib.rs new file mode 100644 index 0000000..57862b4 --- /dev/null +++ b/src/models/clm/src/lib.rs @@ -0,0 +1,24 @@ +//! CLM: the second System1-Omni model engine. +//! +//! A CLM decision is made over embeddings the engine does not compute. A frozen +//! Qwen3-8B encoder sits behind an OpenAI-compatible `/v1/embeddings` endpoint and the +//! engine owns everything after it: load the two projection heads, normalise their +//! output, score candidates by cosine similarity under the trained temperature, and +//! assemble `choice`, `noul` and `score` answers. +//! +//! That split is why this model is implemented second — LAYA's engine owns one forward +//! pass, while this one owns a client to someone else's server. +//! +//! Nothing here needs a GPU. The encoder can be [`embedding::HashingEncoder`], whose +//! vectors are the Python oracle's, so the decision path is checked on a CPU-only machine. +pub mod config; +pub mod embedding; +pub mod scoring; +pub mod serve; +pub mod weights; + +pub use config::{Config, HeadConfig}; +pub use embedding::{Encoder, HashingEncoder, HttpEncoder}; +pub use scoring::{Answer, Kind, Question, answer, confidence, distribution}; +pub use serve::{Decision, Engine, NumberLiterals, Request, to_text_json}; +pub use weights::{Head, Heads, Weights, head_tensors}; diff --git a/src/models/clm/src/scoring.rs b/src/models/clm/src/scoring.rs new file mode 100644 index 0000000..d033735 --- /dev/null +++ b/src/models/clm/src/scoring.rs @@ -0,0 +1,323 @@ +//! The typed decision computation: project embeddings, score candidates, assemble answers. +//! +//! This mirrors `src/clm/engine.py` and `src/clm/schema.py`: a state and each candidate +//! are embedded elsewhere (Qwen3-8B behind `/v1/embeddings`) and arrive here as vectors. +//! The state head sees the state, the action head sees every candidate, both projections +//! are L2-normalised, and the score of a pair is `exp(logit_scale) * cos(state, candidate)`, +//! divided by the request temperature and softmaxed across the question's candidates. +//! +//! The three question types differ only after the distribution exists. +use anyhow::{Result, ensure}; + +use crate::weights::{Head, Heads}; + +/// One question: its type and its candidate keys in the order the caller offered them. +#[derive(Debug, Clone, PartialEq)] +pub struct Question { + pub id: String, + pub kind: Kind, + /// Candidate keys, in order. For `Noul` these are `["false", "true"]`. + pub keys: Vec, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Kind { + Choice, + Noul, + Score, +} + +/// One answer, matching the shapes `client.py` parses. +#[derive(Debug, Clone, PartialEq)] +pub enum Answer { + Choice { + choice: String, + confidence: f32, + probabilities: Vec<(String, f32)>, + }, + Noul { + noul: f32, + }, + Score { + score: f32, + confidence: f32, + /// Level key to level text, in answer order: `answer_from_probs` returns this so a + /// consumer can read a score back as a level without the request. + legend: Vec<(String, String)>, + probabilities: Vec<(String, f32)>, + }, +} + +impl Answer { + /// The discrete label, as `schema.label_of` defines it. + pub fn label(&self) -> String { + match self { + Answer::Choice { choice, .. } => choice.clone(), + Answer::Noul { noul } => if *noul >= 0.5 { "true" } else { "false" }.to_string(), + // The first maximum, as Python's `max` keeps. `max_by` would return the + // last of several equal values, which is a different label for the same + // distribution. + Answer::Score { probabilities, .. } => probabilities + .iter() + .fold(None, |best: Option<&(String, f32)>, kv| match best { + Some(b) if b.1 >= kv.1 => Some(b), + _ => Some(kv), + }) + .map(|(k, _)| k.clone()) + .unwrap_or_default(), + } + } +} + +/// L2-normalise one row in place, the way `torch.nn.functional.normalize` does. +/// +/// Shared with [`crate::embedding`], which normalises what the encoder hands back. +pub fn normalize(row: &mut [f32]) { + let norm = row.iter().map(|v| v * v).sum::().sqrt(); + if norm > 0.0 { + for v in row.iter_mut() { + *v /= norm; + } + } +} + +fn gelu(x: f32) -> f32 { + 0.5 * x * (1.0 + erf(x * std::f32::consts::FRAC_1_SQRT_2)) +} + +fn relu(x: f32) -> f32 { + x.max(0.0) +} + +fn silu(x: f32) -> f32 { + x / (1.0 + (-x).exp()) +} + +/// Abramowitz & Stegun 7.1.26. The coefficients are kept at full precision through named +/// constants: rounding one to a shorter f32 literal changes the value, and +/// `recipe/clm/native/head_oracle.py` carries the same digits so that both sides remain +/// the same arithmetic rather than merely similar. +#[allow(clippy::excessive_precision)] +const ERF_P: f32 = 0.3275911; +#[allow(clippy::excessive_precision)] +const ERF_A1: f32 = 0.254829592; +#[allow(clippy::excessive_precision)] +const ERF_A2: f32 = -0.284496736; +#[allow(clippy::excessive_precision)] +const ERF_A3: f32 = 1.421413741; +#[allow(clippy::excessive_precision)] +const ERF_A4: f32 = -1.453152027; +#[allow(clippy::excessive_precision)] +const ERF_A5: f32 = 1.061405429; + +fn erf(x: f32) -> f32 { + let sign = if x < 0.0 { -1.0 } else { 1.0 }; + let x = x.abs(); + let t = 1.0 / (1.0 + ERF_P * x); + let y = 1.0 + - (((((ERF_A5 * t + ERF_A4) * t + ERF_A3) * t + ERF_A2) * t + ERF_A1) * t * (-x * x).exp()); + sign * y +} + +fn activate(kind: &str, x: f32) -> f32 { + match kind { + "relu" => relu(x), + "silu" => silu(x), + _ => gelu(x), + } +} + +#[inline] +fn linear_into(x: &[f32], w: &[f32], b: &[f32], n: usize, k: usize, out: &mut [f32]) { + for i in 0..n { + let row = &w[i * k..(i + 1) * k]; + let mut acc = b[i]; + for (a, wv) in x.iter().zip(row) { + acc += a * wv; + } + out[i] = acc; + } +} + +fn layer_norm(x: &mut [f32], weight: &[f32], bias: &[f32]) { + let n = x.len() as f32; + let mean = x.iter().sum::() / n; + let var = x.iter().map(|v| (v - mean) * (v - mean)).sum::() / n; + let inv = 1.0 / (var + 1e-5).sqrt(); + for i in 0..x.len() { + x[i] = (x[i] - mean) * inv * weight[i] + bias[i]; + } +} + +/// Project one embedding through a head, returning the L2-normalised vector. +pub fn project(head: &Head, cfg: &crate::config::HeadConfig, x: &[f32]) -> Result> { + ensure!( + x.len() == cfg.hidden_size, + "embedding has {} values, the head expects {}", + x.len(), + cfg.hidden_size + ); + let w = cfg.width; + let mut h = vec![0.0f32; w]; + linear_into( + x, + &head.inp_weight, + &head.inp_bias, + w, + cfg.hidden_size, + &mut h, + ); + for v in h.iter_mut() { + *v = activate(&cfg.activation, *v); + } + if !head.hidden_weight.is_empty() { + let mut hidden = vec![0.0f32; w]; + linear_into( + &h, + &head.hidden_weight, + &head.hidden_bias, + w, + w, + &mut hidden, + ); + if let (Some(nw), Some(nb)) = (&head.norm_weight, &head.norm_bias) { + layer_norm(&mut hidden, nw, nb); + } + for v in hidden.iter_mut() { + *v = activate(&cfg.activation, *v); + } + if cfg.residual { + for i in 0..w { + hidden[i] += h[i]; + } + } + h = hidden; + } + let p = cfg.projection_dim; + let mut out = vec![0.0f32; p]; + linear_into(&h, &head.out_weight, &head.out_bias, p, w, &mut out); + normalize(&mut out); + Ok(out) +} + +/// `softmax(scale * cos / temperature)`, the distribution every question type starts from. +pub fn distribution( + heads: &Heads, + state: &[f32], + candidates: &[Vec], + temperature: f32, +) -> Result> { + ensure!( + !candidates.is_empty(), + "a question needs at least one candidate" + ); + ensure!( + temperature > 0.0 && temperature <= 100.0, + "temperature must be in (0, 100], got {temperature}" + ); + let cfg = &heads.config.head; + let zs = project(&heads.state, cfg, state)?; + let mut logits = Vec::with_capacity(candidates.len()); + for candidate in candidates { + let zc = project(&heads.action, cfg, candidate)?; + let cos: f32 = zs.iter().zip(&zc).map(|(a, b)| a * b).sum(); + logits.push(heads.config.scale() * cos / temperature); + } + let max = logits.iter().copied().fold(f32::MIN, f32::max); + let exp: Vec = logits.iter().map(|v| (v - max).exp()).collect(); + let sum: f32 = exp.iter().sum(); + ensure!(sum > 0.0, "softmax denominator is zero"); + Ok(exp.iter().map(|v| v / sum).collect()) +} + +/// `schema.confidence`: the top probability minus the mean of the rest, clamped to [0, 1]. +/// A single candidate is fully decided by definition. +pub fn confidence(probs: &[f32]) -> f32 { + if probs.len() < 2 { + return 1.0; + } + let (j, top) = + probs.iter().enumerate().fold( + (0usize, f32::MIN), + |best, (i, p)| { + if *p > best.1 { (i, *p) } else { best } + }, + ); + let rest = (probs.len() - 1) as f32; + let mean_rest = probs + .iter() + .enumerate() + .filter(|(i, _)| *i != j) + .map(|(_, p)| *p) + .sum::() + / rest; + (top - mean_rest).clamp(0.0, 1.0) +} + +/// Assemble the answer for one question from its distribution. +/// +/// `texts` is the candidate text per key, which `answer_from_probs` reads out of the +/// question's `criteria`; a `score` answer carries them back as its legend. +pub fn answer(question: &Question, texts: &[String], probs: &[f32]) -> Result { + ensure!( + question.keys.len() == probs.len(), + "{}: {} keys but {} probabilities", + question.id, + question.keys.len(), + probs.len() + ); + let pairs: Vec<(String, f32)> = question + .keys + .iter() + .cloned() + .zip(probs.iter().copied()) + .collect(); + Ok(match question.kind { + Kind::Noul => { + // `answer_from_probs` reads the `true` entry, which is the last key by + // convention; a caller that orders them differently still gets a stable + // answer because the key is looked up by name. + let p = pairs + .iter() + .find(|(k, _)| k == "true") + .map(|(_, p)| *p) + .unwrap_or_else(|| *probs.last().expect("checked non-empty")); + Answer::Noul { noul: p } + } + Kind::Choice => { + let (choice, _) = pairs + .iter() + .fold(None, |best: Option<&(String, f32)>, kv| match best { + Some(b) if b.1 >= kv.1 => Some(b), + _ => Some(kv), + }) + .expect("checked non-empty"); + Answer::Choice { + choice: choice.clone(), + confidence: confidence(probs), + probabilities: pairs, + } + } + Kind::Score => { + // The expected level index, `sum(i * p_i)`. + let score = probs + .iter() + .enumerate() + .map(|(i, p)| i as f32 * p) + .sum::(); + Answer::Score { + score, + confidence: confidence(probs), + // The reference builds the legend from the criteria as `str(i) -> text`, + // which is exactly the keys and the candidate texts it already has. + legend: question + .keys + .iter() + .cloned() + .zip(texts.iter().cloned()) + .collect(), + probabilities: pairs, + } + } + }) +} diff --git a/src/models/clm/src/serve.rs b/src/models/clm/src/serve.rs new file mode 100644 index 0000000..006b6ad --- /dev/null +++ b/src/models/clm/src/serve.rs @@ -0,0 +1,714 @@ +//! The request path: a `/v1/systemone` body to typed answers. +//! +//! Mirrors `src/clm/schema.py` and `src/clm/engine.py`. Two text functions matter and +//! both are reproduced exactly, down to the separator, because the heads were trained on +//! this layout: +//! +//! - the **state head** sees the context and the question joined by a blank line +//! (`state_text`), so a question belongs in `instructions`, not repeated in the state; +//! - the **action head** sees each candidate's own text with nothing prefixed for +//! `choice`, but `": "` for `noul`. +//! +//! A mismatch here does not fail loudly — it shifts every probability — so both are +//! pinned by tests against the reference implementation's own output. +use anyhow::{Context, Result, bail, ensure}; +use serde::Deserialize; +use serde_json::value::RawValue; +use serde_json::{Map, Value}; +use std::collections::HashMap; + +/// The fields of a request that get rendered, kept as the text they arrived as. +/// +/// `serde_json` has no way to hand back what an integer literal said: one that does not +/// fit an `i64` or a `u64` is parsed straight to a double, and `18446744073709551616` +/// comes back as `1.8446744073709552e19`. Python gets an arbitrary-precision `int` +/// instead, and `str` prints every digit. `RawValue` is the only place a literal +/// survives, so the fields that reach the encoder are deserialized twice — once for the +/// structure and once for the text — and rendered from the text. +#[derive(Deserialize)] +struct RawBody { + state: Box, + /// Each question object as its own text. A map, not `serde_json::Map`: with + /// `preserve_order` that one is `IndexMap` and cannot hold another + /// type. The request order comes from the parsed `Value` instead, and this is + /// looked up by id. + #[serde(default)] + questions: std::collections::HashMap>, +} + +/// The two rendered fields of one question. +#[derive(Deserialize)] +struct RawQuestion { + #[serde(default)] + instructions: Option>, + #[serde(default)] + criteria: Option>, +} + +/// What each rendered field arrived as. +#[derive(Debug, Clone, Default)] +struct RawFields { + state: String, + /// Question id to (instructions, criteria). + questions: std::collections::HashMap, Option)>, +} + +impl RawFields { + fn of(line: &str) -> Result { + let raw: RawBody = serde_json::from_str(line).context("read the request text")?; + let mut questions = std::collections::HashMap::with_capacity(raw.questions.len()); + for (id, object) in &raw.questions { + let q: RawQuestion = serde_json::from_str(object.get()) + .with_context(|| format!("read the text of question {id:?}"))?; + questions.insert( + id.clone(), + ( + q.instructions.map(|v| v.get().to_string()), + q.criteria.map(|v| v.get().to_string()), + ), + ); + } + Ok(Self { + state: raw.state.get().to_string(), + questions, + }) + } + + fn question(&self, id: &str) -> (Option<&str>, Option<&str>) { + self.questions + .get(id) + .map(|(i, c)| (i.as_deref(), c.as_deref())) + .unwrap_or((None, None)) + } +} + +/// Render one field from the text it arrived as. Each call gets its own literals, so +/// rendering the same field twice — the state is rendered once per question — is safe, +/// and nothing depends on the order the fields appear in. +fn render_raw(raw: &str) -> Result { + let value: Value = serde_json::from_str(raw).context("render a request field")?; + Ok(render_with(raw, &value)) +} + +/// Render `value` from the text it arrived as. The literals are the value's own, so this +/// can be called per key: a number is always spelled by the literal it was written with, +/// whatever order the keys are read in. +fn render_with(raw: &str, value: &Value) -> String { + render(value, 0, &mut NumberLiterals::of(raw)) +} + +use crate::embedding::Encoder; +use crate::scoring::{self, Answer, Kind, Question}; +use crate::weights::Heads; + +/// The candidate keys of a `noul` question, in the order the reference uses. +pub const NOUL_KEYS: [&str; 2] = ["false", "true"]; + +/// A request body: a state and the questions asked of it. +#[derive(Debug, Clone)] +pub struct Request { + pub state: Value, + pub model: Option, + /// Question id to question object, in insertion order. + pub questions: Vec<(String, QuestionRequest)>, + /// What the rendered fields arrived as, when the caller had the text. `parse` + /// cannot fill this in; `parse_line` can, and it is what keeps an integer literal + /// above `u64::MAX` from being rounded. + raw: Option, +} + +#[derive(Debug, Clone)] +pub struct QuestionRequest { + pub kind: Kind, + pub instructions: String, + pub criteria: Option, +} + +/// One question prepared for the encoder and the scorer. +#[derive(Debug, Clone)] +pub struct Prepared { + pub id: String, + pub question: Question, + /// The text the state head embeds: context and question, blank-line separated. + pub state_text: String, + /// The text the action head embeds per candidate, in `question.keys` order. + pub candidate_texts: Vec, +} + +impl Request { + /// Parse a `/v1/systemone` body from the text it arrived as. + /// + /// This is the entry point a request path should use: numbers are rendered from + /// their literals, which [`Request::parse`] cannot do once the text is gone. + pub fn parse_line(line: &str) -> Result { + let body: Value = serde_json::from_str(line).context("request is not JSON")?; + let mut request = Self::parse(&body)?; + request.raw = Some(RawFields::of(line)?); + Ok(request) + } + + /// Parse a `/v1/systemone` body. Unknown top-level fields are ignored, as the + /// reference does. + pub fn parse(body: &Value) -> Result { + let object = body.as_object().context("body must be an object")?; + let state = object + .get("state") + .context("body must have a state")? + .clone(); + let raw = object + .get("questions") + .and_then(Value::as_object) + .context("body must have a questions object")?; + ensure!(!raw.is_empty(), "questions must not be empty"); + + let mut questions = Vec::with_capacity(raw.len()); + for (id, q) in raw { + let q = q + .as_object() + .with_context(|| format!("question {id:?} is not an object"))?; + let kind = match q.get("type").and_then(Value::as_str) { + Some("choice") => Kind::Choice, + Some("score") => Kind::Score, + Some("noul") => Kind::Noul, + other => bail!("question {id:?}: unknown type {other:?}"), + }; + let instructions = q + .get("instructions") + .map(to_text) + .unwrap_or_default() + .trim() + .to_string(); + questions.push(( + id.clone(), + QuestionRequest { + kind, + instructions, + criteria: q.get("criteria").cloned(), + }, + )); + } + Ok(Self { + state, + model: object + .get("model") + .and_then(Value::as_str) + .map(str::to_owned), + questions, + raw: None, + }) + } + + /// Turn each question into the keys, the state text and the candidate texts. + pub fn prepare(&self) -> Result> { + match &self.raw { + Some(raw) => self.prepare_from_text(raw), + None => Ok(self + .questions + .iter() + .map(|(id, q)| { + let (keys, candidate_texts) = + candidates_with(q, None, &q.instructions, &mut NumberLiterals::default()) + .with_context(|| format!("question {id:?} has invalid criteria"))?; + Ok(Prepared { + id: id.clone(), + question: Question { + id: id.clone(), + kind: q.kind, + keys, + }, + state_text: state_text(&self.state, &q.instructions), + candidate_texts, + }) + }) + .collect::>>()?), + } + } + + /// The same, rendering each field from the text it arrived as. + /// + /// The state is rendered once rather than once per question, and each rendered + /// field gets its own literals, so neither the number of questions nor the order + /// the fields were written in can shift which literal belongs to which number. + fn prepare_from_text(&self, raw: &RawFields) -> Result> { + let mut state_numbers = NumberLiterals::of(&raw.state); + let state = render(&self.state, 0, &mut state_numbers) + .trim() + .to_string(); + + let mut prepared = Vec::with_capacity(self.questions.len()); + for (id, q) in &self.questions { + let (instructions_raw, criteria_raw) = raw.question(id); + let instructions = match instructions_raw { + Some(text) => render_raw(text)?.trim().to_string(), + None => q.instructions.clone(), + }; + let mut numbers = match criteria_raw { + Some(text) => NumberLiterals::of(text), + None => NumberLiterals::default(), + }; + let (keys, candidate_texts) = + candidates_with(q, criteria_raw, &instructions, &mut numbers) + .with_context(|| format!("question {id:?} has invalid criteria"))?; + let state_text = if !state.is_empty() && !instructions.is_empty() { + format!("{state}\n\n{instructions}") + } else if !state.is_empty() { + state.clone() + } else { + instructions + }; + prepared.push(Prepared { + id: id.clone(), + question: Question { + id: id.clone(), + kind: q.kind, + keys, + }, + state_text, + candidate_texts, + }); + } + Ok(prepared) + } +} + +/// Context first, question last — the layout the heads were trained on. +pub fn state_text(state: &Value, instructions: &str) -> String { + state_text_with(state, instructions, &mut NumberLiterals::default()) +} + +fn state_text_with(state: &Value, instructions: &str, numbers: &mut NumberLiterals) -> String { + let s = render(state, 0, numbers).trim().to_string(); + let i = instructions.trim(); + if !s.is_empty() && !i.is_empty() { + format!("{s}\n\n{i}") + } else if !s.is_empty() { + s + } else { + i.to_string() + } +} + +/// Option keys in answer order, and the candidate text per option. +pub fn candidates(q: &QuestionRequest) -> Result<(Vec, Vec)> { + candidates_with(q, None, &q.instructions, &mut NumberLiterals::default()) +} + +/// The same, rendering each criteria value from the text it arrived as. +/// +/// `criteria_raw` is the criteria field's own text when the request had one. Binding each +/// value to the key it was written under is what keeps a literal with its number: a +/// `choice` is read in source order but a `noul` in [`NOUL_KEYS`] order, so +/// `{"true": 1, "false": 2}` has to render as `false: 2` — counting positions gives +/// `false: 1`, which is a different answer, not a different spelling. +/// +/// `instructions` is `to_text` of the question's instructions, rendered from their text +/// when there was one. The `noul` defaults are built from it, so a literal in the +/// statement reaches the candidate too. +fn candidates_with( + q: &QuestionRequest, + criteria_raw: Option<&str>, + instructions: &str, + numbers: &mut NumberLiterals, +) -> Result<(Vec, Vec)> { + // Only an object has values to find by key. Anything else — a `score`'s list, an + // absent or null criteria — renders in the order it was written, which is what the + // cursor is for. The map is keyed, so it does not matter that it is unordered: the + // keys come from the question, not from here. + let by_key: Option> = match (criteria_raw, q.criteria.as_ref()) { + (Some(text), Some(Value::Object(_))) => { + Some(serde_json::from_str(text).context("read the criteria text of a question")?) + } + _ => None, + }; + let value_text = |key: &str, v: &Value, numbers: &mut NumberLiterals| match by_key + .as_ref() + .and_then(|values| values.get(key)) + { + Some(raw) => render_with(raw.get(), v), + None => render(v, 0, numbers), + }; + match q.kind { + Kind::Choice => { + let crit = q + .criteria + .as_ref() + .and_then(Value::as_object) + .context("choice question needs a non-empty 'criteria' object")?; + ensure!( + !crit.is_empty(), + "choice question needs a non-empty 'criteria' object" + ); + let keys: Vec = crit.keys().cloned().collect(); + // The option's own text when one is given, else the key. Nothing is prefixed. + // "Given" is the reference's test — null or the empty string — so an empty + // container is a description that renders to nothing, not a missing one. + let texts = keys + .iter() + .map(|k| match &crit[k] { + v if v.is_null() || v.as_str() == Some("") => k.clone(), + v => value_text(k, v, numbers), + }) + .collect(); + Ok((keys, texts)) + } + Kind::Score => { + let crit = q + .criteria + .as_ref() + .and_then(Value::as_array) + .context("score question needs 'criteria' as an ordered list of levels")?; + ensure!(crit.len() >= 2, "score question needs at least two levels"); + let keys = (0..crit.len()).map(|i| i.to_string()).collect(); + let texts = crit.iter().map(|c| render(c, 0, numbers)).collect(); + Ok((keys, texts)) + } + Kind::Noul => { + let crit = q.criteria.as_ref().and_then(Value::as_object); + let mut texts = Vec::with_capacity(NOUL_KEYS.len()); + for k in NOUL_KEYS { + // The same "given or not" test as `choice`: `crit.get(k)` in `(None, "")`. + let body = match crit.and_then(|c| c.get(k)) { + Some(v) if !v.is_null() && v.as_str() != Some("") => value_text(k, v, numbers), + _ if !instructions.is_empty() => { + if k == "true" { + format!("Yes. This is true: {instructions}") + } else { + format!("No. This is false: {instructions}") + } + } + _ => k.to_string(), + }; + texts.push(format!("{k}: {body}")); + } + Ok((NOUL_KEYS.iter().map(|k| k.to_string()).collect(), texts)) + } + } +} + +/// The number literals of one JSON document, in the order they appear. +/// +/// `serde_json` cannot hand back what an integer literal said. One that does not fit an +/// `i64` or a `u64` is parsed straight to a double — `18446744073709551616` becomes +/// `1.8446744073709552e19` — where `json.loads` gives Python an arbitrary-precision `int` +/// and `str` prints every digit. The literal is the only place those digits survive. +/// +/// A lexer is enough to collect them because a depth-first walk of the parsed value meets +/// numbers in the same order, so the two can be paired up without a second parser. The +/// walk below is that pairing; [`NumberLiterals::of`] documents what keeps them in step. +#[derive(Debug, Clone, Default)] +pub struct NumberLiterals { + literals: Vec, + at: usize, +} + +impl NumberLiterals { + /// Collect the literals of `raw`, which must already have parsed as JSON. + /// + /// Strings are skipped whole, escapes included, so a digit inside one is not a + /// number. Everything else that looks like a number is one: the document is valid + /// JSON, which is what makes a lexer sufficient here. + pub fn of(raw: &str) -> Self { + let bytes = raw.as_bytes(); + let mut literals = Vec::new(); + let mut at = 0; + while at < bytes.len() { + match bytes[at] { + b'"' => { + at += 1; + while at < bytes.len() && bytes[at] != b'"' { + at += if bytes[at] == b'\\' { 2 } else { 1 }; + } + at += 1; + } + b'-' | b'0'..=b'9' => { + let start = at; + while at < bytes.len() + && matches!(bytes[at], b'0'..=b'9' | b'-' | b'+' | b'.' | b'e' | b'E') + { + at += 1; + } + literals.push(raw[start..at].to_string()); + } + _ => at += 1, + } + } + Self { literals, at: 0 } + } + + fn next(&mut self) -> Option<&str> { + let literal = self.literals.get(self.at).map(String::as_str); + self.at += 1; + literal + } + + /// How many literals were collected, for the test that keeps this in step with a walk. + pub fn len(&self) -> usize { + self.literals.len() + } + + pub fn is_empty(&self) -> bool { + self.literals.is_empty() + } +} + +/// Render a state, description or criteria that may be a string, object or array as +/// plain text. Objects become `key: value` fields — top-level fields separated by a blank +/// line, nested ones indented — and arrays become one `- item` line each. Key order is +/// preserved. +pub fn to_text(x: &Value) -> String { + render(x, 0, &mut NumberLiterals::default()) +} + +/// [`to_text`] for a JSON document, rendered from its text so an integer keeps every +/// digit. See [`NumberLiterals`] for why the text has to come along. +pub fn to_text_json(raw: &str) -> Result { + let value: Value = serde_json::from_str(raw).context("render a JSON document")?; + Ok(render(&value, 0, &mut NumberLiterals::of(raw))) +} + +fn render(x: &Value, indent: usize, numbers: &mut NumberLiterals) -> String { + match x { + Value::Null => String::new(), + Value::String(s) => s.clone(), + Value::Bool(true) => "true".to_string(), + Value::Bool(false) => "false".to_string(), + Value::Number(n) => match numbers.next() { + // The literal decides which of `json.loads`'s two types this is, not the + // value: a point or an exponent makes it a float, and anything else is an + // arbitrary-precision `int` that `str` prints digit for digit. + Some(literal) if !literal.contains(['.', 'e', 'E']) => { + // `json.loads("-0")` is the int 0 and `str(0)` is "0". JSON forbids + // leading zeros and a leading "+", so this is the only integer spelling + // `str(int(...))` would change. + if literal == "-0" { + "0".to_string() + } else { + literal.to_string() + } + } + Some(_) => python_float(n.as_f64().expect("a JSON number is an integer or a float")), + // No literal: the caller had only a `Value`, so `serde_json`'s integer types + // are all that is left to tell the two apart. Digits beyond a `u64` are + // already gone by then, which is why the request path carries the text. + None => match n.as_i64() { + Some(i) => i.to_string(), + None => match n.as_u64() { + Some(u) => u.to_string(), + None => { + python_float(n.as_f64().expect("a JSON number is an integer or a float")) + } + }, + }, + }, + Value::Object(map) => { + let pad = " ".repeat(indent); + let parts: Vec = map + .iter() + .map(|(k, v)| { + if is_nonempty_container(v) { + format!("{pad}{k}:\n{}", render(v, indent + 2, numbers)) + } else { + format!("{pad}{k}: {}", render(v, indent, numbers)) + } + }) + .collect(); + parts.join(if indent == 0 { "\n\n" } else { "\n" }) + } + Value::Array(items) => { + let pad = " ".repeat(indent); + let parts: Vec = items + .iter() + .map(|v| { + if is_nonempty_container(v) { + format!("{pad}-\n{}", render(v, indent + 2, numbers)) + } else { + format!("{pad}- {}", render(v, indent, numbers)) + } + }) + .collect(); + parts.join("\n") + } + } +} + +/// `str(float)` as CPython writes it, which is what the reference's `to_text` produces. +/// +/// Rust and Python disagree on both ends of the range: an integral float keeps its `.0` +/// here but not in `serde_json`, and a float outside `[1e-4, 1e16)` is exponent form with +/// a signed, at-least-two-digit exponent (`1e-05`, `1e+16`). The heads are trained on +/// these strings, so a differently spelled number is a different input. +fn python_float(x: f64) -> String { + // The shortest representation that round-trips, which is what `repr` uses. Rust's + // shortest digits agree with CPython's on *how many* there are but not always on the + // last one: for a value exactly between two equally short decimals CPython rounds to + // even and Rust away from zero. Re-rendering at that precision uses the same correct + // rounding as CPython's `%.*e`, so the digits agree. + let shortest = format!("{x:e}"); + let (shortest_mantissa, _) = shortest + .split_once('e') + .expect("a scientific format always writes an exponent"); + let precision = shortest_mantissa + .chars() + .filter(char::is_ascii_digit) + .count() + - 1; + let rendered = format!("{:.*e}", precision, x); + let (mantissa, exponent) = rendered + .split_once('e') + .expect("`{:e}` always writes an exponent"); + let exponent: i32 = exponent.parse().expect("`{:e}` writes a plain exponent"); + let sign = if mantissa.starts_with('-') { "-" } else { "" }; + let digits: String = mantissa.chars().filter(char::is_ascii_digit).collect(); + // CPython writes the exponent when the decimal point lands at or before -4, or past + // 16 digits. + let point = exponent + 1; + let body = if point <= -4 || point > 16 { + let (head, tail) = digits.split_at(1); + let e = point - 1; + let (esign, e) = if e < 0 { ('-', -e) } else { ('+', e) }; + if tail.is_empty() { + format!("{head}e{esign}{e:02}") + } else { + format!("{head}.{tail}e{esign}{e:02}") + } + } else if point <= 0 { + format!("0.{}{digits}", "0".repeat(-point as usize)) + } else if point as usize >= digits.len() { + format!("{digits}{}.0", "0".repeat(point as usize - digits.len())) + } else { + let (head, tail) = digits.split_at(point as usize); + format!("{head}.{tail}") + }; + format!("{sign}{body}") +} + +fn is_nonempty_container(v: &Value) -> bool { + match v { + Value::Object(m) => !m.is_empty(), + Value::Array(a) => !a.is_empty(), + _ => false, + } +} + +/// One decision, plus what the encoder reported spending on it. +#[derive(Debug, Clone)] +pub struct Decision { + /// Answers by question id, in request order. + pub answers: Vec<(String, Answer)>, + pub encoder_tokens: u64, +} + +/// The engine: heads, an encoder and the decision path. +pub struct Engine { + pub heads: Heads, + pub encoder: E, +} + +impl Engine { + pub fn new(heads: Heads, encoder: E) -> Self { + Self { heads, encoder } + } + + /// Answer every question in the request. + /// + /// Every question's state text and every candidate text is embedded in one pass, the + /// way the reference batches them, so the encoder sees one request per decision. + pub fn decide(&self, request: &Request, temperature: f32) -> Result { + ensure!( + temperature > 0.0 && temperature <= 100.0, + "temperature must be in (0, 100]" + ); + let prepared = request.prepare()?; + + let mut texts: Vec = Vec::new(); + for p in &prepared { + texts.push(p.state_text.clone()); + texts.extend(p.candidate_texts.iter().cloned()); + } + let (vectors, encoder_tokens) = self.encoder.embed(&texts)?; + ensure!( + vectors.len() == texts.len(), + "encoder returned {} vectors for {} texts", + vectors.len(), + texts.len() + ); + + // The vectors came back in the order the texts were sent: one state row followed + // by that question's candidate rows, per question. + let mut rows = vectors.into_iter(); + let mut answers = Vec::with_capacity(prepared.len()); + for p in &prepared { + let state = rows.next().context("missing state vector")?; + let mut candidates = Vec::with_capacity(p.candidate_texts.len()); + for _ in &p.candidate_texts { + candidates.push(rows.next().context("missing candidate vector")?); + } + let probs = scoring::distribution(&self.heads, &state, &candidates, temperature) + .with_context(|| format!("question {:?}", p.id))?; + answers.push(( + p.id.clone(), + scoring::answer(&p.question, &p.candidate_texts, &probs)?, + )); + } + + Ok(Decision { + answers, + encoder_tokens, + }) + } +} + +/// The `answers` object of a response, in request order. +pub fn answers_json(answers: &[(String, Answer)]) -> Value { + let mut out = Map::new(); + for (id, answer) in answers { + out.insert(id.clone(), answer_json(answer)); + } + Value::Object(out) +} + +/// The level key to level text map a `score` answer carries, in answer order. +fn legend_json(legend: &[(String, String)]) -> Value { + let mut map = Map::new(); + for (key, text) in legend { + map.insert(key.clone(), Value::String(text.clone())); + } + Value::Object(map) +} + +fn probabilities(pairs: &[(String, f32)]) -> Value { + let mut map = Map::new(); + for (k, p) in pairs { + map.insert(k.clone(), serde_json::json!(p)); + } + Value::Object(map) +} + +/// One answer in the shape `client.py` parses. +pub fn answer_json(answer: &Answer) -> Value { + match answer { + Answer::Choice { + choice, + confidence, + probabilities: p, + } => serde_json::json!({ + "type": "choice", + "choice": choice, + "confidence": confidence, + "probabilities": probabilities(p), + }), + Answer::Noul { noul } => serde_json::json!({"type": "noul", "noul": noul}), + Answer::Score { + score, + confidence, + legend, + probabilities: p, + } => serde_json::json!({ + "type": "score", + "score": score, + "confidence": confidence, + "legend": legend_json(legend), + "probabilities": probabilities(p), + }), + } +} diff --git a/src/models/clm/src/weights.rs b/src/models/clm/src/weights.rs new file mode 100644 index 0000000..5916016 --- /dev/null +++ b/src/models/clm/src/weights.rs @@ -0,0 +1,185 @@ +//! The CLM head tensors: inventory, loading and the expected shapes. +use anyhow::{Context, Result, ensure}; +use memmap2::Mmap; +use safetensors::{Dtype, SafeTensors}; +use std::{collections::BTreeSet, fs::File, path::Path}; + +use crate::config::{Config, HeadConfig}; + +/// The two projection heads. +pub const HEADS: [&str; 2] = ["state_head", "action_head"]; + +/// One head's parameters, all FP32 and row-major. +#[derive(Debug, Clone)] +pub struct Head { + pub inp_weight: Vec, + pub inp_bias: Vec, + pub hidden_weight: Vec, + pub hidden_bias: Vec, + /// Absent when the configuration has no hidden blocks or sets `layernorm: false`; + /// the LayerNorm weight when present. + pub norm_weight: Option>, + pub norm_bias: Option>, + pub out_weight: Vec, + pub out_bias: Vec, +} + +pub struct Heads { + pub state: Head, + pub action: Head, + pub config: Config, +} + +pub struct Weights { + data: Mmap, +} + +impl Weights { + /// The checkpoint must remain immutable while the mapping exists. + pub fn open(path: &Path) -> Result { + let file = File::open(path).with_context(|| format!("open {}", path.display()))?; + // SAFETY: model files are read-only inputs; no mutable mapping is created. + let data = unsafe { Mmap::map(&file)? }; + SafeTensors::deserialize(&data)?; + Ok(Self { data }) + } + + fn tensors(&self) -> Result> { + Ok(SafeTensors::deserialize(&self.data)?) + } + + /// The safetensors metadata, which carries the head configuration. + /// + /// `SafeTensors` keeps its metadata private, so this reads the header directly from + /// the mapping. No tensor data is touched. + pub fn metadata(&self) -> Result>> { + let (_, header) = + SafeTensors::read_metadata(&self.data).context("read the safetensors header")?; + Ok(header.metadata().clone()) + } + + /// Reject omitted, extra or duplicate names in the expected inventory. + pub fn validate_names<'a>(&self, names: impl IntoIterator) -> Result<()> { + let tensors = self.tensors()?; + let mut expected = BTreeSet::new(); + for name in names { + ensure!(expected.insert(name), "duplicate expected tensor: {name}"); + } + let actual: BTreeSet<&str> = tensors.names().into_iter().collect(); + if expected != actual { + let joined = |set: BTreeSet<&str>| set.into_iter().collect::>().join(", "); + let missing = joined(expected.difference(&actual).copied().collect()); + let extra = joined(actual.difference(&expected).copied().collect()); + anyhow::bail!( + "tensor inventory does not match checkpoint: missing [{missing}], unexpected [{extra}]" + ); + } + Ok(()) + } + + /// One tensor as FP32, checking the shape first so a mismatch is an error rather + /// than a reinterpreted buffer. + pub fn f32(&self, name: &str, shape: &[usize]) -> Result> { + let tensors = self.tensors()?; + let t = tensors + .tensor(name) + .with_context(|| format!("read tensor {name}"))?; + ensure!( + t.shape() == shape, + "{name}: expected {shape:?}, got {:?}", + t.shape() + ); + let out = match t.dtype() { + Dtype::F32 => t + .data() + .as_chunks::<4>() + .0 + .iter() + .map(|b| f32::from_le_bytes(*b)) + .collect(), + dt => anyhow::bail!("{name}: unsupported dtype {dt:?}, the export writes FP32"), + }; + Ok(out) + } +} + +/// Names and shapes of every tensor the heads are built from. +pub fn head_tensors(cfg: &HeadConfig) -> Result)>> { + let blocks = cfg.hidden_blocks()?; + let (h, w, p) = (cfg.hidden_size, cfg.width, cfg.projection_dim); + let mut out = Vec::new(); + for head in HEADS { + out.push((format!("{head}.inp.weight"), vec![w, h])); + out.push((format!("{head}.inp.bias"), vec![w])); + for i in 0..blocks { + out.push((format!("{head}.hidden.{i}.weight"), vec![w, w])); + out.push((format!("{head}.hidden.{i}.bias"), vec![w])); + if cfg.layernorm { + out.push((format!("{head}.norms.{i}.weight"), vec![w])); + out.push((format!("{head}.norms.{i}.bias"), vec![w])); + } + } + out.push((format!("{head}.out.weight"), vec![p, w])); + out.push((format!("{head}.out.bias"), vec![p])); + } + Ok(out) +} + +fn load_head(weights: &Weights, name: &str, cfg: &HeadConfig) -> Result { + let (h, w, p) = (cfg.hidden_size, cfg.width, cfg.projection_dim); + let blocks = cfg.hidden_blocks()?; + ensure!( + blocks <= 1, + "the published checkpoint has {blocks} hidden blocks; the loader implements one" + ); + let (hidden_weight, hidden_bias, norm_weight, norm_bias) = if blocks == 1 { + // `head_tensors` only lists `norms.N.*` when `cfg.layernorm` is set, so a + // `layernorm: false` checkpoint is complete without them and loads without them. + let (norm_weight, norm_bias) = if cfg.layernorm { + ( + Some(weights.f32(&format!("{name}.norms.0.weight"), &[w])?), + Some(weights.f32(&format!("{name}.norms.0.bias"), &[w])?), + ) + } else { + (None, None) + }; + ( + weights.f32(&format!("{name}.hidden.0.weight"), &[w, w])?, + weights.f32(&format!("{name}.hidden.0.bias"), &[w])?, + norm_weight, + norm_bias, + ) + } else { + (Vec::new(), Vec::new(), None, None) + }; + Ok(Head { + inp_weight: weights.f32(&format!("{name}.inp.weight"), &[w, h])?, + inp_bias: weights.f32(&format!("{name}.inp.bias"), &[w])?, + hidden_weight, + hidden_bias, + norm_weight, + norm_bias, + out_weight: weights.f32(&format!("{name}.out.weight"), &[p, w])?, + out_bias: weights.f32(&format!("{name}.out.bias"), &[p])?, + }) +} + +impl Heads { + /// Read the config from the metadata, check the inventory, and load every tensor. + pub fn load(weights: &Weights) -> Result { + let config = Config::from_metadata(weights.metadata()?.as_ref())?; + let expected = head_tensors(&config.head)?; + weights.validate_names(expected.iter().map(|(n, _)| n.as_str()))?; + let mut loaded = Vec::new(); + for head in HEADS { + loaded.push(load_head(weights, head, &config.head)?); + } + let action = loaded.pop().expect("two heads"); + let state = loaded.pop().expect("two heads"); + Ok(Self { + state, + action, + config, + }) + } +} diff --git a/tests/clm/checkpoint.rs b/tests/clm/checkpoint.rs new file mode 100644 index 0000000..778b60b --- /dev/null +++ b/tests/clm/checkpoint.rs @@ -0,0 +1,263 @@ +//! CPU checks for the CLM head loader. No GPU, no encoder server, but the frozen export. +use omni_clm::{ + Config, HashingEncoder, Kind, Question, Weights, answer, confidence, distribution, + head_tensors, weights::Heads, +}; +use std::path::PathBuf; + +fn export_dir() -> PathBuf { + PathBuf::from(std::env::var_os("CLM_EXPORT").expect("set CLM_EXPORT to the export directory")) +} + +fn load() -> Heads { + let weights = Weights::open(&export_dir().join("model.safetensors")).unwrap(); + Heads::load(&weights).unwrap() +} + +#[test] +#[ignore = "requires CLM_EXPORT at a converted checkpoint; CPU only"] +fn every_tensor_conversion_matches_the_oracle() { + use sha2::{Digest, Sha256}; + + let dir = export_dir(); + let oracle: Vec = + serde_json::from_slice(&std::fs::read(dir.join("oracle.json")).unwrap()).unwrap(); + let weights = Weights::open(&dir.join("model.safetensors")).unwrap(); + + assert_eq!(oracle.len(), 16, "oracle must cover both heads"); + for row in oracle { + let name = row["name"].as_str().unwrap(); + let shape: Vec = serde_json::from_value(row["shape"].clone()).unwrap(); + let values = weights.f32(name, &shape).unwrap(); + let bytes: Vec = values.iter().flat_map(|v| v.to_le_bytes()).collect(); + assert_eq!( + format!("{:x}", Sha256::digest(&bytes)), + row["f32"].as_str().unwrap(), + "{name} f32" + ); + } +} + +#[test] +#[ignore = "requires CLM_EXPORT at a converted checkpoint; CPU only"] +fn decisions_match_the_reference_implementation() { + let heads = load(); + let oracle: serde_json::Value = + serde_json::from_slice(&std::fs::read(export_dir().join("head-oracle.json")).unwrap()) + .unwrap(); + + for case in oracle["cases"].as_array().unwrap() { + let name = case["name"].as_str().unwrap(); + let keys: Vec = serde_json::from_value(case["keys"].clone()).unwrap(); + let temperature = case["temperature"].as_f64().unwrap() as f32; + let expected: Vec = serde_json::from_value(case["probabilities"].clone()).unwrap(); + + let state = + HashingEncoder::vector(&format!("state::{name}"), heads.config.head.hidden_size); + let candidates: Vec> = keys + .iter() + .map(|k| { + HashingEncoder::vector(&format!("cand::{name}::{k}"), heads.config.head.hidden_size) + }) + .collect(); + + let probs = distribution(&heads, &state, &candidates, temperature).unwrap(); + assert_eq!(probs.len(), expected.len(), "{name}: length"); + for (got, want) in probs.iter().zip(&expected) { + assert!( + (got - want).abs() < 1e-4, + "{name}: probability {got} vs {want}" + ); + } + + let kind = match case["kind"].as_str().unwrap() { + "choice" => Kind::Choice, + "noul" => Kind::Noul, + _ => Kind::Score, + }; + let question = Question { + id: name.to_string(), + kind, + keys, + }; + // The oracle carries keys and probabilities, not candidate texts; the legend a + // score answer is built from is checked in `text.rs`, where `candidates` supplies + // the texts. + let answer = answer(&question, &question.keys, &probs).unwrap(); + match (&answer, case.get("choice")) { + ( + omni_clm::Answer::Choice { + choice, + confidence: c, + .. + }, + Some(want), + ) => { + assert_eq!(choice, want.as_str().unwrap(), "{name}: choice"); + let want_c = case["confidence"].as_f64().unwrap() as f32; + assert!( + (c - want_c).abs() < 1e-4, + "{name}: confidence {c} vs {want_c}" + ); + } + ( + omni_clm::Answer::Score { + score, + confidence: c, + .. + }, + None, + ) => { + let want_s = case["score"].as_f64().unwrap() as f32; + assert!( + (score - want_s).abs() < 1e-4, + "{name}: score {score} vs {want_s}" + ); + let want_c = case["confidence"].as_f64().unwrap() as f32; + assert!( + (c - want_c).abs() < 1e-4, + "{name}: confidence {c} vs {want_c}" + ); + } + (omni_clm::Answer::Noul { noul }, None) => { + let want = case["noul"].as_f64().unwrap() as f32; + assert!((noul - want).abs() < 1e-4, "{name}: noul {noul} vs {want}"); + } + (other, _) => panic!("{name}: unexpected answer {other:?}"), + } + } +} + +#[test] +fn confidence_matches_the_reference_definition() { + // Top minus the mean of the rest, clamped; a single candidate is fully decided. + assert_eq!(confidence(&[1.0]), 1.0); + assert!((confidence(&[0.75, 0.25]) - 0.5).abs() < 1e-6); + assert!((confidence(&[0.5, 0.3, 0.2]) - 0.25).abs() < 1e-6); + assert_eq!(confidence(&[0.4, 0.4, 0.4]), 0.0); +} + +/// `layernorm: false` is a configuration `head_tensors` honours, so the loader must too: +/// a checkpoint that declares no LayerNorm is complete without `norms.*`, and one that +/// declares it is still rejected when the tensors are absent. +#[test] +fn the_loader_follows_the_layernorm_flag() { + let dir = std::env::temp_dir().join("omni-clm-synthetic-heads"); + std::fs::create_dir_all(&dir).unwrap(); + + let plain = synthetic(&dir.join("no-layernorm.safetensors"), false, false); + let heads = Heads::load(&Weights::open(&plain).unwrap()).unwrap(); + assert!(heads.state.norm_weight.is_none()); + assert!(heads.action.norm_weight.is_none()); + assert_eq!(head_tensors(&heads.config.head).unwrap().len(), 12); + + let normed = synthetic(&dir.join("layernorm.safetensors"), true, true); + let heads = Heads::load(&Weights::open(&normed).unwrap()).unwrap(); + assert!(heads.state.norm_weight.is_some()); + assert_eq!(head_tensors(&heads.config.head).unwrap().len(), 16); + + let claiming = synthetic(&dir.join("claims-layernorm.safetensors"), true, false); + let Err(err) = Heads::load(&Weights::open(&claiming).unwrap()) else { + panic!("a checkpoint that declares layernorm but omits norms.* was accepted"); + }; + assert!(err.to_string().contains("norms.0.weight"), "{err}"); +} + +/// A two-head checkpoint small enough to write here, with the LayerNorm tensors present +/// or absent independently of what the configuration declares. +fn synthetic(path: &std::path::Path, layernorm: bool, include_norms: bool) -> PathBuf { + use safetensors::Dtype; + use safetensors::tensor::{TensorView, serialize_to_file}; + use std::collections::HashMap; + + let metadata: HashMap = HashMap::from([ + ("format".to_string(), "clm-heads".to_string()), + ( + "cfg".to_string(), + format!( + r#"{{"hidden_size":4,"projection_dim":2,"width":3,"depth":3,"activation":"gelu","layernorm":{layernorm},"residual":false}}"# + ), + ), + ("hidden_size".to_string(), "4".to_string()), + ("projection_dim".to_string(), "2".to_string()), + ("logit_scale".to_string(), "1.0".to_string()), + ]); + + let (mut names, mut shapes, mut buffers) = (Vec::new(), Vec::new(), Vec::new()); + for head in ["state_head", "action_head"] { + let mut spec = vec![ + (format!("{head}.inp.weight"), vec![3, 4]), + (format!("{head}.inp.bias"), vec![3]), + (format!("{head}.hidden.0.weight"), vec![3, 3]), + (format!("{head}.hidden.0.bias"), vec![3]), + (format!("{head}.out.weight"), vec![2, 3]), + (format!("{head}.out.bias"), vec![2]), + ]; + if include_norms { + spec.push((format!("{head}.norms.0.weight"), vec![3])); + spec.push((format!("{head}.norms.0.bias"), vec![3])); + } + for (name, shape) in spec { + buffers.push(vec![0u8; shape.iter().product::() * 4]); + shapes.push(shape); + names.push(name); + } + } + let tensors: Vec<(String, TensorView)> = names + .iter() + .zip(&shapes) + .zip(&buffers) + .map(|((name, shape), buffer)| { + ( + name.clone(), + TensorView::new(Dtype::F32, shape.clone(), buffer).unwrap(), + ) + }) + .collect(); + + serialize_to_file(tensors, Some(metadata), path).unwrap(); + path.to_path_buf() +} + +/// The exporter writes the geometry at the top level and may leave it out of `cfg`, so a +/// `cfg` without it has to be completed before deserialization rather than after. +#[test] +fn the_top_level_geometry_fills_in_a_cfg_that_omits_it() { + use std::collections::HashMap; + + let geom = r#""width":3,"depth":3,"activation":"gelu","layernorm":false,"residual":false"#; + let meta = |cfg: String| -> HashMap { + HashMap::from([ + ("format".to_string(), "clm-heads".to_string()), + ("cfg".to_string(), cfg), + ("hidden_size".to_string(), "4".to_string()), + ("projection_dim".to_string(), "2".to_string()), + ("logit_scale".to_string(), "1.0".to_string()), + ]) + }; + + let omitted = Config::from_metadata(Some(&meta(format!("{{{geom}}}")))).unwrap(); + assert_eq!(omitted.head.hidden_size, 4); + assert_eq!(omitted.head.projection_dim, 2); + + let both = meta(format!(r#"{{"hidden_size":4,"projection_dim":2,{geom}}}"#)); + assert_eq!(Config::from_metadata(Some(&both)).unwrap().head.width, 3); + + let mut neither = meta(format!("{{{geom}}}")); + neither.remove("hidden_size"); + neither.remove("projection_dim"); + assert!(Config::from_metadata(Some(&neither)).is_err()); +} + +/// `schema.label_of` is `max(p, key=p.__getitem__)`, which keeps the first of several +/// equal values. +#[test] +fn a_tied_score_label_keeps_the_first_level() { + let question = Question { + id: "tie".into(), + kind: Kind::Score, + keys: vec!["0".into(), "1".into(), "2".into()], + }; + let tied = answer(&question, &question.keys, &[0.5, 0.5, 0.0]).unwrap(); + assert_eq!(tied.label(), "0"); +} diff --git a/tests/clm/embedding.rs b/tests/clm/embedding.rs new file mode 100644 index 0000000..d8d0d25 --- /dev/null +++ b/tests/clm/embedding.rs @@ -0,0 +1,31 @@ +//! Tests for the embeddings client, wired in through `#[path]` from +//! `src/models/clm/src/embedding.rs`. +//! +//! `HttpEncoder::body` is private, so CONTRIBUTING's rule for private items +//! applies: the `#[cfg(test)]` and the module path stay in `src/`, the test body +//! lives here. +use super::*; + +/// `embedder.py` sends its `max_tokens` as `truncate_prompt_tokens`; a request without +/// it is truncated differently on the server, or rejected instead of truncated. +#[test] +fn the_request_carries_the_truncation_limit() { + let encoder = + HttpEncoder::new("http://127.0.0.1:1/v1/embeddings", "m", Duration::ZERO).unwrap(); + let texts = vec!["a".to_string()]; + assert_eq!(encoder.body(&texts)["truncate_prompt_tokens"], 2048); + + let unlimited = encoder.with_max_tokens(None); + assert!( + unlimited + .body(&texts) + .get("truncate_prompt_tokens") + .is_none() + ); + + let custom = HttpEncoder::new("http://127.0.0.1:1/v1/embeddings", "m", Duration::ZERO) + .unwrap() + .with_max_tokens(Some(64)); + assert_eq!(custom.body(&texts)["truncate_prompt_tokens"], 64); + assert_eq!(custom.body(&texts)["encoding_format"], "base64"); +} diff --git a/tests/clm/text.rs b/tests/clm/text.rs new file mode 100644 index 0000000..984347f --- /dev/null +++ b/tests/clm/text.rs @@ -0,0 +1,472 @@ +//! The text the two heads see, checked byte-for-byte against the reference. +//! +//! `state_text`, `candidates` and `to_text` decide what reaches the encoder, and getting +//! one separator wrong shifts every probability without failing anything. The oracle is +//! produced by the reference implementation itself +//! (`recipe/clm/native/text_oracle.py`, which imports `clm.schema`), so this compares +//! against the real thing rather than a transcription of it. +use omni_clm::serve::{answer_json, candidates, state_text, to_text}; +use omni_clm::{Kind, Question, Request, answer, serve::QuestionRequest}; +use omni_clm::{NumberLiterals, to_text_json}; +use serde_json::{Map, Value, json}; + +fn oracle() -> Value { + let path = std::env::var_os("CLM_TEXT_ORACLE") + .expect("set CLM_TEXT_ORACLE to the file text_oracle.py writes"); + serde_json::from_slice(&std::fs::read(path).unwrap()).unwrap() +} + +fn states() -> Vec { + vec![ + json!("I was charged twice."), + json!({"body": "Charged twice", "order": 4411, "urgent": true}), + json!({"ticket": {"id": 7, "tags": ["a", "b"]}, "note": null}), + json!([{"k": 1}, {"k": 2}]), + json!({"empty_obj": {}, "empty_arr": [], "n": 0.5}), + json!({"nested": {"deep": {"x": "y"}}}), + // `str(float)` switches to exponent form outside [1e-4, 1e16) and keeps a `.0` on + // an integral float, so these pin the number spelling the reference produces. + json!({"tiny": 1e-5, "smaller": 1e-7, "edge": 1e-4, "round": 1e15, "huge": 1e16, "neg": -1e-6}), + // Parsed from the literal text rather than built from an `f64`, because + // the parse is the part that was wrong. + serde_json::from_str(r#"{"seventeen": 7.8190461323667115, "inexact": 9007199254740993.0}"#) + .unwrap(), + serde_json::from_str(concat!( + r#"{"big": 18446744073709551616, "#, + r#""huge": 340282366920938463463374607431768211456, "#, + r#""negzero": -0, "negbig": -18446744073709551616}"# + )) + .unwrap(), + ] +} + +/// The same states as [`states`], as the JSON text a request arrives as. +fn state_texts() -> Vec<&'static str> { + vec![ + r#""I was charged twice.""#, + r#"{"body": "Charged twice", "order": 4411, "urgent": true}"#, + r#"{"ticket": {"id": 7, "tags": ["a", "b"]}, "note": null}"#, + r#"[{"k": 1}, {"k": 2}]"#, + r#"{"empty_obj": {}, "empty_arr": [], "n": 0.5}"#, + r#"{"nested": {"deep": {"x": "y"}}}"#, + r#"{"tiny": 1e-5, "smaller": 1e-7, "edge": 1e-4, "round": 1e15, "huge": 1e16, "neg": -1e-6}"#, + r#"{"seventeen": 7.8190461323667115, "inexact": 9007199254740993.0}"#, + concat!( + r#"{"big": 18446744073709551616, "#, + r#""huge": 340282366920938463463374607431768211456, "#, + r#""negzero": -0, "negbig": -18446744073709551616}"# + ), + ] +} + +fn questions() -> Vec { + vec![ + QuestionRequest { + kind: Kind::Choice, + instructions: "Which team?".into(), + criteria: Some(json!({"billing": "Charges and refunds", "tech": "Software problems"})), + }, + QuestionRequest { + kind: Kind::Choice, + instructions: "Pick".into(), + criteria: Some(json!({"a": "", "b": null})), + }, + QuestionRequest { + kind: Kind::Score, + instructions: "How urgent?".into(), + criteria: Some(json!(["Not urgent", "Soon", "Now"])), + }, + QuestionRequest { + kind: Kind::Noul, + instructions: "Does the customer ask for a refund?".into(), + criteria: None, + }, + QuestionRequest { + kind: Kind::Noul, + instructions: "Refund?".into(), + criteria: Some(json!({"true": "Yes they do", "false": "No they do not"})), + }, + // An empty container renders to nothing, which is not the same as being absent. + QuestionRequest { + kind: Kind::Choice, + instructions: "Pick a bucket".into(), + criteria: Some(json!({"empty_obj": {}, "empty_list": []})), + }, + QuestionRequest { + kind: Kind::Noul, + instructions: "Is it so?".into(), + criteria: Some(json!({"true": {}, "false": []})), + }, + // Numbers are spelled by `str(float)` in criteria too, not only in the state. + QuestionRequest { + kind: Kind::Score, + instructions: "How much?".into(), + criteria: Some(json!([1e-5, 0.5, 1e16])), + }, + ] +} + +#[test] +#[ignore = "requires CLM_TEXT_ORACLE from recipe/clm/native/text_oracle.py; CPU only"] +fn to_text_matches_the_reference_byte_for_byte() { + let oracle = oracle(); + let expected = oracle["to_text"].as_array().unwrap(); + let got = states(); + assert_eq!( + got.len(), + expected.len(), + "the oracle was built from another case list" + ); + // All but the last: a `Value` cannot hold an integer above `u64::MAX`, so the digits + // are already gone before `to_text` is called. That case is the next loop's. + for (i, state) in got.iter().take(expected.len() - 1).enumerate() { + assert_eq!( + to_text(state), + expected[i].as_str().unwrap(), + "to_text case {i} for {state}" + ); + } + + // The same renderings again, from the text, which is the path a request takes and + // the only one that keeps an integer literal's digits. + for (i, raw) in state_texts().iter().enumerate() { + assert_eq!( + to_text_json(raw).unwrap(), + expected[i].as_str().unwrap(), + "to_text_json case {i}" + ); + } +} + +#[test] +#[ignore = "requires CLM_TEXT_ORACLE from recipe/clm/native/text_oracle.py; CPU only"] +fn state_text_and_candidates_match_the_reference_byte_for_byte() { + let oracle = oracle(); + let cases = oracle["cases"].as_array().unwrap(); + let states = state_texts(); + let questions = questions(); + assert_eq!( + cases.len(), + states.len() * questions.len(), + "the oracle was built from another case list" + ); + + let mut i = 0; + for raw in &states { + for q in &questions { + let case = &cases[i]; + // Through the request path rather than `state_text`/`candidates` on a + // `Value`: the text is what keeps an integer literal's digits, and a `Value` + // has already lost the ones above `u64::MAX`. + let line = format!( + r#"{{"state": {raw}, "questions": {{"q": {}}}}}"#, + question_json(q) + ); + let prepared = Request::parse_line(&line).unwrap().prepare().unwrap(); + assert_eq!( + prepared[0].state_text, + case["state_text"].as_str().unwrap(), + "case {i} state_text" + ); + assert_eq!( + prepared[0].question.keys, + case["keys"] + .as_array() + .unwrap() + .iter() + .map(|v| v.as_str().unwrap().to_string()) + .collect::>(), + "case {i} keys" + ); + assert_eq!( + prepared[0].candidate_texts, + case["candidate_texts"] + .as_array() + .unwrap() + .iter() + .map(|v| v.as_str().unwrap().to_string()) + .collect::>(), + "case {i} candidate_texts" + ); + i += 1; + } + } +} + +/// A question as the JSON text a request carries it in, for the loop above. +fn question_json(q: &QuestionRequest) -> Value { + let mut object = Map::new(); + object.insert( + "type".to_string(), + json!(match q.kind { + Kind::Choice => "choice", + Kind::Score => "score", + Kind::Noul => "noul", + }), + ); + object.insert("instructions".to_string(), json!(q.instructions)); + if let Some(criteria) = &q.criteria { + object.insert("criteria".to_string(), criteria.clone()); + } + Value::Object(object) +} + +#[test] +fn text_construction_handles_the_edges_the_oracle_does_not() { + // No state in the oracle is empty, and no question leaves `instructions` blank. + assert_eq!(state_text(&json!("context"), " "), "context"); + assert_eq!(state_text(&json!(""), "question"), "question"); + // A state that is a bare scalar is not a container, so it renders unindented. + assert_eq!(to_text(&json!(true)), "true"); + assert_eq!(to_text(&json!(null)), ""); +} + +#[test] +fn a_score_answer_carries_its_level_legend() { + // `answer_from_probs` returns the criteria text next to the score so a consumer can + // read a level back without the request; the serialized answer must not drop it. + let request = QuestionRequest { + kind: Kind::Score, + instructions: "How urgent?".into(), + criteria: Some(json!(["Not urgent", "Soon", "Now"])), + }; + let (keys, texts) = candidates(&request).unwrap(); + let question = Question { + id: "urgency".into(), + kind: Kind::Score, + keys, + }; + let answer = answer(&question, &texts, &[0.1, 0.7, 0.2]).unwrap(); + let json = answer_json(&answer); + assert_eq!(json["type"], "score"); + assert_eq!( + json["legend"], + json!({"0": "Not urgent", "1": "Soon", "2": "Now"}) + ); + assert!((json["score"].as_f64().unwrap() - 1.1).abs() < 1e-6); +} + +/// The whole path: a JSON literal to the string the heads see. +/// +/// `serde_json`'s default float parsing is not correctly rounded, so a literal +/// with more digits than a double holds can land on the neighbouring double and +/// render as a different number — `7.8190461323667115` came out +/// `7.819046132366712`, which is a different embedding input. Formatting an +/// already-parsed float correctly does not help; the parse has to be right too, +/// and `float_roundtrip` is what makes it match `json.loads`. +#[test] +fn json_numbers_parse_to_the_same_doubles_as_the_reference() { + // Both columns are what CPython's `json.loads` then `str` produce. + for (literal, bits, text) in [ + ( + "7.8190461323667115", + 0x401f_46b4_0781_b8a4, + "7.8190461323667115", + ), + ( + "9007199254740993.0", + 0x4340_0000_0000_0000, + "9007199254740992.0", + ), + ("0.1", 0x3fb9_9999_9999_999a, "0.1"), + ( + "3.141592653589793238", + 0x4009_21fb_5444_2d18, + "3.141592653589793", + ), + ] { + let parsed: Value = serde_json::from_str(literal).unwrap(); + assert_eq!( + parsed.as_f64().unwrap().to_bits(), + bits, + "parsing {literal}" + ); + assert_eq!(to_text(&parsed), text, "rendering {literal}"); + } +} + +/// The whole path in one test: a request line as `clm-run` reads it, through +/// parsing and preparation, to the exact strings the encoder is asked for. +/// +/// This is the layer the number bug lived at. Checking that `to_text` renders a +/// double correctly says nothing if the double was already the wrong one, so the +/// literals go in as JSON text and the assertion is on what comes out the far end. +#[test] +fn a_request_line_reaches_the_encoder_with_the_reference_text() { + let line = concat!( + r#"{"model":"clm-latest","state":{"seventeen":7.8190461323667115,"#, + r#""inexact":9007199254740993.0,"big":18446744073709551616},"#, + r#""questions":{"q":{"type":"choice","instructions":"Pick","#, + r#""criteria":{"a":7.8190461323667115,"b":18446744073709551616}}}}"# + ); + let request = Request::parse_line(line).unwrap(); + let prepared = request.prepare().unwrap(); + + assert_eq!(prepared.len(), 1); + assert_eq!( + prepared[0].state_text, + "seventeen: 7.8190461323667115\n\ninexact: 9007199254740992.0\n\n\ + big: 18446744073709551616\n\nPick" + ); + assert_eq!(prepared[0].question.keys, ["a", "b"]); + assert_eq!( + prepared[0].candidate_texts, + ["7.8190461323667115", "18446744073709551616"] + ); +} + +/// Whole request bodies, as the text a caller sends them in. +/// +/// The loop above pairs every state with every question, but a `QuestionRequest` holds +/// `instructions` as a `String` and `criteria` as a `Value`, so it cannot express an +/// integer above `u64::MAX` or a statement that is a number rather than a string. These +/// are the requests that need the literal itself, and they go in as text. +#[test] +#[ignore = "requires CLM_TEXT_ORACLE from recipe/clm/native/text_oracle.py; CPU only"] +fn raw_requests_match_the_reference_byte_for_byte() { + let oracle = oracle(); + let cases = oracle["raw_cases"].as_array().unwrap(); + assert!( + !cases.is_empty(), + "the oracle was built from another case list" + ); + + for case in cases { + let line = case["line"].as_str().unwrap(); + let prepared = Request::parse_line(line).unwrap().prepare().unwrap(); + let expected = case["questions"].as_object().unwrap(); + assert_eq!(prepared.len(), expected.len(), "question count for {line}"); + for p in &prepared { + let want = &expected[&p.id]; + assert_eq!( + p.state_text, + want["state_text"].as_str().unwrap(), + "{} state_text for {line}", + p.id + ); + assert_eq!( + p.question.keys, + strings(&want["keys"]), + "{} keys for {line}", + p.id + ); + assert_eq!( + p.candidate_texts, + strings(&want["candidate_texts"]), + "{} candidate_texts for {line}", + p.id + ); + } + } +} + +fn strings(value: &Value) -> Vec { + value + .as_array() + .unwrap() + .iter() + .map(|v| v.as_str().unwrap().to_string()) + .collect() +} + +/// A `noul` question reads its two descriptions in `NOUL_KEYS` order — `false` first — +/// which is not the order they were written in. A literal has to be found by the key it +/// was written under, not by counting: `{"true": 1, "false": 2}` is `false: 2`, and +/// pairing the first literal with the first key gives `false: 1`, a different answer. +#[test] +fn a_noul_literal_follows_its_key_not_its_position() { + let line = concat!( + r#"{"state":"s","questions":{"q":{"type":"noul","instructions":"Is it so?","#, + r#""criteria":{"true":1,"false":2}}}}"# + ); + let prepared = Request::parse_line(line).unwrap().prepare().unwrap(); + assert_eq!(prepared[0].question.keys, ["false", "true"]); + assert_eq!(prepared[0].candidate_texts, ["false: 2", "true: 1"]); +} + +/// The default `noul` candidates are the question's own statement, so a literal in the +/// statement has to reach them as written rather than as a rounded double. The statement +/// is a value like any other here, not necessarily a string. +#[test] +fn a_default_candidate_keeps_the_instruction_literal() { + let line = concat!( + r#"{"state":"s","questions":{"q":{"type":"noul","#, + r#""instructions":18446744073709551616,"criteria":null}}}"# + ); + let prepared = Request::parse_line(line).unwrap().prepare().unwrap(); + assert_eq!( + prepared[0].candidate_texts, + [ + "false: No. This is false: 18446744073709551616", + "true: Yes. This is true: 18446744073709551616", + ] + ); + // The same statement reaches the state head, so both heads see one text. + assert_eq!( + prepared[0].state_text, "s\n\n18446744073709551616", + "the statement is rendered once and reused" + ); +} + +/// Integer literals are arbitrary precision in `json.loads`, so they have to keep every +/// digit here too. +/// +/// A JSON integer above `u64::MAX` has no exact `f64`, and `serde_json` parses one +/// straight to a double: `18446744073709551616` was reaching the encoder as +/// `1.8446744073709552e+19`. The literal decides which of `json.loads`'s two types a +/// number is, not the value it holds. +#[test] +fn integer_literals_keep_every_digit() { + // The expected strings are `str(json.loads(literal))`. + for (literal, text) in [ + ("18446744073709551616", "18446744073709551616"), + ( + "340282366920938463463374607431768211456", + "340282366920938463463374607431768211456", + ), + ("-18446744073709551616", "-18446744073709551616"), + ("0", "0"), + ("-0", "0"), + ("9007199254740993", "9007199254740993"), + // The same digits as a float literal, so it is a float afterwards. + ("9007199254740993.0", "9007199254740992.0"), + // And an exponent is a float however integral it looks. + ("1e5", "100000.0"), + ] { + assert_eq!(to_text_json(literal).unwrap(), text, "rendering {literal}"); + } +} + +/// The lexer and the value walk have to agree on how many numbers a document has, or a +/// literal would be paired with the wrong number. Both are in document order, and this +/// is what says so. +#[test] +fn the_literals_line_up_with_the_values() { + fn numbers(value: &Value) -> usize { + match value { + Value::Number(_) => 1, + Value::Array(items) => items.iter().map(numbers).sum(), + Value::Object(map) => map.values().map(numbers).sum(), + _ => 0, + } + } + for raw in state_texts() { + let value: Value = serde_json::from_str(raw).unwrap(); + let literals = NumberLiterals::of(raw); + assert_eq!(literals.len(), numbers(&value), "counting {raw}"); + } + for raw in [ + r#"{"a": "not a 1 or a 2", "b": [3, {"c": -4.5e-6}], "d": null}"#, + "{\"escaped\": \"quote \\\" then 7 and \\\\ then 8\", \"n\": 9}", + r#"[1, [2, [3]], {"k": 4}]"#, + r#"{}"#, + r#"[]"#, + ] { + let value: Value = serde_json::from_str(raw).unwrap(); + assert_eq!( + NumberLiterals::of(raw).len(), + numbers(&value), + "counting {raw}" + ); + } +}