Download example.py from precisit/one-pass-sv-forms: direct link, hf CLI and curl.
- Browser
- Download file 4.46 kB
-
https://huggingface.co/precisit/one-pass-sv-forms/resolve/d3b7fcdcedf07b511069574f4f4cfc8ccd5564cb/example.py
- Command line
-
hf download hf://precisit/one-pass-sv-forms@d3b7fcdcedf07b511069574f4f4cfc8ccd5564cb/example.py
-
curl -L -o example.py https://huggingface.co/precisit/one-pass-sv-forms/resolve/d3b7fcdcedf07b511069574f4f4cfc8ccd5564cb/example.py
4.46 kB
| """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)] # raw logits; softmax them if you need probabilities | |
| 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 # type: ignore | |
| except ImportError: # pragma: no cover | |
| 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(): # when the script is run outside the repository | |
| 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() | |