quantdiff 0.1.0rc1__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.
- quantdiff/__init__.py +53 -0
- quantdiff/__main__.py +5 -0
- quantdiff/_http.py +151 -0
- quantdiff/_text.py +13 -0
- quantdiff/_version.py +1 -0
- quantdiff/api.py +340 -0
- quantdiff/backends/__init__.py +28 -0
- quantdiff/backends/_common.py +342 -0
- quantdiff/backends/base.py +91 -0
- quantdiff/backends/llamacpp.py +428 -0
- quantdiff/backends/ollama.py +359 -0
- quantdiff/backends/openai_compat.py +338 -0
- quantdiff/cache.py +240 -0
- quantdiff/card.py +1664 -0
- quantdiff/cli.py +377 -0
- quantdiff/discover.py +488 -0
- quantdiff/errors.py +45 -0
- quantdiff/metrics/__init__.py +36 -0
- quantdiff/metrics/codeexec.py +428 -0
- quantdiff/metrics/jsonschema.py +610 -0
- quantdiff/metrics/logit.py +214 -0
- quantdiff/metrics/tasks.py +114 -0
- quantdiff/metrics/textsim.py +66 -0
- quantdiff/metrics/toolcheck.py +99 -0
- quantdiff/png.py +360 -0
- quantdiff/preflight.py +365 -0
- quantdiff/progress.py +283 -0
- quantdiff/py.typed +0 -0
- quantdiff/report.py +780 -0
- quantdiff/runner.py +492 -0
- quantdiff/spec.py +154 -0
- quantdiff/stats.py +226 -0
- quantdiff/suites/__init__.py +462 -0
- quantdiff/suites/data/chat.jsonl +22 -0
- quantdiff/suites/data/code.jsonl +32 -0
- quantdiff/suites/data/json.jsonl +34 -0
- quantdiff/suites/data/scoring.jsonl +41 -0
- quantdiff/suites/data/tools.jsonl +32 -0
- quantdiff/types.py +322 -0
- quantdiff/verdict.py +1513 -0
- quantdiff-0.1.0rc1.dist-info/METADATA +514 -0
- quantdiff-0.1.0rc1.dist-info/RECORD +45 -0
- quantdiff-0.1.0rc1.dist-info/WHEEL +4 -0
- quantdiff-0.1.0rc1.dist-info/entry_points.txt +2 -0
- quantdiff-0.1.0rc1.dist-info/licenses/LICENSE +202 -0
quantdiff/preflight.py
ADDED
|
@@ -0,0 +1,365 @@
|
|
|
1
|
+
"""Pre-flight checks that catch setup mistakes before a long comparison run.
|
|
2
|
+
|
|
3
|
+
Each check is isolated: a server error inside one check becomes a "skip" finding for
|
|
4
|
+
that check instead of aborting the others. Checks that do not apply to the given
|
|
5
|
+
options (no reference, no Hugging Face repo, probe disabled) produce no finding.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import hashlib
|
|
11
|
+
import logging
|
|
12
|
+
import re
|
|
13
|
+
from collections.abc import Callable, Mapping, Sequence
|
|
14
|
+
from typing import Final
|
|
15
|
+
|
|
16
|
+
from quantdiff._http import get_json
|
|
17
|
+
from quantdiff.backends.base import Backend
|
|
18
|
+
from quantdiff.errors import BackendError, SpecError
|
|
19
|
+
from quantdiff.types import BackendKind, Message, PreflightFinding, ServerInfo
|
|
20
|
+
|
|
21
|
+
__all__ = [
|
|
22
|
+
"DEFAULT_CONTEXT_PROBE_TOKENS",
|
|
23
|
+
"HF_BASE_URL",
|
|
24
|
+
"MIN_CONTEXT_PROBE_TOKENS",
|
|
25
|
+
"run_preflight",
|
|
26
|
+
]
|
|
27
|
+
|
|
28
|
+
logger = logging.getLogger(__name__)
|
|
29
|
+
|
|
30
|
+
HF_BASE_URL: Final = "https://huggingface.co"
|
|
31
|
+
DEFAULT_CONTEXT_PROBE_TOKENS: Final = 6000
|
|
32
|
+
MIN_CONTEXT_PROBE_TOKENS: Final = 1000
|
|
33
|
+
"""Smallest long-probe size worth running; 0 disables the probe instead."""
|
|
34
|
+
|
|
35
|
+
_CONTROL_PROBE_TOKENS: Final = 200
|
|
36
|
+
_WORDS_PER_TOKEN: Final = 0.75
|
|
37
|
+
_PROBE_ANSWER_TOKENS: Final = 16
|
|
38
|
+
_HF_TIMEOUT_SECONDS: Final = 15.0
|
|
39
|
+
_HF_MAX_BYTES: Final = 2 * 1024 * 1024
|
|
40
|
+
_HF_REPO_PATTERN: Final = re.compile(r"[A-Za-z0-9][A-Za-z0-9._-]*/[A-Za-z0-9][A-Za-z0-9._-]*")
|
|
41
|
+
|
|
42
|
+
_CONTEXT_FIXES: Final[Mapping[BackendKind, str]] = {
|
|
43
|
+
"ollama": (
|
|
44
|
+
"Set OLLAMA_CONTEXT_LENGTH={size} on the Ollama server, or add "
|
|
45
|
+
"PARAMETER num_ctx {size} to a Modelfile."
|
|
46
|
+
),
|
|
47
|
+
"llamacpp": "Restart llama-server with -c {size} (--ctx-size).",
|
|
48
|
+
"openai": (
|
|
49
|
+
"Raise the context length to at least {size} tokens in the server settings "
|
|
50
|
+
"(LM Studio: Context Length; vLLM: --max-model-len)."
|
|
51
|
+
),
|
|
52
|
+
}
|
|
53
|
+
_TEMPLATE_FIX: Final = (
|
|
54
|
+
"Re-download the model, or pass --chat-template-file with the upstream template."
|
|
55
|
+
)
|
|
56
|
+
_LOGPROBS_FIX: Final = "Serve the model with llama-server or another backend that returns logprobs."
|
|
57
|
+
_TOKENIZER_FIX: Final = "Compare quantizations of the same base model as the reference."
|
|
58
|
+
|
|
59
|
+
_TOKENIZER_PROBES: Final = (
|
|
60
|
+
"The quick brown fox jumps over the lazy dog.",
|
|
61
|
+
" leading spaces,\ttabs\nand newlines\n\n",
|
|
62
|
+
"def add(a: int, b: int) -> int:\n return a + b # sum\n",
|
|
63
|
+
'{"id": 42, "tags": ["x", "y"], "ok": true}',
|
|
64
|
+
"Grüße aus München, café crème brûlée",
|
|
65
|
+
"你好世界。量化模型可以在普通电脑上运行",
|
|
66
|
+
"नमस्ते दुनिया",
|
|
67
|
+
"emoji \U0001f642\U0001f680 and numbers 3.14159 -1024 1e-9",
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
_SYLLABLES: Final = (
|
|
71
|
+
"ka", "lo", "mi", "ru", "te", "va", "zo", "ne",
|
|
72
|
+
"pi", "su", "do", "fa", "gi", "ho", "ju", "be",
|
|
73
|
+
) # fmt: skip
|
|
74
|
+
_SUBJECTS: Final = (
|
|
75
|
+
"The harbor master",
|
|
76
|
+
"A patient gardener",
|
|
77
|
+
"The night librarian",
|
|
78
|
+
"An old cartographer",
|
|
79
|
+
"The village baker",
|
|
80
|
+
"A young engineer",
|
|
81
|
+
"The ferry captain",
|
|
82
|
+
"A careful auditor",
|
|
83
|
+
"The museum guide",
|
|
84
|
+
"A retired teacher",
|
|
85
|
+
"The orchard keeper",
|
|
86
|
+
"A traveling violinist",
|
|
87
|
+
"The station clerk",
|
|
88
|
+
)
|
|
89
|
+
_VERBS: Final = (
|
|
90
|
+
"quietly repaired",
|
|
91
|
+
"carefully counted",
|
|
92
|
+
"slowly painted",
|
|
93
|
+
"proudly displayed",
|
|
94
|
+
"briefly studied",
|
|
95
|
+
"neatly folded",
|
|
96
|
+
"gently polished",
|
|
97
|
+
"patiently sorted",
|
|
98
|
+
"happily described",
|
|
99
|
+
"firmly secured",
|
|
100
|
+
"openly admired",
|
|
101
|
+
)
|
|
102
|
+
_OBJECTS: Final = (
|
|
103
|
+
"a wooden bench",
|
|
104
|
+
"the brass lanterns",
|
|
105
|
+
"several faded maps",
|
|
106
|
+
"a basket of pears",
|
|
107
|
+
"the winter blankets",
|
|
108
|
+
"an iron gate",
|
|
109
|
+
"a stack of letters",
|
|
110
|
+
)
|
|
111
|
+
_PLACES: Final = (
|
|
112
|
+
"near the river before noon",
|
|
113
|
+
"beside the old market square",
|
|
114
|
+
"in the garden after the rain",
|
|
115
|
+
"at the edge of the quiet town",
|
|
116
|
+
"under the tall chestnut trees",
|
|
117
|
+
)
|
|
118
|
+
_SENTENCES_PER_PARAGRAPH: Final = 8
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
def run_preflight(
|
|
122
|
+
backend: Backend,
|
|
123
|
+
*,
|
|
124
|
+
reference: Backend | None = None,
|
|
125
|
+
hf_repo: str | None = None,
|
|
126
|
+
offline: bool = False,
|
|
127
|
+
context_probe_tokens: int = DEFAULT_CONTEXT_PROBE_TOKENS,
|
|
128
|
+
hf_base_url: str = HF_BASE_URL,
|
|
129
|
+
) -> tuple[PreflightFinding, ...]:
|
|
130
|
+
"""Run the pre-flight checks for `backend` and return their findings in order.
|
|
131
|
+
|
|
132
|
+
`context_probe_tokens` sets the size of the long-prompt truncation probe; 0 turns
|
|
133
|
+
the context check off. `hf_repo` (such as "Qwen/Qwen2.5-7B-Instruct") enables the
|
|
134
|
+
chat template comparison. `reference` enables the tokenizer comparison.
|
|
135
|
+
"""
|
|
136
|
+
if hf_repo is not None and not _HF_REPO_PATTERN.fullmatch(hf_repo):
|
|
137
|
+
raise SpecError(f"invalid Hugging Face repo {hf_repo!r}; expected owner/name")
|
|
138
|
+
if context_probe_tokens != 0 and context_probe_tokens < MIN_CONTEXT_PROBE_TOKENS:
|
|
139
|
+
raise SpecError(
|
|
140
|
+
f"context probe must be 0 (off) or at least {MIN_CONTEXT_PROBE_TOKENS} tokens"
|
|
141
|
+
)
|
|
142
|
+
|
|
143
|
+
findings = _isolated("logprobs", lambda: (_check_logprobs(backend),))
|
|
144
|
+
if context_probe_tokens:
|
|
145
|
+
findings += _isolated("context", lambda: _check_context(backend, context_probe_tokens))
|
|
146
|
+
if hf_repo is not None:
|
|
147
|
+
findings += _isolated(
|
|
148
|
+
"template",
|
|
149
|
+
lambda: (_check_template(backend, hf_repo, offline=offline, base_url=hf_base_url),),
|
|
150
|
+
)
|
|
151
|
+
if reference is not None:
|
|
152
|
+
findings += _isolated("tokenizer", lambda: (_check_tokenizer(backend, reference),))
|
|
153
|
+
for finding in findings:
|
|
154
|
+
logger.debug("preflight %s %s: %s", finding.check, finding.severity, finding.message)
|
|
155
|
+
return tuple(findings)
|
|
156
|
+
|
|
157
|
+
|
|
158
|
+
def _isolated(check: str, run: Callable[[], Sequence[PreflightFinding]]) -> list[PreflightFinding]:
|
|
159
|
+
try:
|
|
160
|
+
return list(run())
|
|
161
|
+
except BackendError as exc:
|
|
162
|
+
return [PreflightFinding(check, "skip", f"check could not run: {exc}")]
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
# logprobs --------------------------------------------------------------------------------
|
|
166
|
+
|
|
167
|
+
|
|
168
|
+
def _check_logprobs(backend: Backend) -> PreflightFinding:
|
|
169
|
+
if backend.info().supports_logprobs:
|
|
170
|
+
return PreflightFinding("logprobs", "ok", "logprobs available")
|
|
171
|
+
return PreflightFinding(
|
|
172
|
+
"logprobs", "warn", "logit metrics unavailable; task metrics only", _LOGPROBS_FIX
|
|
173
|
+
)
|
|
174
|
+
|
|
175
|
+
|
|
176
|
+
# context ---------------------------------------------------------------------------------
|
|
177
|
+
|
|
178
|
+
|
|
179
|
+
def _check_context(backend: Backend, probe_tokens: int) -> tuple[PreflightFinding, ...]:
|
|
180
|
+
info = backend.info()
|
|
181
|
+
fix = _CONTEXT_FIXES[info.backend].format(size=_suggested_context(probe_tokens))
|
|
182
|
+
findings = []
|
|
183
|
+
if info.context_length is not None and info.context_length < probe_tokens:
|
|
184
|
+
findings.append(
|
|
185
|
+
PreflightFinding(
|
|
186
|
+
"context",
|
|
187
|
+
"warn",
|
|
188
|
+
f"context window is {info.context_length} tokens, below the "
|
|
189
|
+
f"{probe_tokens}-token probe; long prompts will not fit",
|
|
190
|
+
fix,
|
|
191
|
+
)
|
|
192
|
+
)
|
|
193
|
+
findings.append(_needle_probe(backend, probe_tokens, fix))
|
|
194
|
+
return tuple(findings)
|
|
195
|
+
|
|
196
|
+
|
|
197
|
+
def _needle_probe(backend: Backend, probe_tokens: int, fix: str) -> PreflightFinding:
|
|
198
|
+
"""Check that text at the very start of a long prompt still reaches the model.
|
|
199
|
+
|
|
200
|
+
Servers that silently drop the front of an over-long prompt return fluent answers
|
|
201
|
+
that are simply wrong, so the probe hides a code word at the start and asks for it
|
|
202
|
+
at the end. A short control probe first rules out models that cannot do the task.
|
|
203
|
+
"""
|
|
204
|
+
try:
|
|
205
|
+
control_passed = _recalls_code_word(backend, _CONTROL_PROBE_TOKENS, "control")
|
|
206
|
+
except BackendError as exc:
|
|
207
|
+
return PreflightFinding("context", "skip", f"short probe failed: {exc}")
|
|
208
|
+
if not control_passed:
|
|
209
|
+
return PreflightFinding(
|
|
210
|
+
"context",
|
|
211
|
+
"skip",
|
|
212
|
+
"model could not answer the probe at short length; context check inconclusive",
|
|
213
|
+
)
|
|
214
|
+
try:
|
|
215
|
+
long_passed = _recalls_code_word(backend, probe_tokens, "long")
|
|
216
|
+
except BackendError as exc:
|
|
217
|
+
return PreflightFinding("context", "warn", f"{probe_tokens}-token probe failed: {exc}", fix)
|
|
218
|
+
if not long_passed:
|
|
219
|
+
return PreflightFinding(
|
|
220
|
+
"context", "fail", "front of long prompts is being dropped (silent truncation)", fix
|
|
221
|
+
)
|
|
222
|
+
return PreflightFinding(
|
|
223
|
+
"context", "ok", f"model recalled the start of a {probe_tokens}-token prompt"
|
|
224
|
+
)
|
|
225
|
+
|
|
226
|
+
|
|
227
|
+
def _recalls_code_word(backend: Backend, probe_tokens: int, salt: str) -> bool:
|
|
228
|
+
code_word = _code_word(salt)
|
|
229
|
+
prompt = (
|
|
230
|
+
f"Remember this code word: {code_word}.\n\n"
|
|
231
|
+
f"{_filler(int(probe_tokens * _WORDS_PER_TOKEN))}\n\n"
|
|
232
|
+
"What was the code word given at the very start of this message? "
|
|
233
|
+
"Reply with the code word only."
|
|
234
|
+
)
|
|
235
|
+
result = backend.chat([Message(role="user", content=prompt)], max_tokens=_PROBE_ANSWER_TOKENS)
|
|
236
|
+
return _alphanumeric(code_word) in _alphanumeric(result.text)
|
|
237
|
+
|
|
238
|
+
|
|
239
|
+
def _code_word(salt: str) -> str:
|
|
240
|
+
digest = hashlib.sha256(f"quantdiff-needle-{salt}".encode()).digest()
|
|
241
|
+
word = "".join(_SYLLABLES[byte % len(_SYLLABLES)] for byte in digest[:4])
|
|
242
|
+
number = int.from_bytes(digest[4:6], "big") % 9000 + 1000
|
|
243
|
+
return f"{word.capitalize()}-{number}"
|
|
244
|
+
|
|
245
|
+
|
|
246
|
+
def _filler(word_count: int) -> str:
|
|
247
|
+
"""Return about `word_count` words of varied, deterministic, meaningless prose."""
|
|
248
|
+
paragraphs: list[str] = []
|
|
249
|
+
sentences: list[str] = []
|
|
250
|
+
words = 0
|
|
251
|
+
index = 0
|
|
252
|
+
while words < word_count:
|
|
253
|
+
sentence = (
|
|
254
|
+
f"{_SUBJECTS[index % len(_SUBJECTS)]} {_VERBS[index % len(_VERBS)]} "
|
|
255
|
+
f"{_OBJECTS[index % len(_OBJECTS)]} {_PLACES[index % len(_PLACES)]}."
|
|
256
|
+
)
|
|
257
|
+
sentences.append(sentence)
|
|
258
|
+
words += len(sentence.split())
|
|
259
|
+
index += 1
|
|
260
|
+
if len(sentences) == _SENTENCES_PER_PARAGRAPH:
|
|
261
|
+
paragraphs.append(" ".join(sentences))
|
|
262
|
+
sentences = []
|
|
263
|
+
if sentences:
|
|
264
|
+
paragraphs.append(" ".join(sentences))
|
|
265
|
+
return "\n\n".join(paragraphs)
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
def _alphanumeric(text: str) -> str:
|
|
269
|
+
return re.sub(r"[^a-z0-9]", "", text.lower())
|
|
270
|
+
|
|
271
|
+
|
|
272
|
+
def _suggested_context(probe_tokens: int) -> int:
|
|
273
|
+
"""Smallest power of two that fits the probe plus room for an answer."""
|
|
274
|
+
size = 1024
|
|
275
|
+
while size < probe_tokens + 512:
|
|
276
|
+
size *= 2
|
|
277
|
+
return size
|
|
278
|
+
|
|
279
|
+
|
|
280
|
+
# chat template ---------------------------------------------------------------------------
|
|
281
|
+
|
|
282
|
+
|
|
283
|
+
def _check_template(
|
|
284
|
+
backend: Backend, repo: str, *, offline: bool, base_url: str
|
|
285
|
+
) -> PreflightFinding:
|
|
286
|
+
info = backend.info()
|
|
287
|
+
reason = _template_skip_reason(info, repo, offline=offline)
|
|
288
|
+
if reason is not None:
|
|
289
|
+
return PreflightFinding("template", "skip", reason)
|
|
290
|
+
upstream = _upstream_template(repo, base_url)
|
|
291
|
+
if upstream is None:
|
|
292
|
+
return PreflightFinding(
|
|
293
|
+
"template", "skip", f"upstream {repo} tokenizer_config.json has no chat_template"
|
|
294
|
+
)
|
|
295
|
+
if _squash(upstream) == _squash(info.chat_template or ""):
|
|
296
|
+
return PreflightFinding("template", "ok", f"embedded chat template matches upstream {repo}")
|
|
297
|
+
return PreflightFinding(
|
|
298
|
+
"template", "warn", f"embedded chat template differs from upstream {repo}", _TEMPLATE_FIX
|
|
299
|
+
)
|
|
300
|
+
|
|
301
|
+
|
|
302
|
+
def _template_skip_reason(info: ServerInfo, repo: str, *, offline: bool) -> str | None:
|
|
303
|
+
if info.template_dialect == "go":
|
|
304
|
+
return "Ollama templates are Go templates; compare not supported"
|
|
305
|
+
if info.template_dialect != "jinja":
|
|
306
|
+
return "server does not report a Jinja chat template"
|
|
307
|
+
if offline:
|
|
308
|
+
return f"offline; upstream template for {repo} not fetched"
|
|
309
|
+
if info.chat_template is None:
|
|
310
|
+
return "server did not report its chat template"
|
|
311
|
+
return None
|
|
312
|
+
|
|
313
|
+
|
|
314
|
+
def _upstream_template(repo: str, base_url: str) -> str | None:
|
|
315
|
+
url = f"{base_url.rstrip('/')}/{repo}/raw/main/tokenizer_config.json"
|
|
316
|
+
config = get_json(url, timeout=_HF_TIMEOUT_SECONDS, max_bytes=_HF_MAX_BYTES)
|
|
317
|
+
if not isinstance(config, dict):
|
|
318
|
+
raise BackendError(f"{url} did not return a JSON object")
|
|
319
|
+
template = config.get("chat_template")
|
|
320
|
+
if isinstance(template, str):
|
|
321
|
+
return template
|
|
322
|
+
if isinstance(template, list):
|
|
323
|
+
for entry in template:
|
|
324
|
+
if (
|
|
325
|
+
isinstance(entry, dict)
|
|
326
|
+
and entry.get("name") == "default"
|
|
327
|
+
and isinstance(entry.get("template"), str)
|
|
328
|
+
):
|
|
329
|
+
return str(entry["template"])
|
|
330
|
+
return None
|
|
331
|
+
|
|
332
|
+
|
|
333
|
+
def _squash(text: str) -> str:
|
|
334
|
+
return " ".join(text.split())
|
|
335
|
+
|
|
336
|
+
|
|
337
|
+
# tokenizer -------------------------------------------------------------------------------
|
|
338
|
+
|
|
339
|
+
|
|
340
|
+
def _check_tokenizer(backend: Backend, reference: Backend) -> PreflightFinding:
|
|
341
|
+
for probe in _TOKENIZER_PROBES:
|
|
342
|
+
candidate_ids = backend.tokenize(probe)
|
|
343
|
+
if candidate_ids is None:
|
|
344
|
+
return _tokenizer_skip(backend)
|
|
345
|
+
reference_ids = reference.tokenize(probe)
|
|
346
|
+
if reference_ids is None:
|
|
347
|
+
return _tokenizer_skip(reference)
|
|
348
|
+
if candidate_ids != reference_ids:
|
|
349
|
+
return PreflightFinding(
|
|
350
|
+
"tokenizer",
|
|
351
|
+
"fail",
|
|
352
|
+
"tokenizers differ from the reference; logit metrics are not comparable",
|
|
353
|
+
_TOKENIZER_FIX,
|
|
354
|
+
)
|
|
355
|
+
return PreflightFinding(
|
|
356
|
+
"tokenizer",
|
|
357
|
+
"ok",
|
|
358
|
+
f"tokenizer matches the reference on {len(_TOKENIZER_PROBES)} probe strings",
|
|
359
|
+
)
|
|
360
|
+
|
|
361
|
+
|
|
362
|
+
def _tokenizer_skip(backend: Backend) -> PreflightFinding:
|
|
363
|
+
return PreflightFinding(
|
|
364
|
+
"tokenizer", "skip", f"{backend.spec.label} cannot tokenize; tokenizers not compared"
|
|
365
|
+
)
|
quantdiff/progress.py
ADDED
|
@@ -0,0 +1,283 @@
|
|
|
1
|
+
"""Progress displays for long runs.
|
|
2
|
+
|
|
3
|
+
`make_progress` picks one of two displays for a stream:
|
|
4
|
+
|
|
5
|
+
- On a terminal, one status line is redrawn in place with a bar, percentage, ETA and the
|
|
6
|
+
current step. Each finished step leaves a permanent `done` line so the scrollback reads
|
|
7
|
+
as a short log. A step is finished only when the next one starts, so a step that fails
|
|
8
|
+
never gets a `done` line, and connecting to every server counts as one step whose
|
|
9
|
+
`done` line says how many servers answered.
|
|
10
|
+
- On a pipe or in CI, plain lines are printed when the step changes and at most once per
|
|
11
|
+
tenth of the run, so logs stay short and contain no carriage returns.
|
|
12
|
+
|
|
13
|
+
The ETA divides the units left by the rate measured in the current step. Work units are
|
|
14
|
+
weighted a priori and a scoring unit costs far less than a case, so a rate averaged over
|
|
15
|
+
the whole run is badly off; the recent rate of the step at hand is a much better guide.
|
|
16
|
+
|
|
17
|
+
Output is ASCII only so it renders on any Windows console code page.
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
from __future__ import annotations
|
|
21
|
+
|
|
22
|
+
import os
|
|
23
|
+
import shutil
|
|
24
|
+
import time
|
|
25
|
+
from abc import ABC, abstractmethod
|
|
26
|
+
from collections.abc import Callable
|
|
27
|
+
from typing import Final, TextIO
|
|
28
|
+
|
|
29
|
+
from quantdiff.types import ProgressEvent
|
|
30
|
+
|
|
31
|
+
Clock = Callable[[], float]
|
|
32
|
+
ProgressFn = Callable[[ProgressEvent], None]
|
|
33
|
+
|
|
34
|
+
BAR_WIDTH: Final = 20
|
|
35
|
+
REDRAW_INTERVAL_SECONDS: Final = 0.1
|
|
36
|
+
ETA_MIN_UNITS: Final = 5
|
|
37
|
+
ETA_MIN_SECONDS: Final = 10.0
|
|
38
|
+
"""Units and seconds into a step before its ETA is shown; earlier estimates swing wildly."""
|
|
39
|
+
RATE_WINDOW_EVENTS: Final = 10
|
|
40
|
+
"""The measured rate weights roughly the last this many events."""
|
|
41
|
+
LOG_STEPS: Final = 10
|
|
42
|
+
"""Non-interactive output prints at most once per 1/LOG_STEPS of the run."""
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def make_progress(
|
|
46
|
+
stream: TextIO, *, interactive: bool | None = None, clock: Clock = time.monotonic
|
|
47
|
+
) -> ProgressDisplay:
|
|
48
|
+
"""Return a progress callback writing to `stream`.
|
|
49
|
+
|
|
50
|
+
`interactive` None means a live status line when `stream` is a terminal and TERM is not
|
|
51
|
+
"dumb", and plain log lines otherwise. Call `close()` on the result before printing an
|
|
52
|
+
error so a half-drawn status line does not run into it.
|
|
53
|
+
"""
|
|
54
|
+
if interactive is None:
|
|
55
|
+
interactive = stream.isatty() and os.environ.get("TERM") != "dumb"
|
|
56
|
+
if interactive:
|
|
57
|
+
return StatusLine(stream, clock)
|
|
58
|
+
return StepLog(stream, clock)
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def format_duration(seconds: float) -> str:
|
|
62
|
+
"""Format as 45s, 3m05s or 1h02m."""
|
|
63
|
+
whole = max(0, round(seconds))
|
|
64
|
+
if whole < 60:
|
|
65
|
+
return f"{whole}s"
|
|
66
|
+
minutes, secs = divmod(whole, 60)
|
|
67
|
+
if minutes < 60:
|
|
68
|
+
return f"{minutes}m{secs:02d}s"
|
|
69
|
+
hours, minutes = divmod(minutes, 60)
|
|
70
|
+
return f"{hours}h{minutes:02d}m"
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
class RateEstimator:
|
|
74
|
+
"""Units per second within one step, smoothed over the recent events.
|
|
75
|
+
|
|
76
|
+
Units and seconds are smoothed separately and divided, so a burst of events that
|
|
77
|
+
arrive together (a cache hit) does not count as an infinitely fast interval.
|
|
78
|
+
"""
|
|
79
|
+
|
|
80
|
+
def __init__(self) -> None:
|
|
81
|
+
self._alpha = 2 / (RATE_WINDOW_EVENTS + 1)
|
|
82
|
+
self._started = 0.0
|
|
83
|
+
self._start_units = 0
|
|
84
|
+
self._last = 0.0
|
|
85
|
+
self._last_units = 0
|
|
86
|
+
self._units = 0.0
|
|
87
|
+
self._seconds = 0.0
|
|
88
|
+
|
|
89
|
+
def start(self, now: float, completed: int) -> None:
|
|
90
|
+
"""Begin a new step whose work started at `now` with `completed` units done."""
|
|
91
|
+
self._started = self._last = now
|
|
92
|
+
self._start_units = self._last_units = completed
|
|
93
|
+
self._units = self._seconds = 0.0
|
|
94
|
+
|
|
95
|
+
def update(self, now: float, completed: int) -> None:
|
|
96
|
+
units = max(0, completed - self._last_units)
|
|
97
|
+
seconds = max(0.0, now - self._last)
|
|
98
|
+
self._units += self._alpha * (units - self._units)
|
|
99
|
+
self._seconds += self._alpha * (seconds - self._seconds)
|
|
100
|
+
self._last, self._last_units = now, completed
|
|
101
|
+
|
|
102
|
+
def remaining(self, completed: int, total: int) -> float | None:
|
|
103
|
+
"""Seconds left at the recent rate, or None while too little of the step is done."""
|
|
104
|
+
if total <= 0 or completed >= total:
|
|
105
|
+
return None
|
|
106
|
+
warming_up = (
|
|
107
|
+
self._last - self._started < ETA_MIN_SECONDS
|
|
108
|
+
or self._last_units - self._start_units < ETA_MIN_UNITS
|
|
109
|
+
)
|
|
110
|
+
if warming_up or self._units <= 0 or self._seconds <= 0:
|
|
111
|
+
return None
|
|
112
|
+
return (total - completed) * self._seconds / self._units
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
class ProgressDisplay(ABC):
|
|
116
|
+
"""Shared step tracking: notices when the (phase, model) step changes, times it and
|
|
117
|
+
measures its rate for the ETA."""
|
|
118
|
+
|
|
119
|
+
def __init__(self, stream: TextIO, clock: Clock) -> None:
|
|
120
|
+
self._stream = stream
|
|
121
|
+
self._clock = clock
|
|
122
|
+
self._run_started: float | None = None
|
|
123
|
+
self._step: tuple[str, str] | None = None
|
|
124
|
+
self._step_started = 0.0
|
|
125
|
+
self._step_events = 0
|
|
126
|
+
self._rate = RateEstimator()
|
|
127
|
+
self._last_event: tuple[float, int] | None = None
|
|
128
|
+
|
|
129
|
+
def __call__(self, event: ProgressEvent) -> None:
|
|
130
|
+
now = self._clock()
|
|
131
|
+
if self._run_started is None:
|
|
132
|
+
self._run_started = now
|
|
133
|
+
if event.phase == "done":
|
|
134
|
+
self._finish_step(now)
|
|
135
|
+
self._finish_run(now - self._run_started)
|
|
136
|
+
return
|
|
137
|
+
step = _step_key(event)
|
|
138
|
+
new_step = step != self._step
|
|
139
|
+
if new_step:
|
|
140
|
+
self._finish_step(now)
|
|
141
|
+
self._step = step
|
|
142
|
+
self._step_started = now
|
|
143
|
+
self._step_events = 0
|
|
144
|
+
# Events report work already done, so the step's first unit began at the
|
|
145
|
+
# previous event.
|
|
146
|
+
self._rate.start(*(self._last_event or (now, event.completed)))
|
|
147
|
+
self._step_events += 1
|
|
148
|
+
self._rate.update(now, event.completed)
|
|
149
|
+
self._last_event = (now, event.completed)
|
|
150
|
+
if new_step:
|
|
151
|
+
self._start_step(event, now)
|
|
152
|
+
else:
|
|
153
|
+
self._advance(event, now)
|
|
154
|
+
|
|
155
|
+
@abstractmethod
|
|
156
|
+
def close(self) -> None:
|
|
157
|
+
"""Leave the stream at the start of a clean line. Safe to call more than once."""
|
|
158
|
+
|
|
159
|
+
def _remaining(self, event: ProgressEvent) -> float | None:
|
|
160
|
+
return self._rate.remaining(event.completed, event.total)
|
|
161
|
+
|
|
162
|
+
def _finish_step(self, now: float) -> None:
|
|
163
|
+
if self._step is None:
|
|
164
|
+
return
|
|
165
|
+
self.close()
|
|
166
|
+
phase, model = self._step
|
|
167
|
+
if phase == "connect":
|
|
168
|
+
# The runner sends one connect event per server once that server has answered.
|
|
169
|
+
label = f"connect {self._step_events} server{'' if self._step_events == 1 else 's'}"
|
|
170
|
+
else:
|
|
171
|
+
label = f"{phase} {model}".rstrip()
|
|
172
|
+
self._write_line(f" done {label} ({format_duration(now - self._step_started)})")
|
|
173
|
+
self._step = None
|
|
174
|
+
|
|
175
|
+
def _finish_run(self, elapsed: float) -> None:
|
|
176
|
+
self.close()
|
|
177
|
+
self._write_line(f"Finished in {format_duration(elapsed)}")
|
|
178
|
+
|
|
179
|
+
@abstractmethod
|
|
180
|
+
def _start_step(self, event: ProgressEvent, now: float) -> None:
|
|
181
|
+
"""Show the first event of a new step."""
|
|
182
|
+
|
|
183
|
+
@abstractmethod
|
|
184
|
+
def _advance(self, event: ProgressEvent, now: float) -> None:
|
|
185
|
+
"""Show a later event of the current step, if it is worth showing."""
|
|
186
|
+
|
|
187
|
+
def _write_line(self, text: str) -> None:
|
|
188
|
+
self._stream.write(text + "\n")
|
|
189
|
+
self._stream.flush()
|
|
190
|
+
|
|
191
|
+
|
|
192
|
+
class StatusLine(ProgressDisplay):
|
|
193
|
+
"""A single line redrawn in place, for terminals."""
|
|
194
|
+
|
|
195
|
+
def __init__(self, stream: TextIO, clock: Clock) -> None:
|
|
196
|
+
super().__init__(stream, clock)
|
|
197
|
+
self._drawn_width = 0
|
|
198
|
+
self._last_draw: float | None = None
|
|
199
|
+
|
|
200
|
+
def close(self) -> None:
|
|
201
|
+
if self._drawn_width:
|
|
202
|
+
# Overwrite with spaces rather than an ANSI erase: legacy Windows consoles
|
|
203
|
+
# print escape sequences literally unless virtual terminal mode is enabled.
|
|
204
|
+
self._stream.write("\r" + " " * self._drawn_width + "\r")
|
|
205
|
+
self._stream.flush()
|
|
206
|
+
self._drawn_width = 0
|
|
207
|
+
|
|
208
|
+
def _start_step(self, event: ProgressEvent, now: float) -> None:
|
|
209
|
+
self._draw(event, now)
|
|
210
|
+
|
|
211
|
+
def _advance(self, event: ProgressEvent, now: float) -> None:
|
|
212
|
+
if self._last_draw is None or now - self._last_draw >= REDRAW_INTERVAL_SECONDS:
|
|
213
|
+
self._draw(event, now)
|
|
214
|
+
|
|
215
|
+
def _draw(self, event: ProgressEvent, now: float) -> None:
|
|
216
|
+
width = max(1, shutil.get_terminal_size().columns - 1)
|
|
217
|
+
line = status_text(event, self._remaining(event))[:width]
|
|
218
|
+
padding = " " * max(0, self._drawn_width - len(line))
|
|
219
|
+
self._stream.write("\r" + line + padding)
|
|
220
|
+
self._stream.flush()
|
|
221
|
+
self._drawn_width = len(line)
|
|
222
|
+
self._last_draw = now
|
|
223
|
+
|
|
224
|
+
|
|
225
|
+
class StepLog(ProgressDisplay):
|
|
226
|
+
"""Plain lines for pipes and CI logs."""
|
|
227
|
+
|
|
228
|
+
def __init__(self, stream: TextIO, clock: Clock) -> None:
|
|
229
|
+
super().__init__(stream, clock)
|
|
230
|
+
self._logged_step = -1
|
|
231
|
+
|
|
232
|
+
def close(self) -> None:
|
|
233
|
+
"""Nothing to clean up: every log line already ends with a newline."""
|
|
234
|
+
|
|
235
|
+
def _start_step(self, event: ProgressEvent, now: float) -> None:
|
|
236
|
+
self._log(event)
|
|
237
|
+
|
|
238
|
+
def _advance(self, event: ProgressEvent, now: float) -> None:
|
|
239
|
+
if _log_step(event) > self._logged_step:
|
|
240
|
+
self._log(event)
|
|
241
|
+
|
|
242
|
+
def _log(self, event: ProgressEvent) -> None:
|
|
243
|
+
self._logged_step = _log_step(event)
|
|
244
|
+
parts = [f"[{_percent(event):>4}]", event.phase, event.model, event.detail]
|
|
245
|
+
remaining = self._remaining(event)
|
|
246
|
+
if remaining is not None:
|
|
247
|
+
parts.append(f"(ETA {format_duration(remaining)})")
|
|
248
|
+
self._write_line(" ".join(part for part in parts if part))
|
|
249
|
+
|
|
250
|
+
|
|
251
|
+
def status_text(event: ProgressEvent, remaining: float | None) -> str:
|
|
252
|
+
"""The interactive status line, before truncation to the terminal width."""
|
|
253
|
+
fraction = _fraction(event)
|
|
254
|
+
filled = 0 if fraction is None else round(fraction * BAR_WIDTH)
|
|
255
|
+
parts = [f"[{'#' * filled}{'.' * (BAR_WIDTH - filled)}]", f"{_percent(event):>4}"]
|
|
256
|
+
if remaining is not None:
|
|
257
|
+
parts.append(f"ETA {format_duration(remaining)}")
|
|
258
|
+
parts.extend((event.model, f"{event.phase} {event.detail}".rstrip()))
|
|
259
|
+
return " ".join(part for part in parts if part)
|
|
260
|
+
|
|
261
|
+
|
|
262
|
+
def _step_key(event: ProgressEvent) -> tuple[str, str]:
|
|
263
|
+
# The runner connects to every server before any work; a done line per server would
|
|
264
|
+
# time the wrong one, since each event arrives after its connection attempt.
|
|
265
|
+
if event.phase == "connect":
|
|
266
|
+
return ("connect", "")
|
|
267
|
+
return (event.phase, event.model)
|
|
268
|
+
|
|
269
|
+
|
|
270
|
+
def _fraction(event: ProgressEvent) -> float | None:
|
|
271
|
+
if event.total <= 0:
|
|
272
|
+
return None
|
|
273
|
+
return min(1.0, max(0.0, event.completed / event.total))
|
|
274
|
+
|
|
275
|
+
|
|
276
|
+
def _percent(event: ProgressEvent) -> str:
|
|
277
|
+
fraction = _fraction(event)
|
|
278
|
+
return "--%" if fraction is None else f"{int(fraction * 100)}%"
|
|
279
|
+
|
|
280
|
+
|
|
281
|
+
def _log_step(event: ProgressEvent) -> int:
|
|
282
|
+
fraction = _fraction(event)
|
|
283
|
+
return 0 if fraction is None else int(fraction * LOG_STEPS)
|
quantdiff/py.typed
ADDED
|
File without changes
|