File size: 8,722 Bytes
18c1466
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
"""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)