zeroquantz 0.1.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.
Files changed (105) hide show
  1. zeroquantz/__init__.py +14 -0
  2. zeroquantz/__main__.py +8 -0
  3. zeroquantz/agent/__init__.py +16 -0
  4. zeroquantz/agent/dispatcher.py +520 -0
  5. zeroquantz/agent/intents.py +46 -0
  6. zeroquantz/agent/parser.py +255 -0
  7. zeroquantz/benchmark/__init__.py +7 -0
  8. zeroquantz/benchmark/latency.py +66 -0
  9. zeroquantz/benchmark/memory.py +41 -0
  10. zeroquantz/benchmark/quality.py +38 -0
  11. zeroquantz/benchmark/runner.py +151 -0
  12. zeroquantz/cli/__init__.py +7 -0
  13. zeroquantz/cli/app.py +98 -0
  14. zeroquantz/cli/commands.py +459 -0
  15. zeroquantz/cli/interactive.py +56 -0
  16. zeroquantz/core/__init__.py +7 -0
  17. zeroquantz/core/artifacts.py +179 -0
  18. zeroquantz/core/context.py +127 -0
  19. zeroquantz/core/events.py +30 -0
  20. zeroquantz/core/exceptions.py +105 -0
  21. zeroquantz/core/session.py +202 -0
  22. zeroquantz/core/subenv.py +202 -0
  23. zeroquantz/deploy/__init__.py +25 -0
  24. zeroquantz/deploy/assets.py +161 -0
  25. zeroquantz/deploy/launcher.py +80 -0
  26. zeroquantz/deploy/runtime_env.py +66 -0
  27. zeroquantz/deploy/targets.py +154 -0
  28. zeroquantz/export/__init__.py +8 -0
  29. zeroquantz/export/exporter.py +68 -0
  30. zeroquantz/export/report.py +203 -0
  31. zeroquantz/hardware/__init__.py +15 -0
  32. zeroquantz/hardware/capabilities.py +152 -0
  33. zeroquantz/hardware/detector.py +200 -0
  34. zeroquantz/hardware/gpu.py +31 -0
  35. zeroquantz/models/__init__.py +8 -0
  36. zeroquantz/models/architecture.py +168 -0
  37. zeroquantz/models/downloader.py +161 -0
  38. zeroquantz/models/hf_auth.py +105 -0
  39. zeroquantz/models/inspector.py +249 -0
  40. zeroquantz/models/metadata.py +108 -0
  41. zeroquantz/models/search.py +71 -0
  42. zeroquantz/optimization/__init__.py +22 -0
  43. zeroquantz/optimization/candidate.py +272 -0
  44. zeroquantz/optimization/constraints.py +70 -0
  45. zeroquantz/optimization/fit.py +203 -0
  46. zeroquantz/optimization/pareto.py +66 -0
  47. zeroquantz/optimization/planner.py +297 -0
  48. zeroquantz/optimization/recommender.py +149 -0
  49. zeroquantz/profiling/__init__.py +18 -0
  50. zeroquantz/profiling/calibration.py +74 -0
  51. zeroquantz/profiling/sensitivity.py +234 -0
  52. zeroquantz/quantization/__init__.py +17 -0
  53. zeroquantz/quantization/backends/__init__.py +8 -0
  54. zeroquantz/quantization/backends/bitsandbytes.py +210 -0
  55. zeroquantz/quantization/backends/torchao.py +198 -0
  56. zeroquantz/quantization/base.py +136 -0
  57. zeroquantz/quantization/catalog.py +321 -0
  58. zeroquantz/quantization/config.py +106 -0
  59. zeroquantz/quantization/gguf_pipeline.py +210 -0
  60. zeroquantz/quantization/isolated.py +248 -0
  61. zeroquantz/quantization/memory.py +133 -0
  62. zeroquantz/quantization/native.py +91 -0
  63. zeroquantz/quantization/registry.py +101 -0
  64. zeroquantz/render.py +341 -0
  65. zeroquantz/runtimes/__init__.py +18 -0
  66. zeroquantz/runtimes/base.py +64 -0
  67. zeroquantz/runtimes/compatibility.py +91 -0
  68. zeroquantz/runtimes/registry.py +70 -0
  69. zeroquantz/runtimes/transformers.py +53 -0
  70. zeroquantz/runtimes/vllm.py +83 -0
  71. zeroquantz/tui/__init__.py +13 -0
  72. zeroquantz/tui/app.py +77 -0
  73. zeroquantz/tui/banner.py +47 -0
  74. zeroquantz/tui/screens/__init__.py +25 -0
  75. zeroquantz/tui/screens/confirm.py +41 -0
  76. zeroquantz/tui/screens/execute.py +194 -0
  77. zeroquantz/tui/screens/model_select.py +206 -0
  78. zeroquantz/tui/screens/plan.py +177 -0
  79. zeroquantz/tui/screens/quantize_select.py +272 -0
  80. zeroquantz/tui/screens/settings.py +219 -0
  81. zeroquantz/tui/screens/token.py +94 -0
  82. zeroquantz/tui/screens/welcome.py +128 -0
  83. zeroquantz/tui/screens/workspace.py +175 -0
  84. zeroquantz/tui/styles/app.tcss +424 -0
  85. zeroquantz/tui/widgets/__init__.py +9 -0
  86. zeroquantz/tui/widgets/chip.py +36 -0
  87. zeroquantz/tui/widgets/sidebar.py +107 -0
  88. zeroquantz/tui/widgets/status_bar.py +43 -0
  89. zeroquantz/utils/__init__.py +8 -0
  90. zeroquantz/utils/config.py +46 -0
  91. zeroquantz/utils/env.py +78 -0
  92. zeroquantz/utils/logging.py +73 -0
  93. zeroquantz/utils/metrics.py +98 -0
  94. zeroquantz/utils/paths.py +57 -0
  95. zeroquantz/utils/units.py +134 -0
  96. zeroquantz/verification/__init__.py +17 -0
  97. zeroquantz/verification/logits.py +55 -0
  98. zeroquantz/verification/report.py +186 -0
  99. zeroquantz/verification/weights.py +44 -0
  100. zeroquantz/version.py +8 -0
  101. zeroquantz-0.1.0.dist-info/METADATA +72 -0
  102. zeroquantz-0.1.0.dist-info/RECORD +105 -0
  103. zeroquantz-0.1.0.dist-info/WHEEL +4 -0
  104. zeroquantz-0.1.0.dist-info/entry_points.txt +2 -0
  105. zeroquantz-0.1.0.dist-info/licenses/LICENSE +201 -0
@@ -0,0 +1,78 @@
1
+ """Capture the software/hardware environment for reproducible reports."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import importlib.metadata
6
+ import importlib.util
7
+ import os
8
+ import platform
9
+ from datetime import UTC, datetime
10
+ from typing import TYPE_CHECKING, Any
11
+
12
+ if TYPE_CHECKING:
13
+ from zeroquantz.hardware.capabilities import HardwareProfile
14
+
15
+
16
+ def _pkg_version(name: str) -> str | None:
17
+ try:
18
+ return importlib.metadata.version(name)
19
+ except importlib.metadata.PackageNotFoundError:
20
+ return None
21
+
22
+
23
+ def enable_fast_downloads() -> bool:
24
+ """Turn on the Rust ``hf_transfer`` accelerator for Hub downloads, if present.
25
+
26
+ ``hf_transfer`` is HuggingFace's Rust, multi-threaded, chunked downloader; it
27
+ can multiply throughput on fast links where a single HTTPS connection is the
28
+ bottleneck (on a bandwidth-capped connection it makes little difference).
29
+
30
+ Crucially, ``huggingface_hub`` *raises* at download time if
31
+ ``HF_HUB_ENABLE_HF_TRANSFER=1`` is set but the package is missing, so we only
32
+ set the flag when ``hf_transfer`` is actually importable — and use
33
+ ``setdefault`` so an explicit ``HF_HUB_ENABLE_HF_TRANSFER=0`` from the user's
34
+ environment is always respected. Returns whether acceleration is now active.
35
+ """
36
+ if importlib.util.find_spec("hf_transfer") is None:
37
+ return False
38
+ os.environ.setdefault("HF_HUB_ENABLE_HF_TRANSFER", "1")
39
+ return os.environ.get("HF_HUB_ENABLE_HF_TRANSFER") == "1"
40
+
41
+
42
+ def fast_downloads_active() -> bool:
43
+ """Whether ``hf_transfer`` acceleration is currently enabled in this process."""
44
+ return (
45
+ os.environ.get("HF_HUB_ENABLE_HF_TRANSFER") == "1"
46
+ and importlib.util.find_spec("hf_transfer") is not None
47
+ )
48
+
49
+
50
+ def utc_timestamp() -> str:
51
+ return datetime.now(UTC).isoformat(timespec="seconds")
52
+
53
+
54
+ def capture_environment(hardware: HardwareProfile | None = None) -> dict[str, Any]:
55
+ """A JSON-serializable snapshot of versions, platform, and (optional) GPU."""
56
+ from zeroquantz import __version__
57
+
58
+ env: dict[str, Any] = {
59
+ "zeroquantz_version": __version__,
60
+ "timestamp": utc_timestamp(),
61
+ "python": platform.python_version(),
62
+ "platform": platform.platform(),
63
+ "torch": _pkg_version("torch"),
64
+ "transformers": _pkg_version("transformers"),
65
+ "bitsandbytes": _pkg_version("bitsandbytes"),
66
+ "torchao": _pkg_version("torchao"),
67
+ }
68
+ if hardware is not None:
69
+ env["gpu"] = hardware.gpu_name
70
+ env["vram_gb"] = hardware.total_vram_gb
71
+ env["compute_capability"] = (
72
+ ".".join(map(str, hardware.compute_capability))
73
+ if hardware.compute_capability
74
+ else None
75
+ )
76
+ env["cuda_version"] = hardware.cuda_version
77
+ env["driver_version"] = hardware.driver_version
78
+ return env
@@ -0,0 +1,73 @@
1
+ """Structured logging setup.
2
+
3
+ The interactive UI shows concise messages; detailed logs are written to
4
+ ``~/.zeroquantz/logs/zeroquantz.log``. ``--verbose`` raises the console level to
5
+ INFO and ``--debug`` to DEBUG.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import logging
11
+ from logging.handlers import RotatingFileHandler
12
+ from typing import Literal
13
+
14
+ from zeroquantz.utils.paths import paths
15
+
16
+ _CONFIGURED = False
17
+
18
+ Verbosity = Literal["quiet", "verbose", "debug"]
19
+
20
+
21
+ def get_logger(name: str) -> logging.Logger:
22
+ """Return a namespaced child of the ``zeroquantz`` logger."""
23
+ if not name.startswith("zeroquantz"):
24
+ name = f"zeroquantz.{name}"
25
+ return logging.getLogger(name)
26
+
27
+
28
+ def configure_logging(verbosity: Verbosity = "quiet", *, log_to_file: bool = True) -> None:
29
+ """Configure the root ``zeroquantz`` logger. Safe to call more than once."""
30
+ global _CONFIGURED
31
+ root = logging.getLogger("zeroquantz")
32
+ root.setLevel(logging.DEBUG)
33
+ root.handlers.clear()
34
+
35
+ console_level = {
36
+ "quiet": logging.WARNING,
37
+ "verbose": logging.INFO,
38
+ "debug": logging.DEBUG,
39
+ }[verbosity]
40
+
41
+ console = logging.StreamHandler()
42
+ console.setLevel(console_level)
43
+ console.setFormatter(logging.Formatter("%(levelname)s %(name)s: %(message)s"))
44
+ root.addHandler(console)
45
+
46
+ if log_to_file:
47
+ try:
48
+ log_dir = paths().ensure().logs
49
+ file_handler = RotatingFileHandler(
50
+ log_dir / "zeroquantz.log",
51
+ maxBytes=2 * 1024 * 1024,
52
+ backupCount=3,
53
+ encoding="utf-8",
54
+ )
55
+ file_handler.setLevel(logging.DEBUG)
56
+ file_handler.setFormatter(
57
+ logging.Formatter(
58
+ "%(asctime)s %(levelname)-7s %(name)s %(message)s",
59
+ datefmt="%Y-%m-%dT%H:%M:%S",
60
+ )
61
+ )
62
+ root.addHandler(file_handler)
63
+ except OSError:
64
+ # A read-only home directory should never take the app down.
65
+ pass
66
+
67
+ root.propagate = False
68
+ _CONFIGURED = True
69
+
70
+
71
+ def ensure_configured() -> None:
72
+ if not _CONFIGURED:
73
+ configure_logging("quiet")
@@ -0,0 +1,98 @@
1
+ """Pure numeric comparison metrics (NumPy only).
2
+
3
+ Shared by sensitivity profiling and verification so the definitions live in one
4
+ place. Everything accepts array-likes (Python lists, NumPy arrays, or detached
5
+ tensors converted to NumPy) and returns plain floats.
6
+
7
+ NumPy is imported lazily inside each function rather than at module import time:
8
+ these metrics are only used on the profiling/verification paths, so keeping the
9
+ import local shaves NumPy (~100 ms) off ``zeroquantz`` cold-start, which imports
10
+ this module transitively via the dispatcher. The import is cached after first
11
+ use, so repeated calls pay only a dict lookup.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ from typing import TYPE_CHECKING
17
+
18
+ if TYPE_CHECKING:
19
+ import numpy as np
20
+ from numpy.typing import ArrayLike
21
+
22
+
23
+ def _arr(x: ArrayLike) -> np.ndarray:
24
+ import numpy as np
25
+
26
+ return np.asarray(x, dtype=np.float64)
27
+
28
+
29
+ def mse(a: ArrayLike, b: ArrayLike) -> float:
30
+ """Mean squared error between two equally-shaped arrays."""
31
+ import numpy as np
32
+
33
+ av, bv = _arr(a), _arr(b)
34
+ return float(np.mean((av - bv) ** 2))
35
+
36
+
37
+ def cosine_similarity(a: ArrayLike, b: ArrayLike) -> float:
38
+ """Cosine similarity of two arrays, flattened to vectors."""
39
+ import numpy as np
40
+
41
+ av, bv = _arr(a).ravel(), _arr(b).ravel()
42
+ na = float(np.linalg.norm(av))
43
+ nb = float(np.linalg.norm(bv))
44
+ if na == 0.0 or nb == 0.0:
45
+ return 0.0
46
+ return float(np.dot(av, bv) / (na * nb))
47
+
48
+
49
+ def softmax(logits: ArrayLike, axis: int = -1) -> np.ndarray:
50
+ import numpy as np
51
+
52
+ x = _arr(logits)
53
+ x = x - np.max(x, axis=axis, keepdims=True)
54
+ e = np.exp(x)
55
+ return e / np.sum(e, axis=axis, keepdims=True)
56
+
57
+
58
+ def kl_divergence(p_logits: ArrayLike, q_logits: ArrayLike, *, eps: float = 1e-12) -> float:
59
+ """Mean KL(P || Q) where P, Q are softmaxes of the given logits.
60
+
61
+ Works for 1D (single distribution) or 2D (a batch of rows, averaged) input.
62
+ """
63
+ import numpy as np
64
+
65
+ p = softmax(p_logits)
66
+ q = softmax(q_logits)
67
+ p = np.clip(p, eps, 1.0)
68
+ q = np.clip(q, eps, 1.0)
69
+ kl = np.sum(p * (np.log(p) - np.log(q)), axis=-1)
70
+ return float(np.mean(kl))
71
+
72
+
73
+ def top1_agreement(p_logits: ArrayLike, q_logits: ArrayLike) -> float:
74
+ """Fraction of rows whose argmax agrees between the two logit arrays."""
75
+ import numpy as np
76
+
77
+ p = _arr(p_logits)
78
+ q = _arr(q_logits)
79
+ if p.ndim == 1:
80
+ return float(np.argmax(p) == np.argmax(q))
81
+ return float(np.mean(np.argmax(p, axis=-1) == np.argmax(q, axis=-1)))
82
+
83
+
84
+ def perplexity(nlls: ArrayLike) -> float:
85
+ """Perplexity from a sequence of per-token negative log-likelihoods (nats)."""
86
+ import numpy as np
87
+
88
+ v = _arr(nlls)
89
+ if v.size == 0:
90
+ return float("nan")
91
+ return float(np.exp(np.mean(v)))
92
+
93
+
94
+ def relative_delta(baseline: float, candidate: float) -> float:
95
+ """Signed relative change ``(candidate - baseline) / baseline``."""
96
+ if baseline == 0:
97
+ return float("nan")
98
+ return (candidate - baseline) / baseline
@@ -0,0 +1,57 @@
1
+ """Filesystem locations ZeroQuantz reads and writes.
2
+
3
+ The layout follows the spec's ``~/.zeroquantz`` convention, but the root is
4
+ overridable with the ``ZEROQUANTZ_HOME`` environment variable (useful for tests
5
+ and CI, which point it at a temp directory).
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import os
11
+ from dataclasses import dataclass
12
+ from pathlib import Path
13
+
14
+
15
+ def _default_home() -> Path:
16
+ override = os.environ.get("ZEROQUANTZ_HOME")
17
+ if override:
18
+ return Path(override).expanduser()
19
+ return Path.home() / ".zeroquantz"
20
+
21
+
22
+ @dataclass(frozen=True)
23
+ class ZeroQuantzPaths:
24
+ """Resolved application directories. Created lazily via :meth:`ensure`."""
25
+
26
+ home: Path
27
+
28
+ @classmethod
29
+ def resolve(cls) -> ZeroQuantzPaths:
30
+ return cls(home=_default_home())
31
+
32
+ @property
33
+ def sessions(self) -> Path:
34
+ return self.home / "sessions"
35
+
36
+ @property
37
+ def logs(self) -> Path:
38
+ return self.home / "logs"
39
+
40
+ @property
41
+ def cache(self) -> Path:
42
+ return self.home / "cache"
43
+
44
+ @property
45
+ def config_file(self) -> Path:
46
+ return self.home / "config.toml"
47
+
48
+ def ensure(self) -> ZeroQuantzPaths:
49
+ """Create the standard directories if they do not already exist."""
50
+ for directory in (self.home, self.sessions, self.logs, self.cache):
51
+ directory.mkdir(parents=True, exist_ok=True)
52
+ return self
53
+
54
+
55
+ def paths() -> ZeroQuantzPaths:
56
+ """Return freshly-resolved paths (re-reads ``ZEROQUANTZ_HOME`` each call)."""
57
+ return ZeroQuantzPaths.resolve()
@@ -0,0 +1,134 @@
1
+ """Byte/parameter/precision math and the one place ZeroQuantz defines its units.
2
+
3
+ **Unit convention.** ZeroQuantz reports memory in *binary gigabytes* (GiB,
4
+ ``2**30`` bytes) but labels them ``GB`` in the UI, because that is how GPU VRAM
5
+ and user memory budgets are universally stated ("a 24 GB 4090", "fit it under
6
+ 8 GB"). Keeping VRAM, model size, and user budgets in one consistent binary unit
7
+ is what lets ``estimated_vram_gb <= max_vram_gb`` mean something. Every
8
+ conversion in the codebase goes through :func:`bytes_to_gb` / :func:`gb_to_bytes`
9
+ so the convention lives in exactly one file.
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ from typing import Final
15
+
16
+ BYTES_PER_GIB: Final[int] = 1024**3
17
+ BYTES_PER_MIB: Final[int] = 1024**2
18
+
19
+ # Effective *bits per stored weight* for each precision. Sub-byte quantization
20
+ # schemes carry group-wise metadata (scales, and often zero-points), so their
21
+ # effective footprint is higher than the nominal bit-width. These defaults
22
+ # assume a group size of 128 with fp16 scales, which is the common case; the
23
+ # candidate generator can override them per-config.
24
+ BITS_PER_PARAM: Final[dict[str, float]] = {
25
+ "float32": 32.0,
26
+ "fp32": 32.0,
27
+ "float16": 16.0,
28
+ "fp16": 16.0,
29
+ "bfloat16": 16.0,
30
+ "bf16": 16.0,
31
+ "int8": 8.125, # 8 + fp16 scale per group of 128
32
+ "fp8": 8.0,
33
+ "fp8_e4m3": 8.0,
34
+ "fp8_e5m2": 8.0,
35
+ "int4": 4.25, # 4 + fp16 scale + int4 zero-point over a group of 128
36
+ "nf4": 4.127, # NF4 with double quantization (bitsandbytes)
37
+ "fp4": 4.25,
38
+ "int3": 3.25,
39
+ "int2": 2.25,
40
+ }
41
+
42
+ # Canonical display/short names for the precisions we reason about.
43
+ DTYPE_ALIASES: Final[dict[str, str]] = {
44
+ "torch.float32": "float32",
45
+ "torch.float16": "float16",
46
+ "torch.bfloat16": "bfloat16",
47
+ "float32": "float32",
48
+ "float16": "float16",
49
+ "bfloat16": "bfloat16",
50
+ "f32": "float32",
51
+ "f16": "float16",
52
+ "bf16": "bfloat16",
53
+ "fp32": "float32",
54
+ "fp16": "float16",
55
+ "half": "float16",
56
+ "int8": "int8",
57
+ "i8": "int8",
58
+ "uint8": "int8",
59
+ "int4": "int4",
60
+ "i4": "int4",
61
+ "nf4": "nf4",
62
+ "fp4": "fp4",
63
+ "fp8": "fp8",
64
+ }
65
+
66
+
67
+ def normalize_dtype(dtype: str) -> str:
68
+ """Return a canonical dtype name (e.g. ``"bf16"`` -> ``"bfloat16"``)."""
69
+ key = str(dtype).strip().lower()
70
+ return DTYPE_ALIASES.get(key, key)
71
+
72
+
73
+ def bits_per_param(precision: str) -> float:
74
+ """Effective stored bits per parameter for ``precision``.
75
+
76
+ Raises:
77
+ KeyError: if the precision is unknown. Callers that want a soft failure
78
+ should catch this and treat the config as unsupported.
79
+ """
80
+ key = normalize_dtype(precision)
81
+ if key in BITS_PER_PARAM:
82
+ return BITS_PER_PARAM[key]
83
+ # Also accept the already-canonical short forms used as dict keys above.
84
+ if precision.lower() in BITS_PER_PARAM:
85
+ return BITS_PER_PARAM[precision.lower()]
86
+ raise KeyError(f"unknown precision: {precision!r}")
87
+
88
+
89
+ def bytes_per_param(precision: str) -> float:
90
+ """Effective stored bytes per parameter for ``precision``."""
91
+ return bits_per_param(precision) / 8.0
92
+
93
+
94
+ def bytes_to_gb(num_bytes: float) -> float:
95
+ """Bytes -> binary gigabytes (GiB), the ZeroQuantz canonical unit."""
96
+ return num_bytes / BYTES_PER_GIB
97
+
98
+
99
+ def gb_to_bytes(gb: float) -> float:
100
+ """Binary gigabytes (GiB) -> bytes."""
101
+ return gb * BYTES_PER_GIB
102
+
103
+
104
+ def param_bytes(num_params: float, precision: str) -> float:
105
+ """Total bytes to store ``num_params`` parameters at ``precision``."""
106
+ return num_params * bytes_per_param(precision)
107
+
108
+
109
+ def param_gb(num_params: float, precision: str) -> float:
110
+ """Total GiB to store ``num_params`` parameters at ``precision``."""
111
+ return bytes_to_gb(param_bytes(num_params, precision))
112
+
113
+
114
+ def humanize_params(num_params: float) -> str:
115
+ """Render a parameter count as e.g. ``"8.2B"`` / ``"350M"`` / ``"1.3K"``."""
116
+ n = float(num_params)
117
+ for divisor, suffix in ((1e12, "T"), (1e9, "B"), (1e6, "M"), (1e3, "K")):
118
+ if abs(n) >= divisor:
119
+ return f"{n / divisor:.1f}{suffix}"
120
+ return str(int(n))
121
+
122
+
123
+ def humanize_bytes(num_bytes: float) -> str:
124
+ """Render a byte count with a binary suffix (``"5.3 GB"``, ``"812 MB"``)."""
125
+ n = float(num_bytes)
126
+ for divisor, suffix in (
127
+ (1024**4, "TB"),
128
+ (1024**3, "GB"),
129
+ (1024**2, "MB"),
130
+ (1024, "KB"),
131
+ ):
132
+ if abs(n) >= divisor:
133
+ return f"{n / divisor:.2f} {suffix}"
134
+ return f"{int(n)} B"
@@ -0,0 +1,17 @@
1
+ """Verification: compare a quantized model against its full-precision baseline."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from zeroquantz.verification.report import (
6
+ VerificationReport,
7
+ VerificationThresholds,
8
+ Verifier,
9
+ evaluate_thresholds,
10
+ )
11
+
12
+ __all__ = [
13
+ "VerificationReport",
14
+ "VerificationThresholds",
15
+ "Verifier",
16
+ "evaluate_thresholds",
17
+ ]
@@ -0,0 +1,55 @@
1
+ """Compare logits between a baseline and a quantized model (lazy torch)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import dataclass
6
+ from typing import Any
7
+
8
+ from zeroquantz.utils import metrics
9
+
10
+
11
+ @dataclass
12
+ class LogitComparison:
13
+ samples: int
14
+ logit_cosine: float
15
+ mean_kl: float
16
+ top1_agreement: float
17
+
18
+
19
+ def compare_logits(
20
+ base_model: Any,
21
+ quant_model: Any,
22
+ tokenizer: Any,
23
+ prompts: list[str],
24
+ *,
25
+ max_length: int = 128,
26
+ ) -> LogitComparison:
27
+ """Run both models on ``prompts`` and compare their output logits."""
28
+ import numpy as np
29
+ import torch
30
+
31
+ base_logits: list[np.ndarray] = []
32
+ quant_logits: list[np.ndarray] = []
33
+
34
+ base_device = next(base_model.parameters()).device
35
+ quant_device = next(quant_model.parameters()).device
36
+
37
+ with torch.no_grad():
38
+ for prompt in prompts:
39
+ enc = tokenizer(
40
+ prompt, return_tensors="pt", truncation=True, max_length=max_length
41
+ )
42
+ b = base_model(**{k: v.to(base_device) for k, v in enc.items()})
43
+ q = quant_model(**{k: v.to(quant_device) for k, v in enc.items()})
44
+ # last-token logits
45
+ base_logits.append(b.logits[0, -1].float().cpu().numpy())
46
+ quant_logits.append(q.logits[0, -1].float().cpu().numpy())
47
+
48
+ base_arr = np.stack(base_logits)
49
+ quant_arr = np.stack(quant_logits)
50
+ return LogitComparison(
51
+ samples=len(prompts),
52
+ logit_cosine=round(metrics.cosine_similarity(base_arr, quant_arr), 6),
53
+ mean_kl=round(metrics.kl_divergence(base_arr, quant_arr), 6),
54
+ top1_agreement=round(metrics.top1_agreement(base_arr, quant_arr), 6),
55
+ )