kingjones777's picture
Add files using upload-large-folder tool
18c1466 verified
Raw History Blame Contribute Delete
8.72 kB
"""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)