File size: 7,371 Bytes
4e316d6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
NormalFloat4 (NF4) Block-Wise Weight Quantization.

NF4 is the quantization format from the QLoRA paper. Instead of spacing
quantization levels evenly (like INT8 absmax), NF4 places its 16 levels
at quantiles of the standard normal distribution. Since neural network
weights are approximately normal, this puts more levels near zero (where
most weights cluster) and fewer in the tails — capturing more information
per bit than uniform spacing.

Block-wise quantization computes one absmax scale per block of 64 weights,
finer than per-channel, which reduces quantization error. Each weight is
normalized by its block's absmax, then snapped to the nearest NF4 level.

Two 4-bit indices are packed into one uint8 byte, so the packed tensor is
half the size of the original element count. Combined with 4-bit precision,
NF4 uses ~4.5 bits per weight (4 bits + amortized scale overhead), roughly
a 4x compression vs float16.
"""

import math

import torch
import torch.nn as nn

# 16 quantization levels from the QLoRA paper, computed as quantiles of N(0,1)
# mapped to [-1, 1]. Asymmetric: 8 non-positive (including 0) + 8 positive.
NF4_LEVELS = [
    -1.0,
    -0.6961928009986877,
    -0.5250730514526367,
    -0.39491748809814453,
    -0.28444138169288635,
    -0.18477343022823334,
    -0.09105003625154495,
    0.0,
    0.07958029955625534,
    0.16093020141124725,
    0.24611230194568634,
    0.33791524171829224,
    0.44070982933044434,
    0.5626170039176941,
    0.7229568362236023,
    1.0,
]


def _snap_to_nf4(normalized):
    """Snap each value in normalized tensor (in [-1, 1]) to the nearest NF4 level index."""
    levels = torch.tensor(NF4_LEVELS, dtype=normalized.dtype, device=normalized.device)
    # Shape: (*normalized.shape, 16) — distance to each level
    distances = (normalized.unsqueeze(-1) - levels).abs()
    return distances.argmin(dim=-1).to(torch.uint8)


def _pack_uint4(indices):
    """Pack pairs of 4-bit indices into uint8. Input length must be even."""
    # Cast to int32 before bitwise ops — MPS does not support uint8 bit-shifts.
    indices = indices.reshape(-1).to(torch.int32)
    high = (indices[0::2] << 4) & 0xFF
    low = indices[1::2] & 0xFF
    return (high | low).to(torch.uint8)


def _unpack_uint4(packed, num_elements):
    """Unpack uint8 into 4-bit indices."""
    # Cast to int32 before bitwise ops — MPS does not support uint8 bit-shifts.
    packed_i32 = packed.to(torch.int32)
    high = (packed_i32 >> 4) & 0x0F
    low = packed_i32 & 0x0F
    # Interleave high and low
    unpacked = torch.stack([high, low], dim=-1).reshape(-1).to(torch.uint8)
    return unpacked[:num_elements]


def nf4_quantize(tensor, block_size=64):
    """
    Quantize a float tensor to NF4 using block-wise absmax scaling.

    Args:
        tensor: Float tensor (typically a weight matrix)
        block_size: Number of elements per quantization block

    Returns:
        (packed_uint8, scales)
        - packed_uint8: Two 4-bit NF4 indices packed per byte
        - scales: One absmax scale per block
    """
    if block_size % 2 != 0:
        raise ValueError("block_size must be even for uint4 packing")

    original_shape = tensor.shape
    flat = tensor.reshape(-1).float()
    n = flat.numel()

    # Pad to multiple of block_size
    pad = (block_size - n % block_size) % block_size
    if pad > 0:
        flat = torch.cat([flat, torch.zeros(pad, device=flat.device)])

    # Reshape into blocks
    blocks = flat.reshape(-1, block_size)

    # One absmax scale per block
    scales = blocks.abs().amax(dim=1)
    scales = torch.clamp(scales, min=1e-8)

    # Normalize each block to [-1, 1]
    normalized = blocks / scales.unsqueeze(1)

    # Snap to nearest NF4 level
    indices = _snap_to_nf4(normalized)

    # Pack pairs of 4-bit indices into uint8
    packed = _pack_uint4(indices)

    return packed, scales


def nf4_dequantize(packed, scales, original_shape, block_size=64):
    """
    Reconstruct float tensor from NF4 packed data + scales.

    Args:
        packed: uint8 tensor with two 4-bit NF4 indices per byte
        scales: One absmax scale per block
        original_shape: Original tensor shape before quantization
        block_size: Block size used during quantization

    Returns:
        Reconstructed float tensor with original_shape
    """
    n = math.prod(original_shape)

    # Pad count to match what was used during quantization
    pad = (block_size - n % block_size) % block_size
    total = n + pad

    # Unpack indices
    indices = _unpack_uint4(packed, total)

    # Look up NF4 levels
    levels = torch.tensor(NF4_LEVELS, dtype=torch.float32, device=packed.device)
    values = levels[indices.long()]

    # Reshape into blocks, multiply by scales
    blocks = values.reshape(-1, block_size)
    dequantized = blocks * scales.unsqueeze(1)

    # Flatten and trim padding
    return dequantized.reshape(-1)[:n].reshape(original_shape)


class NF4Linear(nn.Module):
    """Drop-in replacement for nn.Linear that stores NF4-packed weights and dequantizes on forward."""

    def __init__(self, in_features, out_features, packed_weight, scales, weight_shape, block_size, bias=None):
        super().__init__()
        self.in_features = in_features
        self.out_features = out_features
        self.register_buffer("_weight_shape", torch.tensor(list(weight_shape), dtype=torch.long))
        self.register_buffer("_block_size", torch.tensor(block_size, dtype=torch.long))
        self.register_buffer("packed_weight", packed_weight)
        self.register_buffer("scales", scales)
        if bias is not None:
            self.register_buffer("bias", bias)
        else:
            self.bias = None

    @classmethod
    def from_linear(cls, linear, block_size=64):
        """Create an NF4Linear from an existing nn.Linear."""
        weight = linear.weight.data.float()
        packed, scales = nf4_quantize(weight, block_size=block_size)
        bias = linear.bias.data if linear.bias is not None else None
        return cls(linear.in_features, linear.out_features, packed, scales, weight.shape, block_size, bias)

    def forward(self, x):
        w = nf4_dequantize(self.packed_weight, self.scales, tuple(self._weight_shape.tolist()), self._block_size.item()).to(x.dtype)
        return nn.functional.linear(x, w, self.bias)


def quantize_model_nf4(model, block_size=64):
    """Walk the model and replace all nn.Linear layers with NF4Linear.

    Skips Linear layers whose weights are tied to an Embedding (e.g. a
    shared output head), since quantizing them would break the tie and
    produce a mismatch between embedding lookups and output projections.
    """
    # Collect data_ptrs of all Embedding weights so we can detect ties
    tied_ptrs = {
        m.weight.data_ptr()
        for m in model.modules()
        if isinstance(m, nn.Embedding)
    }
    _quantize_nf4_recursive(model, block_size, tied_ptrs)


def _quantize_nf4_recursive(module, block_size, tied_ptrs=frozenset()):
    for name, child in module.named_children():
        if isinstance(child, nn.Linear):
            if child.weight.data_ptr() in tied_ptrs:
                continue  # skip weight-tied layers
            setattr(module, name, NF4Linear.from_linear(child, block_size=block_size))
        else:
            _quantize_nf4_recursive(child, block_size, tied_ptrs)