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..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