File size: 5,683 Bytes
9118991
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Symmetric group-wise RTN fake quantization used by LoopQ.

The paper defines the scalar quantizer in Eq. (4) as

    Q(x; c) = c * clip(round(x / c), q_min, q_max).

It specifies group-wise W4A4/W4A8 quantization with group size 32, but does
not specify several low-level choices.  Those choices are deliberately kept
small and explicit in ``configs/loopq_compatibility_decisions.yaml``.

This module performs QDQ/fake quantization.  Integer packing and kernels are
outside LQ1's scope.
"""

from __future__ import annotations

from dataclasses import dataclass

import torch


SUPPORTED_BITS = frozenset({4, 8})


@dataclass(frozen=True)
class GroupwiseQuantizationResult:
    """Fake-quantized tensor and the group parameters used to produce it."""

    dequantized: torch.Tensor
    integers: torch.Tensor
    scales: torch.Tensor
    bits: int
    group_size: int
    axis: int
    qmin: int
    qmax: int


def _round_half_away_from_zero(x: torch.Tensor) -> torch.Tensor:
    """Deterministic RTN with half-way values rounded away from zero."""

    return torch.copysign(torch.floor(torch.abs(x) + 0.5), x)


def _canonical_axis(ndim: int, axis: int) -> int:
    if ndim == 0:
        raise ValueError("group-wise quantization requires a non-scalar tensor")
    if not -ndim <= axis < ndim:
        raise IndexError(f"axis {axis} is out of bounds for tensor rank {ndim}")
    return axis % ndim


def fake_quantize_groupwise(
    tensor: torch.Tensor,
    *,
    bits: int,
    group_size: int = 32,
    axis: int = -1,
    scales: torch.Tensor | None = None,
    rounding_ste: bool = False,
) -> GroupwiseQuantizationResult:
    """Apply symmetric signed group-wise RTN quantization and dequantization.

    Groups are consecutive chunks on ``axis``.  A final short group is
    calibrated independently and never padded into the returned tensors.
    Unless explicit learned ``scales`` are supplied, each group uses
    ``max(abs(group)) / (2**(bits - 1) - 1)``.  All-zero groups use scale 1,
    producing exact zeros without NaNs.
    """

    if bits not in SUPPORTED_BITS:
        raise ValueError(f"bits must be one of {sorted(SUPPORTED_BITS)}, got {bits}")
    if group_size <= 0:
        raise ValueError(f"group_size must be positive, got {group_size}")
    if not tensor.is_floating_point():
        raise TypeError("tensor must have a floating-point dtype")

    canonical_axis = _canonical_axis(tensor.ndim, axis)
    moved = tensor.movedim(canonical_axis, -1)
    width = moved.shape[-1]
    if width == 0:
        raise ValueError("quantized axis must not be empty")

    qmax = 2 ** (bits - 1) - 1
    # Match the cited FlatQuant symmetric RTN implementation: the signed
    # two's-complement range includes its most-negative code.
    qmin = -(2 ** (bits - 1))
    group_count = (width + group_size - 1) // group_size
    expected_scale_shape = moved.shape[:-1] + (group_count,)

    work = moved.to(torch.float32)
    # Batch groups into one tensor; padding is excluded from returned values
    # and zero padding cannot change an absmax, including the short tail.
    padding = group_count * group_size - width
    grouped = torch.nn.functional.pad(work, (0, padding)).reshape(
        *moved.shape[:-1], group_count, group_size
    )
    if scales is not None:
        group_scales = scales.to(device=tensor.device, dtype=torch.float32)
        if tuple(group_scales.shape) != tuple(expected_scale_shape):
            raise ValueError(
                f"scales shape must be {tuple(expected_scale_shape)}, "
                f"got {tuple(group_scales.shape)}"
            )
        if not torch.isfinite(group_scales).all() or (group_scales <= 0).any():
            raise ValueError("all scales must be finite and strictly positive")
        normalized = grouped / group_scales.unsqueeze(-1)
    else:
        absmax = grouped.abs().amax(dim=-1)
        group_scales = torch.where(absmax == 0, torch.ones_like(absmax), absmax / qmax)
        safe_absmax = torch.where(absmax == 0, torch.ones_like(absmax), absmax)
        normalized = (grouped / safe_absmax.unsqueeze(-1)) * qmax
    rounded = _round_half_away_from_zero(normalized)
    if rounding_ste:
        # Exact rounded forward; unit rounding derivative before clipping.
        # This also differentiates x/scale, giving the scale gradient its
        # quantization-residual term instead of just the integer code.
        rounded = rounded.detach() + (normalized - normalized.detach())
    integer = rounded.clamp(qmin, qmax)
    integers = integer.flatten(-2)[..., :width].movedim(-1, canonical_axis).to(torch.int8)
    dequantized = (integer * group_scales.unsqueeze(-1)).flatten(-2)[..., :width].movedim(-1, canonical_axis)
    dequantized = dequantized.to(dtype=tensor.dtype)
    return GroupwiseQuantizationResult(
        dequantized=dequantized,
        integers=integers,
        scales=group_scales,
        bits=bits,
        group_size=group_size,
        axis=canonical_axis,
        qmin=qmin,
        qmax=qmax,
    )


def quantize_weight(weight: torch.Tensor, *, bits: int = 4) -> GroupwiseQuantizationResult:
    """LoopQ paper-contract weight QDQ: symmetric RTN, group size 32."""

    if bits != 4:
        raise ValueError("LoopQ's evaluated weight precision is 4 bits")
    return fake_quantize_groupwise(weight, bits=bits, group_size=32, axis=-1)


def quantize_activation(
    activation: torch.Tensor,
    *,
    bits: int,
    scales: torch.Tensor | None = None,
) -> GroupwiseQuantizationResult:
    """LoopQ paper-contract activation QDQ for A4 or A8."""

    return fake_quantize_groupwise(
        activation, bits=bits, group_size=32, axis=-1, scales=scales
    )