quantcost 0.2.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
edgellm/hub.py ADDED
@@ -0,0 +1,125 @@
1
+ """Resolve a (model, precision) pair to a local ONNX file, downloading if needed.
2
+
3
+ The default benchmark path deliberately consumes **pre-quantized ONNX already
4
+ published on the Hugging Face Hub** rather than quantizing locally. That is what
5
+ keeps the install light: a contributor needs ``onnxruntime`` and NumPy, not
6
+ PyTorch, ``transformers``, ``optimum`` and a GPU. It also keeps results
7
+ comparable, because every machine scores the *same bytes* of the same artifact
8
+ instead of its own locally-produced quantization.
9
+
10
+ The naming convention here is the one used across the ``onnx-community`` and
11
+ ``transformers.js`` model repos, so thousands of existing models work unchanged.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ import logging
17
+ from dataclasses import dataclass
18
+ from pathlib import Path
19
+
20
+ log = logging.getLogger(__name__)
21
+
22
+ #: Canonical precision label -> filename stem inside the repo's ``onnx/`` folder.
23
+ #: Keys are what the user types; values are the Hub convention.
24
+ PRECISION_FILES: dict[str, str] = {
25
+ "fp32": "model",
26
+ "fp16": "model_fp16",
27
+ "int8": "model_int8",
28
+ "uint8": "model_uint8",
29
+ "q4": "model_q4",
30
+ "q4f16": "model_q4f16",
31
+ "bnb4": "model_bnb4",
32
+ }
33
+
34
+ #: What `--precisions` expands to when the user does not say. fp32 is the honest
35
+ #: baseline every other number is relative to; int8 and q4 are the two choices a
36
+ #: person actually weighs when shipping to a CPU.
37
+ DEFAULT_PRECISIONS: tuple[str, ...] = ("fp32", "int8", "q4")
38
+
39
+ #: Small, widely-mirrored instruct models that carry the full precision set.
40
+ #: The default is the 135M: a complete three-precision sweep stays under a GB of
41
+ #: download and finishes in minutes on a laptop CPU with no accelerator.
42
+ DEFAULT_MODEL = "HuggingFaceTB/SmolLM2-135M-Instruct"
43
+
44
+
45
+ class ModelResolutionError(RuntimeError):
46
+ """A model/precision combination is not available on the Hub."""
47
+
48
+
49
+ @dataclass(frozen=True)
50
+ class ResolvedModel:
51
+ """A downloaded ONNX artifact, ready to hand to ONNX Runtime."""
52
+
53
+ model_id: str
54
+ precision: str
55
+ onnx_path: Path
56
+ size_mb: float
57
+ tokenizer_path: Path
58
+
59
+ @property
60
+ def filename(self) -> str:
61
+ return PRECISION_FILES[self.precision] + ".onnx"
62
+
63
+
64
+ def available_precisions(model_id: str, *, revision: str = "main") -> list[str]:
65
+ """Return the precision labels this repo actually publishes, in canonical order."""
66
+ from huggingface_hub import HfApi
67
+
68
+ files = set(HfApi().list_repo_files(model_id, revision=revision))
69
+ return [p for p, stem in PRECISION_FILES.items() if f"onnx/{stem}.onnx" in files]
70
+
71
+
72
+ def resolve(
73
+ model_id: str,
74
+ precision: str,
75
+ *,
76
+ revision: str = "main",
77
+ cache_dir: str | None = None,
78
+ ) -> ResolvedModel:
79
+ """Download (or reuse from cache) one precision of ``model_id``.
80
+
81
+ Raises :class:`ModelResolutionError` with the list of precisions the repo does
82
+ publish, rather than letting a 404 surface, because "this model has no int4"
83
+ is an ordinary situation a contributor needs to act on.
84
+ """
85
+ from huggingface_hub import hf_hub_download
86
+ from huggingface_hub.errors import EntryNotFoundError
87
+
88
+ if precision not in PRECISION_FILES:
89
+ raise ModelResolutionError(
90
+ f"Unknown precision '{precision}'. Known: {', '.join(PRECISION_FILES)}"
91
+ )
92
+
93
+ stem = PRECISION_FILES[precision]
94
+ kwargs = {"repo_id": model_id, "revision": revision, "cache_dir": cache_dir}
95
+
96
+ try:
97
+ onnx_path = Path(hf_hub_download(filename=f"onnx/{stem}.onnx", **kwargs))
98
+ except EntryNotFoundError as exc:
99
+ have = available_precisions(model_id, revision=revision)
100
+ raise ModelResolutionError(
101
+ f"{model_id} does not publish '{precision}' "
102
+ f"(looked for onnx/{stem}.onnx). Available: {', '.join(have) or 'none'}"
103
+ ) from exc
104
+
105
+ total_bytes = onnx_path.stat().st_size
106
+
107
+ # Models over the 2 GB protobuf limit keep their weights in a sibling
108
+ # ``.onnx_data`` file. ONNX Runtime needs it next to the graph, and it is most
109
+ # of the real on-disk footprint, so it counts toward the reported size.
110
+ try:
111
+ data_path = Path(hf_hub_download(filename=f"onnx/{stem}.onnx_data", **kwargs))
112
+ total_bytes += data_path.stat().st_size
113
+ log.debug("external weights: %s", data_path)
114
+ except EntryNotFoundError:
115
+ pass
116
+
117
+ tokenizer_path = Path(hf_hub_download(filename="tokenizer.json", **kwargs))
118
+
119
+ return ResolvedModel(
120
+ model_id=model_id,
121
+ precision=precision,
122
+ onnx_path=onnx_path,
123
+ size_mb=round(total_bytes / (1024 * 1024), 2),
124
+ tokenizer_path=tokenizer_path,
125
+ )
edgellm/leaderboard.py ADDED
@@ -0,0 +1,210 @@
1
+ """Aggregate every submitted card into a leaderboard.
2
+
3
+ The interesting question is not "which machine is fastest" — that is just a list
4
+ of who owns the newest CPU. It is **"does quantization actually pay off, and where"**.
5
+ So the leaderboard is organised around the per-precision speedup distribution
6
+ across machines, with the raw per-machine rows underneath.
7
+
8
+ Cards flagged as not comparable by :mod:`edgellm.validate` are listed but kept out
9
+ of the aggregate statistics, so a run with a different token budget cannot quietly
10
+ move the headline numbers.
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ import json
16
+ import statistics
17
+ from pathlib import Path
18
+ from typing import Any
19
+
20
+ from edgellm.validate import validate_card
21
+
22
+ BASELINE = "fp32"
23
+
24
+
25
+ def _plural(count: int, noun: str) -> str:
26
+ return f"{count} {noun}" if count == 1 else f"{count} {noun}s"
27
+
28
+
29
+ def _speedups(rows: list[dict]) -> dict[str, float]:
30
+ by_precision = {r["precision"]: r for r in rows}
31
+ base = by_precision.get(BASELINE)
32
+ if not base or not base.get("tokens_per_second"):
33
+ return {}
34
+ return {
35
+ r["precision"]: r["tokens_per_second"] / base["tokens_per_second"]
36
+ for r in rows
37
+ if r["precision"] != BASELINE and r.get("tokens_per_second")
38
+ }
39
+
40
+
41
+ def _ppl_deltas(rows: list[dict]) -> dict[str, float]:
42
+ by_precision = {r["precision"]: r for r in rows}
43
+ base = by_precision.get(BASELINE)
44
+ if not base or not base.get("perplexity"):
45
+ return {}
46
+ return {
47
+ r["precision"]: (r["perplexity"] / base["perplexity"] - 1) * 100
48
+ for r in rows
49
+ if r["precision"] != BASELINE and r.get("perplexity")
50
+ }
51
+
52
+
53
+ def build_leaderboard(card_paths: list[Path]) -> dict[str, Any]:
54
+ """Load, validate and aggregate cards into a JSON-serializable leaderboard."""
55
+ entries: list[dict[str, Any]] = []
56
+ rejected: list[dict[str, str]] = []
57
+
58
+ for path in sorted(card_paths):
59
+ report = validate_card(path)
60
+ if not report.ok:
61
+ rejected.append({"file": path.name, "reason": "; ".join(report.errors)})
62
+ continue
63
+
64
+ card = json.loads(path.read_text())
65
+ machine = card["machine"]
66
+ entries.append(
67
+ {
68
+ "file": path.name,
69
+ "comparable": report.comparable,
70
+ "model_id": card["model_id"],
71
+ "created_utc": card["created_utc"],
72
+ "submitted_by": card.get("submitted_by", ""),
73
+ "notes": card.get("notes", ""),
74
+ "cpu": machine["cpu"],
75
+ "arch": machine["arch"],
76
+ "os": machine["os"],
77
+ "physical_cores": machine["physical_cores"],
78
+ "ram_gb": machine["ram_gb"],
79
+ "provider": machine["provider"],
80
+ "threads": machine["intra_op_threads"],
81
+ "thread_policy": machine.get("thread_policy", "unknown"),
82
+ "onnxruntime": machine["onnxruntime"],
83
+ "rows": card["rows"],
84
+ "speedups": _speedups(card["rows"]),
85
+ "ppl_deltas": _ppl_deltas(card["rows"]),
86
+ }
87
+ )
88
+
89
+ aggregate = _aggregate(e for e in entries if e["comparable"])
90
+ return {
91
+ "entries": sorted(entries, key=lambda e: (e["model_id"], e["cpu"])),
92
+ "aggregate": aggregate,
93
+ "rejected": rejected,
94
+ "machine_count": len({e["cpu"] for e in entries}),
95
+ "model_count": len({e["model_id"] for e in entries}),
96
+ }
97
+
98
+
99
+ def _aggregate(entries) -> dict[str, Any]:
100
+ """Per-(model, precision) speedup and quality-cost distribution across machines."""
101
+ buckets: dict[tuple[str, str], dict[str, list[float]]] = {}
102
+ for entry in entries:
103
+ for precision, speedup in entry["speedups"].items():
104
+ bucket = buckets.setdefault((entry["model_id"], precision), {"speedup": [], "ppl": []})
105
+ bucket["speedup"].append(speedup)
106
+ for precision, delta in entry["ppl_deltas"].items():
107
+ bucket = buckets.setdefault((entry["model_id"], precision), {"speedup": [], "ppl": []})
108
+ bucket["ppl"].append(delta)
109
+
110
+ out = []
111
+ for (model_id, precision), values in sorted(buckets.items()):
112
+ speedups = values["speedup"]
113
+ ppls = values["ppl"]
114
+ out.append(
115
+ {
116
+ "model_id": model_id,
117
+ "precision": precision,
118
+ "machines": len(speedups),
119
+ "speedup_median": round(statistics.median(speedups), 3) if speedups else None,
120
+ "speedup_min": round(min(speedups), 3) if speedups else None,
121
+ "speedup_max": round(max(speedups), 3) if speedups else None,
122
+ "slower_than_baseline": sum(1 for s in speedups if s < 0.95),
123
+ "ppl_delta_median_pct": round(statistics.median(ppls), 2) if ppls else None,
124
+ }
125
+ )
126
+ return {"by_model_precision": out}
127
+
128
+
129
+ def render_leaderboard_markdown(board: dict[str, Any]) -> str:
130
+ """The committed LEADERBOARD.md."""
131
+ lines = [
132
+ "# Leaderboard",
133
+ "",
134
+ "<!-- Generated by `quantcost leaderboard`. Do not edit by hand:",
135
+ " every number here is derived from a card in results/community/. -->",
136
+ "",
137
+ f"**{_plural(board['machine_count'], 'machine')} · "
138
+ f"{_plural(board['model_count'], 'model')} · "
139
+ f"{_plural(len(board['entries']), 'card')} submitted**",
140
+ "",
141
+ "## Does quantization pay off?",
142
+ "",
143
+ "Speedup is throughput relative to that machine's own fp32 baseline, so the",
144
+ "column compares quantization choices rather than hardware budgets.",
145
+ "",
146
+ ]
147
+
148
+ aggregate = board["aggregate"]["by_model_precision"]
149
+ if aggregate:
150
+ lines += [
151
+ "| Model | Precision | Machines | Median speedup | Range | "
152
+ "Slower than fp32 | Median PPL change |",
153
+ "| --- | --- | --- | --- | --- | --- | --- |",
154
+ ]
155
+ for row in aggregate:
156
+ span = (
157
+ f"{row['speedup_min']:.2f}x – {row['speedup_max']:.2f}x"
158
+ if row["speedup_min"] is not None
159
+ else "—"
160
+ )
161
+ median = f"{row['speedup_median']:.2f}x" if row["speedup_median"] else "—"
162
+ ppl = (
163
+ f"{row['ppl_delta_median_pct']:+.1f}%"
164
+ if row["ppl_delta_median_pct"] is not None
165
+ else "—"
166
+ )
167
+ slower = f"{row['slower_than_baseline']}/{row['machines']}" if row["machines"] else "—"
168
+ lines.append(
169
+ f"| `{row['model_id']}` | {row['precision']} | {row['machines']} | "
170
+ f"{median} | {span} | {slower} | {ppl} |"
171
+ )
172
+ else:
173
+ lines.append("_No comparable cards yet._")
174
+
175
+ lines += ["", "## Per-machine results", ""]
176
+ for entry in board["entries"]:
177
+ flag = "" if entry["comparable"] else " _(not ranked — non-default settings)_"
178
+ credit = f" — submitted by {entry['submitted_by']}" if entry["submitted_by"] else ""
179
+ lines += [
180
+ f"### {entry['cpu']} · {entry['os']} ({entry['arch']}){flag}",
181
+ "",
182
+ f"`{entry['model_id']}` · {entry['physical_cores']} physical cores · "
183
+ f"{entry['ram_gb']:.0f} GB RAM · "
184
+ f"{entry['threads']} threads ({entry['thread_policy']}) · "
185
+ f"onnxruntime {entry['onnxruntime']} · {entry['provider']}{credit}",
186
+ "",
187
+ "| Precision | Size (MB) | tok/s | Speedup | Peak RAM (MB) | Perplexity |",
188
+ "| --- | --- | --- | --- | --- | --- |",
189
+ ]
190
+ for row in entry["rows"]:
191
+ speed = entry["speedups"].get(row["precision"])
192
+ speed_text = (
193
+ "baseline" if row["precision"] == BASELINE else (f"{speed:.2f}x" if speed else "—")
194
+ )
195
+ ppl = f"{row['perplexity']:.2f}" if row.get("perplexity") else "—"
196
+ lines.append(
197
+ f"| {row['precision']} | {row['size_mb']:.0f} | {row['tokens_per_second']:.2f} | "
198
+ f"{speed_text} | {row['peak_ram_mb']:.0f} | {ppl} |"
199
+ )
200
+ if entry["notes"]:
201
+ lines += ["", f"> {entry['notes']}"]
202
+ lines.append("")
203
+
204
+ if board["rejected"]:
205
+ lines += ["## Rejected cards", ""]
206
+ for item in board["rejected"]:
207
+ lines.append(f"- `{item['file']}`: {item['reason']}")
208
+ lines.append("")
209
+
210
+ return "\n".join(lines)
edgellm/models.py ADDED
@@ -0,0 +1,91 @@
1
+ """Model + tokenizer loading.
2
+
3
+ :class:`ModelLoader` wraps Hugging Face ``transformers`` so the rest of the code
4
+ never touches ``AutoModelForCausalLM`` directly and device/dtype resolution lives
5
+ in exactly one place.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import logging
11
+ from dataclasses import dataclass
12
+
13
+ import torch
14
+ from transformers import (
15
+ AutoModelForCausalLM,
16
+ AutoTokenizer,
17
+ PreTrainedModel,
18
+ PreTrainedTokenizerBase,
19
+ )
20
+
21
+ from edgellm.config import ModelConfig
22
+
23
+ logger = logging.getLogger(__name__)
24
+
25
+ _DTYPES: dict[str, torch.dtype] = {
26
+ "float32": torch.float32,
27
+ "float16": torch.float16,
28
+ "bfloat16": torch.bfloat16,
29
+ }
30
+
31
+
32
+ @dataclass
33
+ class LoadedModel:
34
+ """A loaded model bundled with its tokenizer and resolved placement."""
35
+
36
+ model: PreTrainedModel
37
+ tokenizer: PreTrainedTokenizerBase
38
+ device: str
39
+ dtype: torch.dtype
40
+
41
+
42
+ class ModelLoader:
43
+ """Load a causal-LM + tokenizer from a :class:`ModelConfig`."""
44
+
45
+ def __init__(self, config: ModelConfig) -> None:
46
+ self.config = config
47
+
48
+ def resolve_device(self) -> str:
49
+ """Turn ``device: auto`` into a concrete device string for this machine."""
50
+ requested = self.config.device
51
+ if requested != "auto":
52
+ return requested
53
+ if torch.cuda.is_available():
54
+ return "cuda"
55
+ if torch.backends.mps.is_available():
56
+ return "mps"
57
+ return "cpu"
58
+
59
+ def resolve_dtype(self) -> torch.dtype:
60
+ try:
61
+ return _DTYPES[self.config.dtype]
62
+ except KeyError as exc:
63
+ raise ValueError(
64
+ f"Unsupported dtype '{self.config.dtype}'. Choose one of {list(_DTYPES)}."
65
+ ) from exc
66
+
67
+ def load(self) -> LoadedModel:
68
+ """Load the model and tokenizer and move them onto the resolved device."""
69
+ device = self.resolve_device()
70
+ dtype = self.resolve_dtype()
71
+ logger.info("Loading %s (dtype=%s, device=%s)", self.config.id, self.config.dtype, device)
72
+
73
+ tokenizer = AutoTokenizer.from_pretrained(
74
+ self.config.id,
75
+ revision=self.config.revision,
76
+ trust_remote_code=self.config.trust_remote_code,
77
+ )
78
+ if tokenizer.pad_token is None and tokenizer.eos_token is not None:
79
+ # Small decoder-only models often ship without a pad token.
80
+ tokenizer.pad_token = tokenizer.eos_token
81
+
82
+ model = AutoModelForCausalLM.from_pretrained(
83
+ self.config.id,
84
+ revision=self.config.revision,
85
+ dtype=dtype,
86
+ trust_remote_code=self.config.trust_remote_code,
87
+ )
88
+ model.to(device)
89
+ model.eval()
90
+
91
+ return LoadedModel(model=model, tokenizer=tokenizer, device=device, dtype=dtype)
edgellm/ort_lite.py ADDED
@@ -0,0 +1,236 @@
1
+ """A torch-free ONNX Runtime decoder: greedy generation with a KV cache, in NumPy.
2
+
3
+ Why this exists instead of ``optimum.onnxruntime``: Optimum is the right tool when
4
+ you are exporting and quantizing models, but it pulls in PyTorch and
5
+ ``transformers``, which together are a multi-gigabyte install and the single
6
+ biggest reason a stranger abandons a benchmark run. The decode loop for a
7
+ causal LM is about a hundred lines of NumPy, so the default path implements it
8
+ directly and depends only on ``onnxruntime``, ``numpy`` and ``tokenizers``.
9
+
10
+ The graph contract (verified against the ``onnx-community`` / ``transformers.js``
11
+ exports, which is what the Hub publishes for thousands of models):
12
+
13
+ inputs input_ids int64 [batch, sequence]
14
+ attention_mask int64 [batch, total_sequence]
15
+ position_ids int64 [batch, sequence]
16
+ past_key_values.{i}.key T [batch, kv_heads, past_sequence, head_dim]
17
+ past_key_values.{i}.value T ... same
18
+ outputs logits float [batch, sequence, vocab]
19
+ present.{i}.key T [batch, kv_heads, total_sequence, head_dim]
20
+ present.{i}.value T ... same
21
+
22
+ Layer count, KV-head count, head dim and the cache dtype are all read off the
23
+ session rather than from ``config.json``, so the runner is model-agnostic and
24
+ works for fp16 graphs (whose cache is float16) without special-casing.
25
+ """
26
+
27
+ from __future__ import annotations
28
+
29
+ import logging
30
+ import re
31
+ import time
32
+
33
+ import numpy as np
34
+
35
+ from edgellm.base import GenerationResult, InferenceRunner
36
+ from edgellm.config import GenerationConfig
37
+
38
+ log = logging.getLogger(__name__)
39
+
40
+ _ORT_TO_NUMPY = {
41
+ "tensor(float)": np.float32,
42
+ "tensor(float16)": np.float16,
43
+ "tensor(bfloat16)": np.float32, # ORT has no bf16 numpy view; feed fp32 zeros.
44
+ "tensor(int64)": np.int64,
45
+ "tensor(int32)": np.int32,
46
+ }
47
+
48
+ _PAST_RE = re.compile(r"^past_key_values\.(\d+)\.(key|value)$")
49
+
50
+
51
+ class OnnxLiteRunner(InferenceRunner):
52
+ """Greedy ONNX Runtime generation with a KV cache, implemented in NumPy."""
53
+
54
+ def __init__(
55
+ self,
56
+ onnx_path: str,
57
+ tokenizer_path: str,
58
+ *,
59
+ provider: str = "CPUExecutionProvider",
60
+ name: str = "ort-lite",
61
+ intra_op_threads: int | None = None,
62
+ ) -> None:
63
+ import onnxruntime as ort
64
+ from tokenizers import Tokenizer
65
+
66
+ available = ort.get_available_providers()
67
+ if provider not in available:
68
+ raise RuntimeError(
69
+ f"Execution provider '{provider}' is not available in this ONNX "
70
+ f"Runtime build (have: {', '.join(available)})."
71
+ )
72
+
73
+ # Thread count is the single largest lever on CPU throughput, so it is
74
+ # always pinned to a concrete number and recorded in the result card.
75
+ # Leaving it at ORT's default would store 0 ("decide for me"), which
76
+ # makes two cards look identical while having run very differently.
77
+ # The *fast* cores only, not every physical core. On a heterogeneous CPU
78
+ # one thread landing on an efficiency core gates the whole parallel
79
+ # region; see edgellm.card.performance_cores for the measurements.
80
+ if intra_op_threads is None:
81
+ from edgellm.card import performance_cores
82
+
83
+ intra_op_threads, self.thread_policy = performance_cores()
84
+ else:
85
+ self.thread_policy = "explicit"
86
+
87
+ opts = ort.SessionOptions()
88
+ opts.intra_op_num_threads = intra_op_threads
89
+ opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
90
+
91
+ self.name = name
92
+ self.provider = provider
93
+ self.session = ort.InferenceSession(onnx_path, opts, providers=[provider])
94
+ self.tokenizer = Tokenizer.from_file(tokenizer_path)
95
+ self.intra_op_threads = opts.intra_op_num_threads
96
+
97
+ self._inspect_graph()
98
+
99
+ # ------------------------------------------------------------------ setup
100
+
101
+ def _inspect_graph(self) -> None:
102
+ """Read layer count, KV geometry and cache dtype off the graph itself."""
103
+ inputs = {i.name: i for i in self.session.get_inputs()}
104
+ self.input_names = set(inputs)
105
+
106
+ layers = sorted(
107
+ int(m.group(1))
108
+ for name in inputs
109
+ if (m := _PAST_RE.match(name)) and m.group(2) == "key"
110
+ )
111
+ if not layers:
112
+ raise ValueError(
113
+ "This ONNX graph exposes no 'past_key_values.*' inputs, so it is not a "
114
+ "KV-cached decoder export. Use a model from onnx-community or any repo "
115
+ "exported with Optimum's `--task text-generation-with-past`."
116
+ )
117
+ if layers != list(range(len(layers))):
118
+ raise ValueError(f"Non-contiguous KV cache layer indices: {layers}")
119
+ self.num_layers = len(layers)
120
+
121
+ spec = inputs["past_key_values.0.key"]
122
+ shape = spec.shape
123
+ if len(shape) != 4 or not isinstance(shape[1], int) or not isinstance(shape[3], int):
124
+ raise ValueError(
125
+ f"Unexpected KV cache shape {shape}; expected "
126
+ "[batch, kv_heads, past_sequence, head_dim] with static head dims."
127
+ )
128
+ self.kv_heads, self.head_dim = int(shape[1]), int(shape[3])
129
+ self.cache_dtype = _ORT_TO_NUMPY.get(spec.type, np.float32)
130
+
131
+ self.vocab_size = None
132
+ for out in self.session.get_outputs():
133
+ if out.name == "logits" and isinstance(out.shape[-1], int):
134
+ self.vocab_size = int(out.shape[-1])
135
+
136
+ self.uses_position_ids = "position_ids" in self.input_names
137
+
138
+ log.debug(
139
+ "graph: %d layers, %d kv heads, head_dim %d, cache %s, position_ids %s",
140
+ self.num_layers,
141
+ self.kv_heads,
142
+ self.head_dim,
143
+ np.dtype(self.cache_dtype).name,
144
+ self.uses_position_ids,
145
+ )
146
+
147
+ def _empty_cache(self, batch: int = 1) -> dict[str, np.ndarray]:
148
+ """Zero-length KV cache: the graph's signal that this is the prefill step."""
149
+ empty = np.zeros((batch, self.kv_heads, 0, self.head_dim), dtype=self.cache_dtype)
150
+ feed = {}
151
+ for i in range(self.num_layers):
152
+ feed[f"past_key_values.{i}.key"] = empty
153
+ feed[f"past_key_values.{i}.value"] = empty
154
+ return feed
155
+
156
+ # ------------------------------------------------------------- inference
157
+
158
+ def _forward(
159
+ self,
160
+ input_ids: np.ndarray,
161
+ past: dict[str, np.ndarray],
162
+ past_len: int,
163
+ ) -> tuple[np.ndarray, dict[str, np.ndarray]]:
164
+ """One step (prefill or decode). Returns logits and the next cache."""
165
+ batch, seq = input_ids.shape
166
+ total = past_len + seq
167
+
168
+ feed: dict[str, np.ndarray] = {
169
+ "input_ids": input_ids,
170
+ "attention_mask": np.ones((batch, total), dtype=np.int64),
171
+ **past,
172
+ }
173
+ if self.uses_position_ids:
174
+ feed["position_ids"] = np.arange(past_len, total, dtype=np.int64)[None, :].repeat(
175
+ batch, axis=0
176
+ )
177
+
178
+ outputs = self.session.run(None, feed)
179
+ names = [o.name for o in self.session.get_outputs()]
180
+ by_name = dict(zip(names, outputs, strict=True))
181
+
182
+ logits = by_name["logits"]
183
+ next_past = {
184
+ f"past_key_values.{i}.{kind}": by_name[f"present.{i}.{kind}"]
185
+ for i in range(self.num_layers)
186
+ for kind in ("key", "value")
187
+ }
188
+ return logits, next_past
189
+
190
+ def forward_logits(self, input_ids: np.ndarray) -> np.ndarray:
191
+ """Single full forward pass with no cache — what perplexity scoring needs."""
192
+ logits, _ = self._forward(input_ids, self._empty_cache(input_ids.shape[0]), 0)
193
+ return logits
194
+
195
+ def generate(self, prompt: str, generation: GenerationConfig) -> GenerationResult:
196
+ """Greedy decode ``max_new_tokens`` tokens, timing the whole loop.
197
+
198
+ Decoding is always greedy here regardless of ``generation.temperature``:
199
+ a benchmark needs a fixed token count and a deterministic path, and
200
+ sampling would make throughput depend on the RNG. ``min_new_tokens`` is
201
+ implied — generation never stops early on EOS, so every backend and
202
+ precision is measured over exactly the same amount of work.
203
+ """
204
+ ids = self.tokenizer.encode(prompt).ids
205
+ input_ids = np.asarray([ids], dtype=np.int64)
206
+ prompt_tokens = len(ids)
207
+ want = generation.max_new_tokens
208
+
209
+ start = time.perf_counter()
210
+
211
+ logits, past = self._forward(input_ids, self._empty_cache(), 0)
212
+ past_len = prompt_tokens
213
+ next_id = int(np.argmax(logits[0, -1]))
214
+ generated = [next_id]
215
+
216
+ for _ in range(want - 1):
217
+ logits, past = self._forward(np.asarray([[next_id]], dtype=np.int64), past, past_len)
218
+ past_len += 1
219
+ next_id = int(np.argmax(logits[0, -1]))
220
+ generated.append(next_id)
221
+
222
+ latency_s = time.perf_counter() - start
223
+
224
+ text = self.tokenizer.decode(generated, skip_special_tokens=True)
225
+ return GenerationResult(
226
+ backend=self.name,
227
+ prompt=prompt,
228
+ text=text,
229
+ prompt_tokens=prompt_tokens,
230
+ generated_tokens=len(generated),
231
+ latency_s=latency_s,
232
+ tokens_per_second=len(generated) / latency_s if latency_s > 0 else 0.0,
233
+ )
234
+
235
+ def close(self) -> None:
236
+ self.session = None