"""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", ]