Qwen3.5-9B-SpeedX9-GDN32 / quantization.py
summerMC's picture
Upload Qwen3.5 UNI MAX 9B checkpoint with benchmark results
efc38b5 verified
Raw History Blame Contribute Delete
7.9 kB
from __future__ import annotations
"""Runtime-only quantization policies for UNI MAX.
UNI MAX v1.1 distinguishes two very different notions:
* ``full`` quantization (MAX TURBO compatible): every dense MLP projection plus
the LM head. This is fast, but it changes the hidden trajectory that writes
future GDN recurrent states. A cache copied from a BF16-prefill model is
therefore *not* generally a valid state for this quantized dynamical system.
* ``state_safe`` quantization (UNI MAX): only the LM head and the MLP of the
final decoder layer. These operators execute strictly after the final
persistent GDN state update for each token. Hence they cannot change any
conv/recurrent cache tensor. If their greedy token agrees with the exact
model, the next-token recurrent state remains exactly synchronized.
Quantization is inference-only and is never written back into the checkpoint.
"""
from dataclasses import dataclass
from typing import Callable
import torch
from torch import nn
@dataclass(frozen=True)
class QuantizationReport:
mode: str
applied: bool
reason: str | None
targeted_linears: int
cuda_capability: tuple[int, int] | None
policy: str = "full"
def _cuda_capability() -> tuple[int, int] | None:
if not torch.cuda.is_available():
return None
return tuple(int(x) for x in torch.cuda.get_device_capability())
def _is_dense_mlp_linear(module: nn.Module, fqn: str) -> bool:
if not isinstance(module, nn.Linear):
return False
return ".mlp." in fqn and fqn.rsplit(".", 1)[-1] in {"gate_proj", "up_proj", "down_proj"}
def _target_dense_decode_linears(module: nn.Module, fqn: str) -> bool:
"""MAX-compatible full dense target set; native GDN projections excluded."""
if not isinstance(module, nn.Linear):
return False
return fqn == "lm_head" or _is_dense_mlp_linear(module, fqn)
def _final_layer_index(model: nn.Module) -> int:
cfg = getattr(model, "config", None)
n = getattr(cfg, "num_hidden_layers", None)
if n is not None:
return int(n) - 1
layers = getattr(getattr(model, "model", None), "layers", None)
if layers is None:
raise ValueError("cannot determine final decoder layer index")
return len(layers) - 1
def make_state_safe_filter(model: nn.Module) -> Callable[[nn.Module, str], bool]:
"""Return the causal-state-compatible UNI quantization filter.
The eligible set is exactly:
* ``lm_head``
* ``model.layers.<last>.mlp.{gate_proj,up_proj,down_proj}``
No module whose output can influence a persistent recurrent-state write is
eligible.
"""
last = _final_layer_index(model)
prefix = f"model.layers.{last}.mlp."
def _filter(module: nn.Module, fqn: str) -> bool:
if not isinstance(module, nn.Linear):
return False
if fqn == "lm_head":
return True
return fqn.startswith(prefix) and fqn.rsplit(".", 1)[-1] in {
"gate_proj",
"up_proj",
"down_proj",
}
return _filter
def count_target_linears(model: nn.Module) -> int:
return sum(1 for fqn, mod in model.named_modules() if _target_dense_decode_linears(mod, fqn))
def count_state_safe_linears(model: nn.Module) -> int:
filt = make_state_safe_filter(model)
return sum(1 for fqn, mod in model.named_modules() if filt(mod, fqn))
def state_safe_target_names(model: nn.Module) -> tuple[str, ...]:
filt = make_state_safe_filter(model)
return tuple(fqn for fqn, mod in model.named_modules() if filt(mod, fqn))
def available_quant_modes() -> tuple[str, ...]:
modes = ["exact"]
try:
import torchao # noqa: F401
modes.extend(["fp8", "int8"])
except Exception:
pass
return tuple(modes)
def _apply_quantization(model: nn.Module, mode: str, *, policy: str) -> QuantizationReport:
mode = str(mode).lower().strip()
policy = str(policy).lower().strip()
capability = _cuda_capability()
if policy == "state_safe":
filter_fn = make_state_safe_filter(model)
targeted = count_state_safe_linears(model)
elif policy == "full":
filter_fn = _target_dense_decode_linears
targeted = count_target_linears(model)
else:
return QuantizationReport(mode, False, f"unknown quantization policy: {policy}", 0, capability, policy)
if mode in {"", "none", "exact", "bf16"}:
return QuantizationReport("exact", False, None, targeted, capability, policy)
try:
from torchao.quantization import quantize_
except Exception as exc: # pragma: no cover - runtime dependent
return QuantizationReport(
mode,
False,
f"torchao unavailable: {type(exc).__name__}: {exc}",
targeted,
capability,
policy,
)
if targeted == 0:
return QuantizationReport(mode, False, "no eligible dense linear layers found", 0, capability, policy)
try:
if mode == "fp8":
if capability is None or capability < (8, 9):
return QuantizationReport(
mode,
False,
f"FP8 requires CUDA SM 8.9+, got {capability}",
targeted,
capability,
policy,
)
from torchao.quantization import Float8DynamicActivationFloat8WeightConfig, PerTensor
config = Float8DynamicActivationFloat8WeightConfig(granularity=PerTensor())
elif mode == "int8":
from torchao.quantization import Int8WeightOnlyConfig
config = Int8WeightOnlyConfig()
else:
return QuantizationReport(
mode,
False,
f"unknown quantization mode: {mode}",
targeted,
capability,
policy,
)
quantize_(model, config, filter_fn=filter_fn)
return QuantizationReport(mode, True, None, targeted, capability, policy)
except Exception as exc: # pragma: no cover - hardware/runtime specific
return QuantizationReport(
mode,
False,
f"{type(exc).__name__}: {exc}",
targeted,
capability,
policy,
)
def apply_max_quantization(model: nn.Module, mode: str) -> QuantizationReport:
"""Legacy MAX TURBO policy: all MLP projections + LM head."""
return _apply_quantization(model, mode, policy="full")
def apply_uni_quantization(model: nn.Module, mode: str) -> QuantizationReport:
"""UNI MAX state-compatible policy: final MLP + LM head only."""
return _apply_quantization(model, mode, policy="state_safe")
@torch.inference_mode()
def greedy_token_trace(model, input_ids: torch.Tensor, steps: int = 32) -> torch.Tensor:
steps = int(steps)
if steps < 1:
return torch.empty((input_ids.shape[0], 0), dtype=torch.long, device=input_ids.device)
cache = model.make_recurrent_cache()
out = model(input_ids=input_ids, past_key_values=cache, use_cache=True, logits_to_keep=1)
token = out.logits[:, -1, :].argmax(dim=-1, keepdim=True)
tokens = [token]
for _ in range(steps - 1):
token = model.greedy_step(token, cache)
tokens.append(token)
return torch.cat(tokens, dim=1)
def token_agreement(reference: torch.Tensor, candidate: torch.Tensor) -> tuple[float, int]:
if reference.shape != candidate.shape:
raise ValueError(f"token trace shape mismatch: {reference.shape} vs {candidate.shape}")
if reference.numel() == 0:
return 1.0, 0
eq = reference.eq(candidate)
agreement = float(eq.float().mean().item())
flat = eq.reshape(-1).tolist()
prefix = 0
for ok in flat:
if not ok:
break
prefix += 1
return agreement, prefix