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 Download src/bigru_t/quantization/w8a8_smoothquant.py from PowerMachine/BiGRU_T_version: direct link, hf CLI and curl.
- Browser
- Download file 16.3 kB
-
https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/quantization/w8a8_smoothquant.py
- Command line
-
hf download hf://PowerMachine/BiGRU_T_version/src/bigru_t/quantization/w8a8_smoothquant.py
-
curl -L -o w8a8_smoothquant.py https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/quantization/w8a8_smoothquant.py
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. | |
| """ | |
| def forward(ctx, x, scale, qmin, qmax): | |
| x_q = torch.round(x / scale).clamp(qmin, qmax) | |
| x_deq = x_q * scale | |
| return x_deq | |
| 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)) | |
| def alpha(self) -> torch.Tensor: | |
| """α atual = sigmoid(eta_alpha) ∈ (0, 1).""" | |
| return torch.sigmoid(self.eta_alpha) | |
| 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) | |
| 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) | |
| 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", | |
| ] | |