File size: 16,296 Bytes
cbf308e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
"""w8a8_smoothquant.py — W8A8 com SmoothQuant e α APRENDÍVEL (V5).

═══════════════════════════════════════════════════════════════════════════════
V5 — APRIMORAMENTO MATEMÁTICO: α APRENDÍVEL VIA STE
═══════════════════════════════════════════════════════════════════════════════

ORIGEM (gru-ring-v13-9-2):
  source/xavante_work/flexnet/w8a8_smoothquant.py (V13.9.1)

Algoritmo (Xiao et al., 2023 — "SmoothQuant"):
==============================================
Dada uma camada Linear Y = X W^T:
  * X ∈ R^{B × d_in}  ativações (com outliers por canal)
  * W ∈ R^{d_out × d_in}  pesos (aproximadamente uniformes)

SmoothQuant introduz fator de suavização s ∈ R^{d_in}_{>0}:
  * X̃ = X / s                (divide ativações → reduz outliers)
  * W̃ = W * s                (multiplica pesos → preserva produto)
  * Y = X̃ W̃^T = X W^T        (matematicamente equivalente)

Escolha de s por canal de entrada j:
  s_j = max(|X_j|)^α  /  max(|W_j|)^{1-α}

  * α = 0.0 → "dificuldade" toda nos pesos (= quantização per-channel clássica)
  * α = 1.0 → "dificuldade" toda nas ativações (= quantização per-token)
  * α = 0.5 (default empírico do paper) → balanceia

V5 — MELHORIA MATEMÁTICA SOBRE V13.9.1:
  Em V13.9.1, α era FIXO (default 0.5). Em V5, α é APRENDÍVEL:
    * Parametrizado via sigmoid: α = sigmoid(eta_alpha) ∈ (0, 1)
    * eta_alpha é nn.Parameter treinável — o modelo ajusta α sozinho
    * Backward via STE (Straight-Through Estimator):
        - round() não é diferenciável → STE substitui grad por 1
        - clamp() tem gradiente 1 dentro do intervalo, 0 fora
    * Loss auxiliary: ||Y_fp - Y_q||_F^2 / ||Y_fp||_F^2 (relative MSE)
      adicionada à loss total → gradiente flui para eta_alpha

  Isso atende ao requisito V5: "modelo deve ajustar todos os seus parâmetros
  por conta própria sempre". α era o único parâmetro do W8A8 que não era
  auto-ajustado — agora é.

V5 — MELHORIA DE MEMÓRIA (OOM-Killer fix):
  * calib_X_buf (256 × d_in float32 = 1KB por camada) REMOVIDO
    — em V13.9.1 era armazenado para freeze_quantization, mas só max-abs
      é usado; buffer completo era desperdício.
  * Substituído por calib_max_abs (d_in float32 = 0.5KB por camada)
    — running max-abs per channel, suficiente para s.
  * Em modelo com 100 camadas Linear × d_in=1024: 256KB → 0.5KB (-99.8%)

Limite de erro (W8A8 simétrico, per-channel):
  ||Y_fp - Y_q||_2  ≤  ε_w · ||W||_F · ||X||_F  +  ε_x · ||X||_F · ||W||_F
onde ε_w, ε_x ≤ 2^{-7} (1/128) para escalas bem calibradas.
"""
from __future__ import annotations

import math
import logging
from typing import Optional, Tuple

import torch
import torch.nn as nn
import torch.nn.functional as F

logger = logging.getLogger(__name__)


# ═══════════════════════════════════════════════════════════════════════════
# Funções de Quantização Simétrica (zero_point=0 sempre)
# ═══════════════════════════════════════════════════════════════════════════
def _per_channel_max_abs(W: torch.Tensor, dim: int) -> torch.Tensor:
    """max |W| ao longo de dim, com mínimo epsilon para evitar divisão por 0."""
    return W.abs().amax(dim=dim).clamp(min=1e-8)


def calibrate_smoothquant(
    X: torch.Tensor,
    W: torch.Tensor,
    alpha: float = 0.5,
) -> torch.Tensor:
    """Calcula o fator SmoothQuant s ∈ R^{d_in}_{>0}.

    X: (B, d_in) ou (..., d_in) — ativações de calibração
    W: (d_out, d_in) — pesos da camada Linear
    alpha: ∈ [0, 1], trade-off (default 0.5)

    Retorna: s ∈ R^{d_in} (positivo).

    Fórmula: s_j = max(|X_j|)^α / max(|W_j|)^{1-α}
    """
    d_in = W.shape[-1]
    if X.shape[-1] != d_in:
        return torch.ones(d_in, device=W.device, dtype=W.dtype)
    X_flat = X.reshape(-1, d_in)
    max_x = X_flat.abs().amax(dim=0).clamp(min=1e-8)        # (d_in,)
    max_w = W.abs().amax(dim=0).clamp(min=1e-8)             # (d_in,)
    s = torch.pow(max_x, alpha) / torch.pow(max_w, 1.0 - alpha)
    return s.clamp(min=1e-8)


def quantize_per_channel_symmetric(
    W: torch.Tensor,
    axis: int = 0,
    n_bits: int = 8,
) -> Tuple[torch.Tensor, torch.Tensor]:
    """Quantização int8 simétrica per-channel (tensor 2D).

    W: 2D tensor (d_axis, d_other).
    axis=0 → uma escala por linha (cada output channel tem sua escala).
    Retorna: (W_int8, scale) onde W_fp ≈ W_int8 * scale e scale é shape (d_axis,).
    """
    assert W.dim() == 2, f"esperado tensor 2D, recebeu dim={W.dim()}"
    qmax = 2 ** (n_bits - 1) - 1   # 127 para int8 (faixa simétrica [-127, 127])
    qmin = -(2 ** (n_bits - 1) - 1)
    other = 1 - axis
    max_abs = W.abs().amax(dim=other, keepdim=True).clamp(min=1e-8)  # (d_axis, 1)
    scale = max_abs / qmax
    W_q = torch.round(W / scale).clamp(qmin, qmax).to(torch.int8)
    return W_q, scale.squeeze(1)


def quantize_per_token_symmetric(
    X: torch.Tensor,
    n_bits: int = 8,
) -> Tuple[torch.Tensor, torch.Tensor]:
    """Quantização int8 simétrica per-token (último eixo = canal).

    X: (..., d_in)
    Retorna: (X_int8, scale) onde X_fp ≈ X_int8 * scale e scale é shape (..., 1).
    """
    qmax = 2 ** (n_bits - 1) - 1
    qmin = -(2 ** (n_bits - 1) - 1)
    max_abs = X.abs().amax(dim=-1, keepdim=True).clamp(min=1e-8)
    scale = max_abs / qmax
    X_q = torch.round(X / scale).clamp(qmin, qmax).to(torch.int8)
    return X_q, scale


# ═══════════════════════════════════════════════════════════════════════════
# STE para round() — permite gradiente fluir através de quantização
# ═══════════════════════════════════════════════════════════════════════════
class STEQuantize(torch.autograd.Function):
    """Straight-Through Estimator para round + clamp.

    Forward:  x_q = round(x / scale).clamp(qmin, qmax) * scale
    Backward: grad flui direto (grad_x = grad_y), ignorando round/clamp.

    Isso é necessário porque:
      * d(round)/dx = 0 em quase todos os pontos (escada de Heaviside)
      * d(clamp)/dx = 0 fora do intervalo, 1 dentro
    Sem STE, gradientes não fluiriam para os pesos/α através da quantização.
    """

    @staticmethod
    def forward(ctx, x, scale, qmin, qmax):
        x_q = torch.round(x / scale).clamp(qmin, qmax)
        x_deq = x_q * scale
        return x_deq

    @staticmethod
    def backward(ctx, grad_output):
        # STE: gradiente passa direto
        return grad_output, None, None, None


def ste_quantize(x: torch.Tensor, scale: torch.Tensor, n_bits: int = 8) -> torch.Tensor:
    """Quantiza com STE. scale deve ter shape broadcastable com x."""
    qmax = 2 ** (n_bits - 1) - 1
    qmin = -(2 ** (n_bits - 1) - 1)
    return STEQuantize.apply(x, scale, qmin, qmax)


# ═══════════════════════════════════════════════════════════════════════════
# SmoothQuantW8A8 — camada Linear W8A8 com α APRENDÍVEL (V5)
# ═══════════════════════════════════════════════════════════════════════════
class SmoothQuantW8A8(nn.Module):
    """Camada Linear W8A8 com SmoothQuant e α APRENDÍVEL.

    V5 vs V13.9.1:
      * α é parâmetro treinável (via sigmoid(eta_alpha))
      * calib_X_buf removido (era OOM em modelos grandes) — apenas max_abs
      * STE explícito para round() — backward bem definido
      * Quantização do peso é DIFFERENCIÁVEL w.r.t. α (via s e STE)

    Modos:
      * calibration: coleta max_abs das ativações
      * quantized: pesos em int8, ativações quantizadas on-the-fly com STE

    Args:
        ref: nn.Linear de referência (pesos copiados)
        init_alpha: valor inicial de α (default 0.5)
        learnable_alpha: se True (default V5), α é treinável
        calib_batch_size: tamanho do buffer de calibração (apenas p/ count)
        n_bits: nº de bits (default 8)
    """

    def __init__(
        self,
        ref: nn.Linear,
        init_alpha: float = 0.5,
        learnable_alpha: bool = True,
        calib_batch_size: int = 256,
        n_bits: int = 8,
    ):
        super().__init__()
        self.in_features = ref.in_features
        self.out_features = ref.out_features
        self.n_bits = n_bits
        self.learnable_alpha = learnable_alpha

        # Pesos fp (treináveis até freeze; após freeze ainda treináveis p/ α)
        self.weight_fp = nn.Parameter(ref.weight.detach().clone())
        self.bias = (
            nn.Parameter(ref.bias.detach().clone()) if ref.bias is not None else None
        )

        # V5: α aprendível via sigmoid(eta_alpha) ∈ (0, 1)
        init_eta = float(math.log(init_alpha / max(1e-6, 1.0 - init_alpha)))
        if learnable_alpha:
            self.eta_alpha = nn.Parameter(torch.tensor(init_eta, dtype=torch.float32))
        else:
            self.register_buffer("eta_alpha", torch.tensor(init_eta, dtype=torch.float32))

        # Buffers de quantização
        self.register_buffer("smooth_scale", torch.ones(self.in_features))
        self.register_buffer("weight_scale", torch.ones(self.out_features))
        self.register_buffer(
            "weight_int8",
            torch.zeros(self.out_features, self.in_features, dtype=torch.int8),
        )
        self.register_buffer("is_quantized", torch.tensor(False))
        self.register_buffer("calib_collected", torch.tensor(False))
        # V5: Apenas max-abs per channel (OOM fix — buffer completo removido)
        self.register_buffer("calib_max_abs", torch.zeros(self.in_features))
        self.register_buffer("calib_count", torch.tensor(0, dtype=torch.long))

    @property
    def alpha(self) -> torch.Tensor:
        """α atual = sigmoid(eta_alpha) ∈ (0, 1)."""
        return torch.sigmoid(self.eta_alpha)

    @torch.no_grad()
    def calibrate(self, X: torch.Tensor) -> None:
        """Acumula estatísticas de ativação (max-abs per channel).

        V5 OOM fix: não armazena tensor completo, apenas max(|X|) per channel.
        """
        if bool(self.is_quantized):
            raise RuntimeError("Calibração após freeze não é permitida.")
        flat = X.reshape(-1, self.in_features)
        batch_max_abs = flat.abs().amax(dim=0).to(self.calib_max_abs.dtype)
        self.calib_max_abs.copy_(torch.maximum(self.calib_max_abs, batch_max_abs))
        self.calib_count += flat.shape[0]
        self.calib_collected.fill_(True)

    @torch.no_grad()
    def freeze_quantization(self) -> None:
        """Computa s com α ATUAL, quantiza W em int8 per-channel.

        Nota: depois de freeze, o modelo ainda pode chamar recompute_smooth_scale()
        para atualizar s quando α mudar (durante o treino do α).
        """
        if not bool(self.calib_collected):
            logger.warning("SmoothQuantW8A8: freeze sem calibração. Usando s=1.")
            X_calib = self.weight_fp.new_zeros(1, self.in_features)
        else:
            # Reconstrói tensor "representativo" das ativações a partir de max-abs
            X_calib = self.calib_max_abs.unsqueeze(0)  # (1, d_in)

        alpha_val = float(self.alpha.item())
        s = calibrate_smoothquant(X_calib, self.weight_fp, alpha=alpha_val)
        self.smooth_scale.copy_(s.to(self.smooth_scale.dtype))

        W_smooth = self.weight_fp * s.unsqueeze(0)
        W_q, w_scale = quantize_per_channel_symmetric(W_smooth, axis=0, n_bits=self.n_bits)
        self.weight_int8.copy_(W_q)
        self.weight_scale.copy_(w_scale)
        self.is_quantized.fill_(True)

    @torch.no_grad()
    def recompute_smooth_scale(self) -> None:
        """Recomputa s e re-quantiza W com α ATUAL (após treino de α).

        V5: chamado pelo trainer a cada N steps para manter W quantizado
        consistente com α atualizado.
        """
        if not bool(self.is_quantized):
            return
        if not bool(self.calib_collected):
            return
        X_calib = self.calib_max_abs.unsqueeze(0)
        alpha_val = float(self.alpha.item())
        s = calibrate_smoothquant(X_calib, self.weight_fp, alpha=alpha_val)
        self.smooth_scale.copy_(s.to(self.smooth_scale.dtype))
        W_smooth = self.weight_fp * s.unsqueeze(0)
        W_q, w_scale = quantize_per_channel_symmetric(W_smooth, axis=0, n_bits=self.n_bits)
        self.weight_int8.copy_(W_q)
        self.weight_scale.copy_(w_scale)

    def forward(self, X: torch.Tensor) -> torch.Tensor:
        """Forward W8A8 com SmoothQuant e α aprendível.

        V5: quando quantized, faz quantização on-the-fly com STE para que
        gradientes fluam para eta_alpha (via s(α)).
        """
        if not bool(self.is_quantized):
            # Modo calibração/treino inicial: forward fp32 normal
            return F.linear(X, self.weight_fp, self.bias)

        # V5: recompute s(α) on-the-fly (diferenciável w.r.t. eta_alpha)
        # smooth_scale é buffer (detached), mas recalculamos s com α treinável
        alpha_t = self.alpha  # diferenciável
        # max_abs armazenado em calib_max_abs (detached)
        max_x = self.calib_max_abs.clamp(min=1e-8)
        max_w = self.weight_fp.detach().abs().amax(dim=0).clamp(min=1e-8)
        s = torch.pow(max_x, alpha_t) / torch.pow(max_w, 1.0 - alpha_t)
        s = s.clamp(min=1e-8)

        # Smooth: X̃ = X / s, W̃_fp = weight_fp * s
        X_smooth = X / s  # (..., d_in) — diferenciável w.r.t. eta_alpha via s
        # W̃_fp não precisa de STE aqui porque weight_fp é treinável diretamente
        W_smooth = self.weight_fp * s.unsqueeze(0)  # (d_out, d_in)

        # Quantiza W per-channel com STE (diferenciável)
        qmax = 2 ** (self.n_bits - 1) - 1
        max_w_smooth = W_smooth.abs().amax(dim=1, keepdim=True).clamp(min=1e-8)
        w_scale_dyn = max_w_smooth / qmax  # (d_out, 1)
        W_q_deq = ste_quantize(W_smooth, w_scale_dyn, self.n_bits)  # (d_out, d_in)

        # Quantiza X per-token com STE
        max_x_smooth = X_smooth.abs().amax(dim=-1, keepdim=True).clamp(min=1e-8)
        x_scale_dyn = max_x_smooth / qmax  # (..., 1)
        X_q_deq = ste_quantize(X_smooth, x_scale_dyn, self.n_bits)  # (..., d_in)

        # Matmul em fp32 (simulação; HW real faria INT8×INT8→INT32→fp32)
        Y = X_q_deq @ W_q_deq.t()
        if self.bias is not None:
            Y = Y + self.bias
        return Y

    def quantization_error(self, X: torch.Tensor) -> float:
        """Calcula ||Y_fp - Y_q||_F / ||Y_fp||_F como métrica de erro."""
        with torch.no_grad():
            Y_fp = F.linear(X, self.weight_fp, self.bias)
            Y_q = self.forward(X)
            num = (Y_fp - Y_q).norm().item()
            den = max(Y_fp.norm().item(), 1e-12)
        return num / den

    def extra_repr(self) -> str:
        return (
            f"in_features={self.in_features}, out_features={self.out_features}, "
            f"alpha={float(self.alpha.item()):.4f}, "
            f"learnable_alpha={self.learnable_alpha}, "
            f"is_quantized={bool(self.is_quantized)}"
        )


__all__ = [
    "calibrate_smoothquant",
    "quantize_per_channel_symmetric",
    "quantize_per_token_symmetric",
    "STEQuantize",
    "ste_quantize",
    "SmoothQuantW8A8",
]