"""Weight-only symmetric per-output-channel INT8 linear. Scales stay float32 across dtype casts. ``module.to(dtype=torch.bfloat16)`` (and ``.bfloat16()`` / ``.half()`` / ``.to(device, dtype)``) must not touch them; device moves still do. The int8 weight codes are likewise dtype-stable. """ from __future__ import annotations import torch import torch.nn.functional as F from torch import nn # Leaf names of Linear modules whose 2-D weights are quantized. # Exact match: the routers are `gate` / `image_gate` / `audio_gate`, NOT `gate_proj`. QUANT_LEAVES = frozenset( {"query_key_value", "dense", "gate_proj", "up_proj", "down_proj"} ) QUANT_RULE = ( "Quantize ONLY 2-D .weight tensors under model.model.layers. whose owning " "module's leaf name is exactly one of query_key_value, dense, gate_proj, " "up_proj, down_proj. Everything else stays byte-identical BF16: embeddings, " "lm_head, all norms, the vision tower, linear_proj, and the three routers " "(modules named gate, image_gate, audio_gate — leaf match, not a substring). " "Per-output-channel symmetric: scale = absmax/127, " "q = clamp(round(w/scale), -127, 127). All-zero rows: scale 1.0, q 0." ) # Real checkpoint keys look like `model.model.layers.N...`. A bare # `layers.N...` name is the same stack with the root prefix omitted (tests). _DECODER_LAYER_PREFIXES = ((), ("model", "model")) def _weight_leaf(tensor_name: str) -> str | None: """Owning module's leaf name if `tensor_name` ends in `.weight`, else None.""" if not isinstance(tensor_name, str) or not tensor_name.endswith(".weight"): return None module = tensor_name[: -len(".weight")] if not module: return None return module.rsplit(".", 1)[-1] def _under_decoder_layers(tensor_name: str) -> bool: """True when the tensor lives under the MLLM decoder `model.model.layers` stack. `layers` must be its own path component, followed by a layer index. The components before it must be empty or end in `model.model` — so a vision tower that happens to contain the substring "layers" is not selected, and `gate` is never selected just because `gate_proj` contains those letters. """ parts = tensor_name.split(".") for i, part in enumerate(parts): if part != "layers": continue if i + 1 >= len(parts) or not parts[i + 1].isdigit(): continue prefix = tuple(parts[:i]) if prefix in _DECODER_LAYER_PREFIXES: return True if len(prefix) >= 2 and prefix[-2:] == ("model", "model"): return True return False def quant_rule_leaf(tensor_name: str) -> str | None: """Leaf name if the name matches the quantize rule, ignoring rank. Returns None when the tensor is not a candidate. A candidate whose rank is not 2 is a hard error for the stream (see quantize_stream); ``is_quantizable`` itself returns False for that case. """ leaf = _weight_leaf(tensor_name) if leaf not in QUANT_LEAVES: return None if not _under_decoder_layers(tensor_name): return None return leaf def is_quantizable(tensor_name: str, shape) -> bool: """True only for 2-D quantize-rule weights. See ``QUANT_RULE``.""" if quant_rule_leaf(tensor_name) is None: return False try: rank = len(shape) except TypeError: return False return rank == 2 def quantize_weight(weight: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """Per-output-channel symmetric int8. ``scale = absmax(row) / 127``, ``q = clamp(round(w / scale), -127, 127)``. An all-zero row gets scale 1.0 and q 0 (no div-by-zero, no NaN/Inf). """ if weight.ndim != 2: raise ValueError( f"quantize_weight expects a 2-D weight, got shape {tuple(weight.shape)}" ) wf = weight.detach().to(dtype=torch.float32) absmax = wf.abs().amax(dim=1) scale = absmax / 127.0 zero = scale == 0 # All-zero rows would divide by 0. Force scale 1 and q 0 instead of NaN. scale = torch.where(zero, torch.ones_like(scale), scale) q = torch.round(wf / scale[:, None]).clamp(-127, 127).to(dtype=torch.int8) q = torch.where(zero[:, None], torch.zeros_like(q), q) return q.contiguous(), scale.to(dtype=torch.float32).contiguous() def _scale_name(weight_name: str) -> str: if not weight_name.endswith(".weight"): raise ValueError(f"not a weight tensor name: {weight_name}") return weight_name[: -len("weight")] + "scale" class Int8Linear(nn.Module): """``F.linear`` on a weight dequantized from int8 + per-row float32 scale. ``weight`` is int8 ``[out, in]``, ``scale`` is float32 ``[out]``, ``bias`` (optional) keeps the source dtype. All three are buffers. """ def __init__(self, weight: torch.Tensor, scale: torch.Tensor, bias: torch.Tensor | None): super().__init__() if weight.dtype != torch.int8 or weight.ndim != 2: raise ValueError( f"weight must be int8 [out, in], got dtype={weight.dtype} shape={tuple(weight.shape)}" ) if scale.dtype != torch.float32 or tuple(scale.shape) != (weight.shape[0],): raise ValueError( f"scale must be float32 [{weight.shape[0]}], got dtype={scale.dtype} shape={tuple(scale.shape)}" ) if bias is not None: if bias.ndim != 1 or bias.shape[0] != weight.shape[0]: raise ValueError( f"bias must be [{weight.shape[0]}], got shape={tuple(bias.shape)}" ) self.in_features = int(weight.shape[1]) self.out_features = int(weight.shape[0]) self.register_buffer("weight", weight) self.register_buffer("scale", scale) self.register_buffer("bias", bias) def _apply(self, fn, *args, **kwargs): # Pull dtype-stable buffers out before Module._apply. Putting them back # with only a device move (never fn's dtype cast) keeps scale float32 # and weight int8. Bias is left in the dict so it follows the cast. saved: dict[str, torch.Tensor] = {} for name in ("weight", "scale"): buf = self._buffers.get(name, None) if buf is not None: saved[name] = buf self._buffers[name] = None try: out = super()._apply(fn, *args, **kwargs) finally: for name, buf in saved.items(): self._buffers[name] = _move_device_keep_dtype(buf, fn) return out def forward(self, x: torch.Tensor) -> torch.Tensor: # One dequant in fp32, one cast to the activation dtype, then linear. w = (self.weight.float() * self.scale[:, None]).to(dtype=x.dtype) return F.linear(x, w, self.bias) @classmethod def from_linear(cls, linear: nn.Linear) -> "Int8Linear": if not isinstance(linear, nn.Linear): raise TypeError(f"from_linear expects nn.Linear, got {type(linear).__name__}") q, scale = quantize_weight(linear.weight.data) if linear.bias is None: bias = None else: bias = linear.bias.detach().clone() return cls(q, scale, bias) @classmethod def shell( cls, in_features: int, out_features: int, bias: bool, bias_dtype: torch.dtype, device, ) -> "Int8Linear": """Empty buffers (for ``meta``). Does not read or write any weight values.""" dev = torch.device(device) if not isinstance(device, torch.device) else device weight = torch.empty((out_features, in_features), dtype=torch.int8, device=dev) scale = torch.empty((out_features,), dtype=torch.float32, device=dev) if bias: bias_t: torch.Tensor | None = torch.empty( (out_features,), dtype=bias_dtype, device=dev ) else: bias_t = None return cls(weight, scale, bias_t) def extra_repr(self) -> str: return ( f"in_features={self.in_features}, out_features={self.out_features}, " f"bias={self.bias is not None}" ) def _move_device_keep_dtype(buf: torch.Tensor, fn) -> torch.Tensor: """Apply only the device change implied by ``fn``, preserving ``buf``'s dtype and values. Probed with a 0-element tensor so a dtype cast cannot round the real scale. """ try: probe = torch.empty((), dtype=buf.dtype, device=buf.device) moved = fn(probe) except Exception: return buf if not torch.is_tensor(moved) or moved.device == buf.device: return buf return buf.to(device=moved.device)