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.
- zeroquantz/__init__.py +14 -0
- zeroquantz/__main__.py +8 -0
- zeroquantz/agent/__init__.py +16 -0
- zeroquantz/agent/dispatcher.py +520 -0
- zeroquantz/agent/intents.py +46 -0
- zeroquantz/agent/parser.py +255 -0
- zeroquantz/benchmark/__init__.py +7 -0
- zeroquantz/benchmark/latency.py +66 -0
- zeroquantz/benchmark/memory.py +41 -0
- zeroquantz/benchmark/quality.py +38 -0
- zeroquantz/benchmark/runner.py +151 -0
- zeroquantz/cli/__init__.py +7 -0
- zeroquantz/cli/app.py +98 -0
- zeroquantz/cli/commands.py +459 -0
- zeroquantz/cli/interactive.py +56 -0
- zeroquantz/core/__init__.py +7 -0
- zeroquantz/core/artifacts.py +179 -0
- zeroquantz/core/context.py +127 -0
- zeroquantz/core/events.py +30 -0
- zeroquantz/core/exceptions.py +105 -0
- zeroquantz/core/session.py +202 -0
- zeroquantz/core/subenv.py +202 -0
- zeroquantz/deploy/__init__.py +25 -0
- zeroquantz/deploy/assets.py +161 -0
- zeroquantz/deploy/launcher.py +80 -0
- zeroquantz/deploy/runtime_env.py +66 -0
- zeroquantz/deploy/targets.py +154 -0
- zeroquantz/export/__init__.py +8 -0
- zeroquantz/export/exporter.py +68 -0
- zeroquantz/export/report.py +203 -0
- zeroquantz/hardware/__init__.py +15 -0
- zeroquantz/hardware/capabilities.py +152 -0
- zeroquantz/hardware/detector.py +200 -0
- zeroquantz/hardware/gpu.py +31 -0
- zeroquantz/models/__init__.py +8 -0
- zeroquantz/models/architecture.py +168 -0
- zeroquantz/models/downloader.py +161 -0
- zeroquantz/models/hf_auth.py +105 -0
- zeroquantz/models/inspector.py +249 -0
- zeroquantz/models/metadata.py +108 -0
- zeroquantz/models/search.py +71 -0
- zeroquantz/optimization/__init__.py +22 -0
- zeroquantz/optimization/candidate.py +272 -0
- zeroquantz/optimization/constraints.py +70 -0
- zeroquantz/optimization/fit.py +203 -0
- zeroquantz/optimization/pareto.py +66 -0
- zeroquantz/optimization/planner.py +297 -0
- zeroquantz/optimization/recommender.py +149 -0
- zeroquantz/profiling/__init__.py +18 -0
- zeroquantz/profiling/calibration.py +74 -0
- zeroquantz/profiling/sensitivity.py +234 -0
- zeroquantz/quantization/__init__.py +17 -0
- zeroquantz/quantization/backends/__init__.py +8 -0
- zeroquantz/quantization/backends/bitsandbytes.py +210 -0
- zeroquantz/quantization/backends/torchao.py +198 -0
- zeroquantz/quantization/base.py +136 -0
- zeroquantz/quantization/catalog.py +321 -0
- zeroquantz/quantization/config.py +106 -0
- zeroquantz/quantization/gguf_pipeline.py +210 -0
- zeroquantz/quantization/isolated.py +248 -0
- zeroquantz/quantization/memory.py +133 -0
- zeroquantz/quantization/native.py +91 -0
- zeroquantz/quantization/registry.py +101 -0
- zeroquantz/render.py +341 -0
- zeroquantz/runtimes/__init__.py +18 -0
- zeroquantz/runtimes/base.py +64 -0
- zeroquantz/runtimes/compatibility.py +91 -0
- zeroquantz/runtimes/registry.py +70 -0
- zeroquantz/runtimes/transformers.py +53 -0
- zeroquantz/runtimes/vllm.py +83 -0
- zeroquantz/tui/__init__.py +13 -0
- zeroquantz/tui/app.py +77 -0
- zeroquantz/tui/banner.py +47 -0
- zeroquantz/tui/screens/__init__.py +25 -0
- zeroquantz/tui/screens/confirm.py +41 -0
- zeroquantz/tui/screens/execute.py +194 -0
- zeroquantz/tui/screens/model_select.py +206 -0
- zeroquantz/tui/screens/plan.py +177 -0
- zeroquantz/tui/screens/quantize_select.py +272 -0
- zeroquantz/tui/screens/settings.py +219 -0
- zeroquantz/tui/screens/token.py +94 -0
- zeroquantz/tui/screens/welcome.py +128 -0
- zeroquantz/tui/screens/workspace.py +175 -0
- zeroquantz/tui/styles/app.tcss +424 -0
- zeroquantz/tui/widgets/__init__.py +9 -0
- zeroquantz/tui/widgets/chip.py +36 -0
- zeroquantz/tui/widgets/sidebar.py +107 -0
- zeroquantz/tui/widgets/status_bar.py +43 -0
- zeroquantz/utils/__init__.py +8 -0
- zeroquantz/utils/config.py +46 -0
- zeroquantz/utils/env.py +78 -0
- zeroquantz/utils/logging.py +73 -0
- zeroquantz/utils/metrics.py +98 -0
- zeroquantz/utils/paths.py +57 -0
- zeroquantz/utils/units.py +134 -0
- zeroquantz/verification/__init__.py +17 -0
- zeroquantz/verification/logits.py +55 -0
- zeroquantz/verification/report.py +186 -0
- zeroquantz/verification/weights.py +44 -0
- zeroquantz/version.py +8 -0
- zeroquantz-0.1.0.dist-info/METADATA +72 -0
- zeroquantz-0.1.0.dist-info/RECORD +105 -0
- zeroquantz-0.1.0.dist-info/WHEEL +4 -0
- zeroquantz-0.1.0.dist-info/entry_points.txt +2 -0
- 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
|