| """Minimal inference example for precisit/sv0-forms. |
| |
| Two routes: the Core ML package (no PyTorch needed) and the PyTorch checkpoint through the |
| vendored Cua-S1 loader. Both take the same byte-level inputs; the encoder is: |
| |
| ids = utf-8 bytes, truncated to the limit, each byte + 1, zero-padded |
| context: 224 bytes | option: 96 bytes | up to 40 options (the export's ceiling) |
| |
| Run: python example.py (Core ML, needs coremltools) |
| python example.py --torch (PyTorch, needs torch + safetensors + the vendored package) |
| """ |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| from pathlib import Path |
|
|
| HERE = Path(__file__).resolve().parent |
| CONTEXT_BYTES, OPTION_BYTES, MAX_OPTIONS = 224, 96, 40 |
|
|
| TASK = "UPPGIFT fyll i formuläret från dokumentet och skicka sedan in" |
| FORM = "Exempelkliniken - Ny patientregistrering" |
| ELEMENT = ('Edit "E-post" value=""',) |
| OPTIONS = [ |
| "fyll Förnamn: Anna", |
| "fyll Efternamn: Lindqvist", |
| "fyll E-post: anna.lindqvist@exempel.invalid", |
| "fyll Telefon: 070-341 22 87", |
| "kryssa", |
| "klicka", |
| "hoppa över", |
| ] |
|
|
|
|
| def context_string() -> str: |
| role, label, state = ELEMENT[0], "E-post", 'value=""' |
| return f"{TASK}\nFORM {FORM}\nELEMENT {role} \"{label}\" {state}" |
|
|
|
|
| def byte_ids(text: str, length: int) -> list[int]: |
| return [byte + 1 for byte in text.encode("utf-8", errors="replace")[:length]] |
|
|
|
|
| def pad(ids: list[int], length: int) -> list[int]: |
| return ids + [0] * (length - len(ids)) |
|
|
|
|
| def via_coreml() -> None: |
| import coremltools as ct |
| import numpy as np |
|
|
| package = HERE / "coreml" / "sv0_forms_int8_options40.mlpackage" |
| model = ct.models.MLModel(str(package), compute_units=ct.ComputeUnit.CPU_AND_NE) |
| context = np.array([pad(byte_ids(context_string(), CONTEXT_BYTES), CONTEXT_BYTES)], dtype=np.int32) |
| options = np.zeros((1, MAX_OPTIONS, OPTION_BYTES), dtype=np.int32) |
| for index, option in enumerate(OPTIONS): |
| encoded = pad(byte_ids(option, OPTION_BYTES), OPTION_BYTES) |
| options[0, index] = np.array(encoded, dtype=np.int32) |
| mask = np.zeros((1, MAX_OPTIONS), dtype=np.int32) |
| mask[0, : len(OPTIONS)] = 1 |
| output = model.predict({"context_ids": context, "option_ids": options, "option_mask": mask}) |
| scores = output["logits"][0][: len(OPTIONS)] |
| best = int(scores.argmax()) |
| print(f"chosen: {OPTIONS[best]!r} (logit {float(scores[best]):.3f})") |
| print("logits:", {option: round(float(score), 3) for option, score in zip(OPTIONS, scores)}) |
|
|
|
|
| def via_torch() -> None: |
| import torch |
| from huggingface_hub import hf_hub_download |
|
|
| try: |
| from cua_s1.model import ChoiceExample, load_checkpoint, select_device |
| except ImportError: |
| raise SystemExit( |
| "The PyTorch route needs the vendored Cua-S1 loader: clone the toolkit repository " |
| "(precisit/one-pass-specialists) and put its `vendor/` directory on sys.path, or import this " |
| "repository's Core ML packages instead — they need no PyTorch." |
| ) |
|
|
| weights = HERE / "sv0-forms.safetensors" |
| if not weights.exists(): |
| weights = Path(hf_hub_download("precisit/one-pass-sv-forms", "sv0-forms.safetensors")) |
| hf_hub_download("precisit/one-pass-sv-forms", "sv0-forms.json", local_dir=weights.parent) |
| device = select_device("auto") |
| model, collator, _ = load_checkpoint(weights, device) |
| batch = collator([ChoiceExample(context=context_string(), options=tuple(OPTIONS), label=0)]) |
| with torch.no_grad(): |
| scores = model({key: value.to(device) for key, value in batch.items()})[0].softmax(-1) |
| best = int(scores.argmax()) |
| print(f"chosen: {OPTIONS[best]!r} (probability {float(scores[best]):.3f})") |
|
|
|
|
| def metadata() -> None: |
| print(json.dumps(json.loads((HERE / "sv0-forms.json").read_text(encoding="utf-8"))["metadata"], indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--torch", action="store_true", help="use the PyTorch checkpoint instead of Core ML") |
| parser.add_argument("--metadata", action="store_true", help="print the training metadata and exit") |
| arguments = parser.parse_args() |
| if arguments.metadata: |
| metadata() |
| elif arguments.torch: |
| via_torch() |
| else: |
| via_coreml() |
|
|