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,234 @@
1
+ """Layer-sensitivity profiling.
2
+
3
+ Estimates how much each transformer sub-module (attention / MLP) degrades model
4
+ outputs when pushed to low precision, so the mixed-precision planner can keep the
5
+ sensitive ones higher. Two paths:
6
+
7
+ * :meth:`SensitivityProfiler.heuristic` — a zero-cost prior (edge layers are most
8
+ sensitive). Always available; used offline and as a fallback.
9
+ * :meth:`SensitivityProfiler.profile` — the real measurement (spec §11): baseline
10
+ outputs, then per-module simulated quantization, measure deviation, restore.
11
+ Requires torch + transformers and is imported lazily.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ from typing import TYPE_CHECKING, Any
17
+
18
+ from pydantic import BaseModel, Field
19
+
20
+ from zeroquantz.core.exceptions import DependencyError
21
+ from zeroquantz.utils import metrics
22
+ from zeroquantz.utils.logging import get_logger
23
+
24
+ if TYPE_CHECKING:
25
+ from collections.abc import Callable
26
+
27
+ from zeroquantz.models.metadata import ModelProfile
28
+ from zeroquantz.profiling.calibration import CalibrationDataset
29
+
30
+ log = get_logger(__name__)
31
+
32
+
33
+ def layer_sensitivity_prior(layer_idx: int, num_layers: int, kind: str = "attn") -> float:
34
+ """Prior sensitivity in [0, 1]: layers near input/output are most sensitive,
35
+ attention slightly more than MLP."""
36
+ if num_layers <= 1:
37
+ base = 1.0
38
+ else:
39
+ edge_distance = min(layer_idx, num_layers - 1 - layer_idx)
40
+ base = 1.0 - edge_distance / ((num_layers - 1) / 2)
41
+ base = max(0.0, min(1.0, base))
42
+ return round(base * (1.0 if kind == "attn" else 0.9), 4)
43
+
44
+
45
+ class LayerSensitivity(BaseModel):
46
+ layer: str
47
+ score: float # normalized 0..1, higher = more sensitive
48
+ raw: float | None = None # raw deviation metric, if measured
49
+ metric: str = "prior"
50
+
51
+
52
+ class SensitivityProfile(BaseModel):
53
+ model_id: str
54
+ metric: str = "prior"
55
+ target_precision: str = "int4"
56
+ layers: list[LayerSensitivity] = Field(default_factory=list)
57
+
58
+ def as_lookup(self) -> dict[str, float]:
59
+ return {ls.layer: ls.score for ls in self.layers}
60
+
61
+ def ranked(self) -> list[LayerSensitivity]:
62
+ return sorted(self.layers, key=lambda ls: ls.score, reverse=True)
63
+
64
+ def top(self, n: int) -> list[LayerSensitivity]:
65
+ return self.ranked()[:n]
66
+
67
+
68
+ def _normalize_scores(raw: dict[str, float]) -> dict[str, float]:
69
+ if not raw:
70
+ return {}
71
+ lo, hi = min(raw.values()), max(raw.values())
72
+ if hi <= lo:
73
+ return {k: 0.5 for k in raw}
74
+ return {k: round((v - lo) / (hi - lo), 4) for k, v in raw.items()}
75
+
76
+
77
+ class SensitivityProfiler:
78
+ """Produce a :class:`SensitivityProfile` for a model."""
79
+
80
+ @staticmethod
81
+ def heuristic(model: ModelProfile) -> SensitivityProfile:
82
+ layers = model.num_layers or 0
83
+ entries: list[LayerSensitivity] = []
84
+ for i in range(layers):
85
+ for kind, mod in (("attn", "self_attn"), ("mlp", "mlp")):
86
+ entries.append(
87
+ LayerSensitivity(
88
+ layer=f"model.layers.{i}.{mod}",
89
+ score=layer_sensitivity_prior(i, layers, kind),
90
+ metric="prior",
91
+ )
92
+ )
93
+ return SensitivityProfile(model_id=model.model_id, metric="prior", layers=entries)
94
+
95
+ @staticmethod
96
+ def profile(
97
+ model_ref: str,
98
+ calibration: CalibrationDataset | None = None,
99
+ *,
100
+ metric: str = "mse",
101
+ target_bits: int = 4,
102
+ group_size: int = 128,
103
+ max_layers: int | None = None,
104
+ max_length: int = 128,
105
+ progress: Callable[[str, float], None] | None = None,
106
+ ) -> SensitivityProfile:
107
+ """Measure per-module sensitivity (requires torch + transformers).
108
+
109
+ For each layer's attention and MLP block, all linear weights are replaced
110
+ with a simulated-quantized round-trip, the deviation of the output logits
111
+ from baseline is measured, and the original weights are restored.
112
+ """
113
+ try:
114
+ import torch
115
+ from transformers import AutoModelForCausalLM, AutoTokenizer
116
+ except ImportError as exc:
117
+ raise DependencyError.for_extra(
118
+ "transformers", "torch", purpose="profile layer sensitivity"
119
+ ) from exc
120
+
121
+ from zeroquantz.profiling.calibration import default_calibration
122
+
123
+ calibration = calibration or default_calibration()
124
+
125
+ if progress:
126
+ progress("Loading model", 0.02)
127
+ tokenizer = AutoTokenizer.from_pretrained(model_ref)
128
+ model = AutoModelForCausalLM.from_pretrained(
129
+ model_ref, torch_dtype=torch.float16, device_map="auto"
130
+ )
131
+ model.eval()
132
+ device = next(model.parameters()).device
133
+
134
+ prompts = calibration.sample(min(len(calibration), 8))
135
+ enc = tokenizer(
136
+ prompts,
137
+ return_tensors="pt",
138
+ padding=True,
139
+ truncation=True,
140
+ max_length=max_length,
141
+ ).to(device)
142
+
143
+ with torch.no_grad():
144
+ baseline = model(**enc).logits.float().cpu().numpy()
145
+
146
+ blocks = SensitivityProfiler._decoder_blocks(model)
147
+ if max_layers is not None:
148
+ blocks = blocks[:max_layers]
149
+
150
+ raw: dict[str, float] = {}
151
+ total = len(blocks) * 2 or 1
152
+ step = 0
153
+ with torch.no_grad():
154
+ for layer_idx, block in blocks:
155
+ for kind, submodule_name in (("attn", "self_attn"), ("mlp", "mlp")):
156
+ submodule = getattr(block, submodule_name, None)
157
+ if submodule is None:
158
+ step += 1
159
+ continue
160
+ saved = SensitivityProfiler._fake_quantize_submodule(
161
+ submodule, torch, target_bits, group_size
162
+ )
163
+ logits = model(**enc).logits.float().cpu().numpy()
164
+ raw[f"model.layers.{layer_idx}.{submodule_name}"] = _deviation(
165
+ baseline, logits, metric
166
+ )
167
+ SensitivityProfiler._restore(submodule, saved)
168
+ step += 1
169
+ if progress:
170
+ progress(f"layer {layer_idx} {kind}", step / total)
171
+
172
+ normalized = _normalize_scores(raw)
173
+ entries = [
174
+ LayerSensitivity(layer=k, score=v, raw=raw[k], metric=metric)
175
+ for k, v in normalized.items()
176
+ ]
177
+ return SensitivityProfile(
178
+ model_id=model_ref,
179
+ metric=metric,
180
+ target_precision=f"int{target_bits}",
181
+ layers=entries,
182
+ )
183
+
184
+ # ---- torch helpers (only reached on the real path) ----------------------
185
+
186
+ @staticmethod
187
+ def _decoder_blocks(model: Any) -> list[tuple[int, Any]]:
188
+ for attr in ("model", "transformer", "gpt_neox"):
189
+ inner = getattr(model, attr, None)
190
+ layers = getattr(inner, "layers", None) if inner is not None else None
191
+ if layers is not None:
192
+ return list(enumerate(layers))
193
+ layers = getattr(model, "layers", None)
194
+ return list(enumerate(layers)) if layers is not None else []
195
+
196
+ @staticmethod
197
+ def _fake_quantize_submodule(submodule, torch, bits: int, group_size: int) -> dict:
198
+ saved: dict[str, Any] = {}
199
+ for name, param in submodule.named_parameters(recurse=True):
200
+ if param.dim() != 2: # only quantize linear weight matrices
201
+ continue
202
+ saved[name] = param.detach().clone()
203
+ param.copy_(_fake_quantize_tensor(param.detach(), torch, bits, group_size))
204
+ return saved
205
+
206
+ @staticmethod
207
+ def _restore(submodule, saved: dict) -> None:
208
+ params = dict(submodule.named_parameters(recurse=True))
209
+ for name, value in saved.items():
210
+ params[name].copy_(value)
211
+
212
+
213
+ def _fake_quantize_tensor(weight, torch, bits: int, group_size: int):
214
+ """Symmetric per-group fake quantization round-trip (dequantized result)."""
215
+ qmax = 2 ** (bits - 1) - 1
216
+ w = weight.float()
217
+ out_features, in_features = w.shape
218
+ gs = group_size if group_size and in_features % group_size == 0 else in_features
219
+ w = w.reshape(out_features, in_features // gs, gs)
220
+ scale = w.abs().amax(dim=-1, keepdim=True) / max(qmax, 1)
221
+ scale = torch.where(scale == 0, torch.ones_like(scale), scale)
222
+ q = torch.clamp(torch.round(w / scale), -qmax - 1, qmax)
223
+ deq = (q * scale).reshape(out_features, in_features)
224
+ return deq.to(weight.dtype)
225
+
226
+
227
+ def _deviation(baseline, candidate, metric: str) -> float:
228
+ if metric == "cosine":
229
+ return 1.0 - metrics.cosine_similarity(baseline, candidate)
230
+ if metric == "kl":
231
+ b = baseline.reshape(-1, baseline.shape[-1])
232
+ c = candidate.reshape(-1, candidate.shape[-1])
233
+ return metrics.kl_divergence(b, c)
234
+ return metrics.mse(baseline, candidate)
@@ -0,0 +1,17 @@
1
+ """Quantization backends, configs, and the backend registry."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from zeroquantz.quantization.base import MemoryEstimate, QuantizationBackend, QuantizationResult
6
+ from zeroquantz.quantization.config import QuantizationConfig, QuantMethod
7
+ from zeroquantz.quantization.registry import BackendRegistry, default_registry
8
+
9
+ __all__ = [
10
+ "BackendRegistry",
11
+ "MemoryEstimate",
12
+ "QuantMethod",
13
+ "QuantizationBackend",
14
+ "QuantizationConfig",
15
+ "QuantizationResult",
16
+ "default_registry",
17
+ ]
@@ -0,0 +1,8 @@
1
+ """Built-in quantization backends."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from zeroquantz.quantization.backends.bitsandbytes import BitsAndBytesBackend
6
+ from zeroquantz.quantization.backends.torchao import TorchAOBackend
7
+
8
+ __all__ = ["BitsAndBytesBackend", "TorchAOBackend"]
@@ -0,0 +1,210 @@
1
+ """bitsandbytes backend: LLM.int8() weight/activation INT8 and 4-bit NF4.
2
+
3
+ bitsandbytes quantizes *at load time* via ``transformers``' ``BitsAndBytesConfig``,
4
+ so :meth:`quantize` takes a model id / local path rather than a live module.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import importlib.util
10
+ import time
11
+ from typing import TYPE_CHECKING, Any
12
+
13
+ from zeroquantz.core.exceptions import DependencyError, QuantizationError, ZeroQuantzError
14
+ from zeroquantz.quantization import memory
15
+ from zeroquantz.quantization.base import (
16
+ MemoryEstimate,
17
+ QuantizationBackend,
18
+ QuantizationResult,
19
+ )
20
+ from zeroquantz.quantization.config import QuantizationConfig, QuantMethod
21
+ from zeroquantz.utils import units
22
+ from zeroquantz.utils.logging import get_logger
23
+
24
+ if TYPE_CHECKING:
25
+ from collections.abc import Callable
26
+
27
+ from zeroquantz.hardware.capabilities import HardwareProfile
28
+ from zeroquantz.models.metadata import ModelProfile
29
+ from zeroquantz.optimization.constraints import OptimizationGoal
30
+
31
+ log = get_logger(__name__)
32
+
33
+
34
+ class BitsAndBytesBackend(QuantizationBackend):
35
+ name = "bitsandbytes"
36
+ description = "LLM.int8() 8-bit and 4-bit NF4 quantization (quantize-on-load)."
37
+ methods = (QuantMethod.BNB_NF4.value, QuantMethod.BNB_INT8.value)
38
+
39
+ def is_available(self) -> bool:
40
+ return all(
41
+ importlib.util.find_spec(m) is not None
42
+ for m in ("bitsandbytes", "torch", "transformers")
43
+ )
44
+
45
+ def availability_error(self) -> ZeroQuantzError | None:
46
+ if self.is_available():
47
+ return None
48
+ missing = next(
49
+ m
50
+ for m in ("bitsandbytes", "torch", "transformers")
51
+ if importlib.util.find_spec(m) is None
52
+ )
53
+ extra = "quant" if missing == "bitsandbytes" else "torch"
54
+ return DependencyError.for_extra(
55
+ missing, extra, purpose="quantize with bitsandbytes"
56
+ )
57
+
58
+ def supports(
59
+ self,
60
+ model: ModelProfile,
61
+ hardware: HardwareProfile,
62
+ runtime: str | None = None,
63
+ config: QuantizationConfig | None = None,
64
+ ) -> bool:
65
+ if not hardware.cuda_available:
66
+ return False # bitsandbytes is CUDA-only in practice
67
+ method = config.method if config else None
68
+ if method == QuantMethod.BNB_INT8.value:
69
+ return hardware.supports("int8")
70
+ if method == QuantMethod.BNB_NF4.value:
71
+ return hardware.supports("int4")
72
+ # No specific method requested: available if the GPU can do either.
73
+ return hardware.supports("int8") or hardware.supports("int4")
74
+
75
+ def default_config(
76
+ self,
77
+ method: str,
78
+ model: ModelProfile,
79
+ hardware: HardwareProfile,
80
+ goal: OptimizationGoal | None = None,
81
+ ) -> QuantizationConfig:
82
+ compute_dtype = "bfloat16" if hardware.supports_bf16 else "float16"
83
+ if method == QuantMethod.BNB_NF4.value:
84
+ return QuantizationConfig(
85
+ backend=self.name,
86
+ method=method,
87
+ weight_bits=4,
88
+ activation_bits=16,
89
+ group_size=64,
90
+ compute_dtype=compute_dtype,
91
+ double_quant=True,
92
+ quant_type="nf4",
93
+ skip_modules=["lm_head"],
94
+ )
95
+ if method == QuantMethod.BNB_INT8.value:
96
+ return QuantizationConfig(
97
+ backend=self.name,
98
+ method=method,
99
+ weight_bits=8,
100
+ activation_bits=8,
101
+ group_size=None,
102
+ compute_dtype=compute_dtype,
103
+ skip_modules=["lm_head"],
104
+ extra={"llm_int8_threshold": 6.0},
105
+ )
106
+ raise QuantizationError(f"bitsandbytes has no method '{method}'.")
107
+
108
+ def estimate_memory(
109
+ self, model: ModelProfile, config: QuantizationConfig
110
+ ) -> MemoryEstimate:
111
+ return MemoryEstimate(
112
+ weight_gb=memory.estimate_weight_gb(model, config),
113
+ vram_gb=memory.estimate_vram_gb(model, config),
114
+ )
115
+
116
+ def quantize(
117
+ self,
118
+ model: str | Any,
119
+ config: QuantizationConfig,
120
+ *,
121
+ progress: Callable[[str, float], None] | None = None,
122
+ ) -> QuantizationResult:
123
+ err = self.availability_error()
124
+ if err is not None:
125
+ raise err
126
+ if not isinstance(model, str):
127
+ raise QuantizationError(
128
+ "bitsandbytes quantizes at load time.",
129
+ detail="Pass a model id or local path, not an already-loaded module.",
130
+ suggestions=["zeroquantz quantize Qwen/Qwen3-8B --backend bitsandbytes --method nf4"],
131
+ )
132
+
133
+ import torch
134
+ from transformers import AutoModelForCausalLM, BitsAndBytesConfig
135
+
136
+ if progress:
137
+ progress("Building bitsandbytes config", 0.05)
138
+ compute_dtype = getattr(torch, config.compute_dtype, torch.float16)
139
+ if config.method == QuantMethod.BNB_NF4.value:
140
+ bnb_config = BitsAndBytesConfig(
141
+ load_in_4bit=True,
142
+ bnb_4bit_quant_type=config.quant_type or "nf4",
143
+ bnb_4bit_use_double_quant=config.double_quant,
144
+ bnb_4bit_compute_dtype=compute_dtype,
145
+ )
146
+ elif config.method == QuantMethod.BNB_INT8.value:
147
+ bnb_config = BitsAndBytesConfig(
148
+ load_in_8bit=True,
149
+ llm_int8_threshold=float(config.extra.get("llm_int8_threshold", 6.0)),
150
+ llm_int8_skip_modules=config.skip_modules or None,
151
+ )
152
+ else:
153
+ raise QuantizationError(f"bitsandbytes has no method '{config.method}'.")
154
+
155
+ if progress:
156
+ progress("Loading and quantizing weights", 0.1)
157
+ started = time.perf_counter()
158
+ try:
159
+ module = AutoModelForCausalLM.from_pretrained(
160
+ model,
161
+ quantization_config=bnb_config,
162
+ device_map="auto",
163
+ low_cpu_mem_usage=True,
164
+ )
165
+ except Exception as exc: # pragma: no cover - requires GPU + weights
166
+ raise QuantizationError(
167
+ f"bitsandbytes quantization of '{model}' failed.",
168
+ detail=str(exc),
169
+ suggestions=["zeroquantz recommend " + str(model), "check available VRAM"],
170
+ ) from exc
171
+ elapsed = time.perf_counter() - started
172
+ if progress:
173
+ progress("Quantization complete", 1.0)
174
+
175
+ return QuantizationResult(
176
+ backend=self.name,
177
+ method=config.method,
178
+ config=config,
179
+ model=module,
180
+ elapsed_s=elapsed,
181
+ quantized_size_gb=_module_size_gb(module),
182
+ )
183
+
184
+ def export(self, result: QuantizationResult, output_dir: str) -> str:
185
+ if result.model is None:
186
+ raise QuantizationError("Nothing to export: no quantized model in result.")
187
+ result.model.save_pretrained(output_dir)
188
+ result.output_dir = output_dir
189
+ return output_dir
190
+
191
+
192
+ def _module_size_gb(module: Any) -> float | None:
193
+ try:
194
+ import torch
195
+
196
+ total = 0
197
+ for p in module.parameters():
198
+ total += p.numel() * _element_size(p)
199
+ for b in module.buffers():
200
+ total += b.numel() * _element_size(b)
201
+ return round(units.bytes_to_gb(total), 2)
202
+ except Exception: # pragma: no cover - torch dependent
203
+ return None
204
+
205
+
206
+ def _element_size(tensor: Any) -> float:
207
+ try:
208
+ return tensor.element_size()
209
+ except Exception:
210
+ return 4.0
@@ -0,0 +1,198 @@
1
+ """TorchAO backend: INT8 and INT4 weight-only quantization.
2
+
3
+ TorchAO quantizes a *live* module in place via ``torchao.quantization.quantize_``,
4
+ so :meth:`quantize` accepts either a model id/path (which it loads) or an
5
+ already-loaded module.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import importlib.util
11
+ import time
12
+ from typing import TYPE_CHECKING, Any
13
+
14
+ from zeroquantz.core.exceptions import DependencyError, QuantizationError, ZeroQuantzError
15
+ from zeroquantz.quantization import memory
16
+ from zeroquantz.quantization.base import (
17
+ MemoryEstimate,
18
+ QuantizationBackend,
19
+ QuantizationResult,
20
+ )
21
+ from zeroquantz.quantization.config import QuantizationConfig, QuantMethod
22
+ from zeroquantz.utils.logging import get_logger
23
+
24
+ if TYPE_CHECKING:
25
+ from collections.abc import Callable
26
+
27
+ from zeroquantz.hardware.capabilities import HardwareProfile
28
+ from zeroquantz.models.metadata import ModelProfile
29
+ from zeroquantz.optimization.constraints import OptimizationGoal
30
+
31
+ log = get_logger(__name__)
32
+
33
+
34
+ class TorchAOBackend(QuantizationBackend):
35
+ name = "torchao"
36
+ description = "PyTorch-native INT8/INT4 weight-only quantization (torchao)."
37
+ methods = (QuantMethod.TORCHAO_INT8.value, QuantMethod.TORCHAO_INT4.value)
38
+
39
+ def is_available(self) -> bool:
40
+ return all(
41
+ importlib.util.find_spec(m) is not None for m in ("torchao", "torch")
42
+ )
43
+
44
+ def availability_error(self) -> ZeroQuantzError | None:
45
+ if self.is_available():
46
+ return None
47
+ missing = "torchao" if importlib.util.find_spec("torchao") is None else "torch"
48
+ extra = "quant" if missing == "torchao" else "torch"
49
+ return DependencyError.for_extra(missing, extra, purpose="quantize with TorchAO")
50
+
51
+ def supports(
52
+ self,
53
+ model: ModelProfile,
54
+ hardware: HardwareProfile,
55
+ runtime: str | None = None,
56
+ config: QuantizationConfig | None = None,
57
+ ) -> bool:
58
+ if not hardware.cuda_available:
59
+ return False
60
+ method = config.method if config else None
61
+ if method == QuantMethod.TORCHAO_INT4.value:
62
+ return hardware.supports("int4")
63
+ if method == QuantMethod.TORCHAO_INT8.value:
64
+ return hardware.supports("int8")
65
+ return hardware.supports("int8")
66
+
67
+ def default_config(
68
+ self,
69
+ method: str,
70
+ model: ModelProfile,
71
+ hardware: HardwareProfile,
72
+ goal: OptimizationGoal | None = None,
73
+ ) -> QuantizationConfig:
74
+ compute_dtype = "bfloat16" if hardware.supports_bf16 else "float16"
75
+ if method == QuantMethod.TORCHAO_INT4.value:
76
+ return QuantizationConfig(
77
+ backend=self.name,
78
+ method=method,
79
+ weight_bits=4,
80
+ activation_bits=16,
81
+ group_size=128,
82
+ compute_dtype=compute_dtype,
83
+ skip_modules=["lm_head"],
84
+ )
85
+ if method == QuantMethod.TORCHAO_INT8.value:
86
+ return QuantizationConfig(
87
+ backend=self.name,
88
+ method=method,
89
+ weight_bits=8,
90
+ activation_bits=16,
91
+ group_size=None,
92
+ compute_dtype=compute_dtype,
93
+ skip_modules=["lm_head"],
94
+ )
95
+ raise QuantizationError(f"TorchAO has no method '{method}'.")
96
+
97
+ def estimate_memory(
98
+ self, model: ModelProfile, config: QuantizationConfig
99
+ ) -> MemoryEstimate:
100
+ return MemoryEstimate(
101
+ weight_gb=memory.estimate_weight_gb(model, config),
102
+ vram_gb=memory.estimate_vram_gb(model, config),
103
+ )
104
+
105
+ def quantize(
106
+ self,
107
+ model: str | Any,
108
+ config: QuantizationConfig,
109
+ *,
110
+ progress: Callable[[str, float], None] | None = None,
111
+ ) -> QuantizationResult:
112
+ err = self.availability_error()
113
+ if err is not None:
114
+ raise err
115
+
116
+ import torch
117
+
118
+ if progress:
119
+ progress("Loading model", 0.1)
120
+ module = model
121
+ if isinstance(model, str):
122
+ from transformers import AutoModelForCausalLM
123
+
124
+ compute_dtype = getattr(torch, config.compute_dtype, torch.bfloat16)
125
+ module = AutoModelForCausalLM.from_pretrained(
126
+ model, torch_dtype=compute_dtype, device_map="cuda"
127
+ )
128
+
129
+ if progress:
130
+ progress("Applying TorchAO quantization", 0.4)
131
+ started = time.perf_counter()
132
+ try:
133
+ self._apply_quantization(module, config)
134
+ except QuantizationError:
135
+ raise
136
+ except Exception as exc: # pragma: no cover - requires GPU
137
+ raise QuantizationError(
138
+ "TorchAO quantization failed.",
139
+ detail=str(exc),
140
+ suggestions=["update torchao", "zeroquantz recommend " + str(model)],
141
+ ) from exc
142
+ elapsed = time.perf_counter() - started
143
+ if progress:
144
+ progress("Quantization complete", 1.0)
145
+
146
+ return QuantizationResult(
147
+ backend=self.name,
148
+ method=config.method,
149
+ config=config,
150
+ model=module,
151
+ elapsed_s=elapsed,
152
+ )
153
+
154
+ def _apply_quantization(self, module: Any, config: QuantizationConfig) -> None:
155
+ from torchao.quantization import quantize_
156
+
157
+ # Prefer the modern config-object API; fall back to the factory functions
158
+ # used by older torchao releases.
159
+ if config.method == QuantMethod.TORCHAO_INT4.value:
160
+ try:
161
+ from torchao.quantization import Int4WeightOnlyConfig
162
+
163
+ quantize_(module, Int4WeightOnlyConfig(group_size=config.group_size or 128))
164
+ return
165
+ except ImportError:
166
+ from torchao.quantization import int4_weight_only
167
+
168
+ quantize_(module, int4_weight_only(group_size=config.group_size or 128))
169
+ return
170
+ if config.method == QuantMethod.TORCHAO_INT8.value:
171
+ try:
172
+ from torchao.quantization import Int8WeightOnlyConfig
173
+
174
+ quantize_(module, Int8WeightOnlyConfig())
175
+ return
176
+ except ImportError:
177
+ from torchao.quantization import int8_weight_only
178
+
179
+ quantize_(module, int8_weight_only())
180
+ return
181
+ raise QuantizationError(f"TorchAO has no method '{config.method}'.")
182
+
183
+ def export(self, result: QuantizationResult, output_dir: str) -> str:
184
+ if result.model is None:
185
+ raise QuantizationError("Nothing to export: no quantized model in result.")
186
+ # TorchAO models save through the standard HF API when loaded via
187
+ # transformers; otherwise fall back to a state-dict checkpoint.
188
+ try:
189
+ result.model.save_pretrained(output_dir, safe_serialization=False)
190
+ except AttributeError: # pragma: no cover - non-HF module
191
+ import torch
192
+
193
+ from pathlib import Path
194
+
195
+ Path(output_dir).mkdir(parents=True, exist_ok=True)
196
+ torch.save(result.model.state_dict(), Path(output_dir) / "model.pt")
197
+ result.output_dir = output_dir
198
+ return output_dir