BiGRU_T_version / src /bigru_t /quantization /w8a8_smoothquant.py
PowerMachine's picture
V5: W8A8 aprimorado (α aprendível via STE) + bug fix apply_w8a8_to_model + 4 hiperparâmetros auto-ajustáveis (T, τ, λ_ent, init_gate) + OOM-Killer fixes + num_layers_hyp=8 fixo
cbf308e verified
Raw History Blame Contribute Delete
16.3 kB
"""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",
]