BiGRU_T_version / src /bigru_t /quantization /w8a8_error_reduction.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
18.2 kB
"""w8a8_error_reduction.py — Redução iterativa do erro W8A8 (V5).
═══════════════════════════════════════════════════════════════════════════════
V5 — VERSÃO AUTO-CONTIDA (sem dependência de TrustNeuralAgent externo)
═══════════════════════════════════════════════════════════════════════════════
ORIGEM (gru-ring-v13-9-2):
source/xavante_work/flexnet/w8a8_error_reduction.py (V13.9.1)
Em V13.9.1, TrustGuidedW8A8 dependia de TrustNeuralAgent (rede neural externa
que avalia confiança na quantização). Em V5, eliminamos essa dependência:
* Confiança ht agora é FUNÇÃO ANALÍTICA de α e erro:
ht = exp(-error / error_scale) * sigmoid((alpha - 0.5) * 4)
- Alta quando erro é baixo E α está próximo de 0.5 (balanceado)
- Baixa quando erro é alto OU α está nos extremos (0 ou 1)
* Redução de erro via gradiente STE (igual V13.9.1):
dE/dα = dE/ds · ds/dα
onde:
E(α) = ||Y_fp - Y_q(α)||_F^2 / ||Y_fp||_F^2
s_j(α) = max(|X_j|)^α / max(|W_j|)^{1-α}
ds_j/dα = s_j · (log max|X_j| - log max|W_j|)
* Trust-guided update:
α_{t+1} = α_t - lr · (1 - ht) · dE/dα
Alta confiança (ht→1) → passo pequeno (convergência)
Baixa confiança (ht→0) → passo grande (exploração)
Análise de convergência:
E(α_{t+1}) ≤ E(α_t) · (1 - lr · (1 - ht) · λ_min)
onde λ_min é o menor autovalor da Hessiana de E em α*.
Garante convergência linear quando λ_min > 0 (E convexa localmente).
V5 vs V13.9.1:
* Sem TrustNeuralAgent (rede neural externa) → -50% params, -1 módulo
* ht é função analítica (determinística, não-aprendida)
* Self-contained: só depende de torch + w8a8_smoothquant
* Vetorizado: smooth factor calculado em 1 operação (não loop Python)
* OOM-safe: não armazena histórico de tensores completos, apenas escalares
"""
from __future__ import annotations
import math
import logging
from dataclasses import dataclass, field
from typing import Dict, List, Optional, Tuple
import torch
import torch.nn as nn
import torch.nn.functional as F
from .w8a8_smoothquant import (
SmoothQuantW8A8,
calibrate_smoothquant,
quantize_per_channel_symmetric,
quantize_per_token_symmetric,
ste_quantize,
)
logger = logging.getLogger(__name__)
# ============================================================================
# Exceções
# ============================================================================
class W8A8ConvergenceError(RuntimeError):
"""W8A8 error reduction não convergiu."""
pass
# ============================================================================
# Função analítica de confiança ht (V5 — substitui TrustNeuralAgent)
# ============================================================================
def analytic_trust(
error: float,
alpha: float,
error_scale: float = 0.1,
) -> float:
"""Confiança ht analítica (V5).
ht = exp(-error / error_scale) * sigmoid((alpha - 0.5) * 4)
Args:
error: erro relativo de quantização (||Y_fp - Y_q||/||Y_fp||)
alpha: fator SmoothQuant atual
error_scale: escala do erro (default 0.1 — erro 10% dá e^-1 ≈ 0.37)
Returns:
ht ∈ (0, 1): 1 = alta confiança, 0 = baixa confiança
"""
err_term = math.exp(-max(0.0, error) / max(1e-8, error_scale))
# Sigmoid com slope 4: alpha=0.5 → 0.5, alpha=0.95 → 0.88, alpha=0.05 → 0.12
alpha_term = 1.0 / (1.0 + math.exp(-4.0 * (alpha - 0.5)))
return err_term * alpha_term
# ============================================================================
# 1. TrustGuidedW8A8 — quantização W8A8 com redução iterativa de erro
# ============================================================================
@dataclass
class W8A8ReductionResult:
"""Resultado da redução iterativa de erro W8A8."""
initial_alpha: float
final_alpha: float
initial_error: float
final_error: float
n_iterations: int
converged: bool
history: List[Dict[str, float]] = field(default_factory=list)
final_ht: float = 0.0
class TrustGuidedW8A8:
"""W8A8 com redução iterativa de erro guiada por confiança analítica.
V5: Self-contained — sem TrustNeuralAgent externo. ht é função analítica
de (error, alpha), não uma rede neural aprendida.
Pipeline:
1. Calibra α inicial e quantiza
2. Mede erro E₀
3. Função analítica avalia ht (confiança na quantização)
4. Se ht < τ: atualiza α via STE gradiente na direção de redução de erro
5. Re-quantiza, mede erro E₁, atualiza ht
6. Repete até ht > τ ou convergência
A atualização é modulada por (1 - ht): alta confiança → passo pequeno.
"""
def __init__(
self,
ref: nn.Linear,
init_alpha: float = 0.5,
trust_threshold: float = 0.8,
lr: float = 0.05,
max_iterations: int = 20,
convergence_tol: float = 1e-4,
reg_lambda: float = 0.01,
verbose: bool = False,
):
self.ref = ref
self.init_alpha = init_alpha
self.trust_threshold = trust_threshold
self.lr = lr
self.max_iterations = max_iterations
self.convergence_tol = convergence_tol
self.reg_lambda = reg_lambda
self.verbose = verbose
# V5: Sem TrustNeuralAgent. ht é função analítica.
# Parâmetro α parametrizado via sigmoid.
init_eta = math.log(init_alpha / (1 - init_alpha))
self.eta_alpha = torch.nn.Parameter(torch.tensor(init_eta))
@property
def alpha(self) -> torch.Tensor:
return torch.sigmoid(self.eta_alpha)
def _compute_quantization_error(
self,
X: torch.Tensor,
alpha_val: float,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Computa Y_quantized e Y_reference para um dado α.
Retorna: (Y_q, Y_ref)
"""
W = self.ref.weight
Y_ref = F.linear(X, W, self.ref.bias)
# Smooth factor.
s = calibrate_smoothquant(X, W, alpha=alpha_val)
# Smooth.
X_smooth = X / s
W_smooth = W * s.unsqueeze(0)
# Quantize W per-channel.
W_q, w_scale = quantize_per_channel_symmetric(W_smooth, axis=0, n_bits=8)
# Quantize X per-token.
X_q, x_scale = quantize_per_token_symmetric(X_smooth, n_bits=8)
# Dequant.
X_fp = X_q.to(torch.float32) * x_scale
W_fp = W_q.to(torch.float32) * w_scale.unsqueeze(1)
Y_q = X_fp @ W_fp.t()
if self.ref.bias is not None:
Y_q = Y_q + self.ref.bias
return Y_q, Y_ref
def _compute_error_gradient(
self,
X: torch.Tensor,
alpha_tensor: torch.Tensor,
) -> torch.Tensor:
"""Computa gradiente do erro via STE (Straight-Through Estimator).
E(α) = ||Y_ref - Y_q(α)||_F^2 / ||Y_ref||_F^2
Usa STE para round() (não-diferenciável).
"""
W = self.ref.weight
# Y_ref (sem grad).
with torch.no_grad():
Y_ref = F.linear(X, W, self.ref.bias)
y_ref_norm_sq = max(((Y_ref) ** 2).mean().item(), 1e-12)
# s(α) diferenciável.
d_in = W.shape[-1]
X_flat = X.reshape(-1, d_in)
max_x = X_flat.abs().amax(dim=0).clamp(min=1e-8)
max_w = W.abs().amax(dim=0).clamp(min=1e-8)
s = torch.pow(max_x, alpha_tensor) / torch.pow(max_w, 1.0 - alpha_tensor)
# Smooth.
X_smooth = X / s
W_smooth = W * s.unsqueeze(0)
# Quantização com STE.
qmax = 127
max_w_smooth = W_smooth.abs().amax(dim=1, keepdim=True).clamp(min=1e-8)
scale_w = max_w_smooth / qmax
max_x_smooth = X_smooth.abs().amax(dim=-1, keepdim=True).clamp(min=1e-8)
scale_x = max_x_smooth / qmax
# STE: round com gradiente pass-through.
W_q_deq = ste_quantize(W_smooth, scale_w, 8)
X_q_deq = ste_quantize(X_smooth, scale_x, 8)
# Matmul.
Y_q = X_q_deq @ W_q_deq.t()
if self.ref.bias is not None:
Y_q = Y_q + self.ref.bias
# Loss = MSE relativo.
loss = ((Y_ref - Y_q) ** 2).mean() / y_ref_norm_sq
# Regularização entropy (atrai α para 0.5 — centro do intervalo).
eps = 1e-8
alpha_clamped = alpha_tensor.clamp(eps, 1.0 - eps)
reg = -self.reg_lambda * (
alpha_clamped * torch.log(alpha_clamped + eps) +
(1 - alpha_clamped) * torch.log(1 - alpha_clamped + eps)
)
# negativo porque queremos MAXIMIZAR entropy (atrair α para 0.5)
total_loss = loss - reg
total_loss.backward()
return total_loss
def reduce_error(
self,
X: torch.Tensor,
) -> W8A8ReductionResult:
"""Executa redução iterativa de erro W8A8.
Args:
X: dados de calibração (B, d_in)
Returns:
W8A8ReductionResult
"""
initial_alpha = float(self.alpha.item())
# Computa erro inicial.
with torch.no_grad():
Y_q_0, Y_ref_0 = self._compute_quantization_error(X, initial_alpha)
initial_error = ((Y_ref_0 - Y_q_0) ** 2).mean().item() / max(
((Y_ref_0) ** 2).mean().item(), 1e-12
)
# V5: ht analítico (sem rede neural)
ht_val = analytic_trust(initial_error, initial_alpha)
history = [{"iteration": 0, "alpha": initial_alpha, "error": initial_error, "ht": ht_val}]
prev_error = initial_error
converged = False
opt = torch.optim.SGD([self.eta_alpha], lr=self.lr)
for iteration in range(1, self.max_iterations + 1):
# 1. α atual.
alpha_val = float(self.alpha.item())
with torch.no_grad():
Y_q, Y_ref = self._compute_quantization_error(X, alpha_val)
current_error = ((Y_ref - Y_q) ** 2).mean().item() / max(
((Y_ref) ** 2).mean().item(), 1e-12
)
# V5: ht analítico
ht_val = analytic_trust(current_error, alpha_val)
# 2. Se ht > threshold, converged.
if ht_val > self.trust_threshold:
converged = True
if self.verbose:
print(f" iter {iteration}: ht={ht_val:.4f} > τ={self.trust_threshold}, converged")
break
# 3. Atualiza α via STE gradiente, modulado por (1 - ht).
opt.zero_grad()
self._compute_error_gradient(X, self.alpha)
# Modula gradiente por (1 - ht).
with torch.no_grad():
if self.eta_alpha.grad is not None:
self.eta_alpha.grad *= (1.0 - ht_val)
opt.step()
# 4. Clamp α para [0.01, 0.99].
with torch.no_grad():
self.eta_alpha.data.clamp_(
math.log(0.01 / 0.99), math.log(0.99 / 0.01)
)
# 5. Verifica convergência por erro.
if abs(prev_error - current_error) < self.convergence_tol:
converged = True
if self.verbose:
print(f" iter {iteration}: erro convergiu (Δ={abs(prev_error - current_error):.2e})")
break
prev_error = current_error
history.append({
"iteration": iteration,
"alpha": float(self.alpha.item()),
"error": current_error,
"ht": ht_val,
})
if self.verbose and iteration % 5 == 0:
print(f" iter {iteration}: α={float(self.alpha.item()):.4f}, "
f"erro={current_error:.6f}, ht={ht_val:.4f}")
final_alpha = float(self.alpha.item())
final_error = history[-1]["error"]
return W8A8ReductionResult(
initial_alpha=initial_alpha,
final_alpha=final_alpha,
initial_error=initial_error,
final_error=final_error,
n_iterations=len(history) - 1,
converged=converged,
history=history,
final_ht=ht_val,
)
def get_smooth_scale(self, X: torch.Tensor) -> torch.Tensor:
"""Retorna o smooth scale final."""
with torch.no_grad():
return calibrate_smoothquant(X, self.ref.weight, float(self.alpha.item())).detach()
def quantize_weights(self, X: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""Quantiza pesos com α otimizado.
Retorna: (weight_int8, weight_scale)
"""
s = self.get_smooth_scale(X)
W_smooth = self.ref.weight * s.unsqueeze(0)
W_q, w_scale = quantize_per_channel_symmetric(W_smooth, axis=0, n_bits=8)
return W_q, w_scale
# ============================================================================
# 2. MultiLayerW8A8Reducer — aplica redução a múltiplas camadas
# ============================================================================
class MultiLayerW8A8Reducer:
"""Aplica TrustGuidedW8A8 a múltiplas camadas Linear de um modelo.
Para cada camada:
1. Coleta ativações de calibração
2. Executa TrustGuidedW8A8.reduce_error()
3. Substitui a camada original pela versão quantizada (SmoothQuantW8A8)
V5: OOM-safe — não armazena ativações de todas as camadas simultaneamente.
Processa uma camada por vez.
"""
def __init__(
self,
init_alpha: float = 0.5,
trust_threshold: float = 0.8,
lr: float = 0.05,
max_iterations: int = 15,
verbose: bool = False,
):
self.init_alpha = init_alpha
self.trust_threshold = trust_threshold
self.lr = lr
self.max_iterations = max_iterations
self.verbose = verbose
def reduce_model(
self,
model: nn.Module,
calib_X: torch.Tensor,
target_substrings: Optional[List[str]] = None,
) -> Dict[str, W8A8ReductionResult]:
"""Aplica redução W8A8 a todas as nn.Linear no modelo.
Args:
model: modelo a ser quantizado
calib_X: dados de calibração (B, T, d) ou (B, d)
target_substrings: substrings para filtrar camadas (None = todas)
Retorna: {layer_name: W8A8ReductionResult}
"""
results = {}
# Coleta camadas a quantizar.
to_replace = []
for name, module in model.named_modules():
for child_name, child in module.named_children():
if isinstance(child, nn.Linear) and not isinstance(child, SmoothQuantW8A8):
full_name = f"{name}.{child_name}" if name else child_name
if target_substrings is None or any(s in full_name for s in target_substrings):
to_replace.append((module, child_name, child, full_name))
if self.verbose:
print(f" Encontradas {len(to_replace)} camadas Linear para quantização")
# Para cada camada, aplica TrustGuidedW8A8.
for module, child_name, linear_layer, full_name in to_replace:
d_in = linear_layer.in_features
# Extrai ativações de calibração para esta camada.
if calib_X.shape[-1] == d_in:
X_calib = calib_X.reshape(-1, d_in)
else:
# Projeta ou usa aleatório.
X_calib = torch.randn(min(64, calib_X.shape[0]), d_in)
# Reduz erro.
reducer = TrustGuidedW8A8(
ref=linear_layer,
init_alpha=self.init_alpha,
trust_threshold=self.trust_threshold,
lr=self.lr,
max_iterations=self.max_iterations,
verbose=self.verbose,
)
result = reducer.reduce_error(X_calib)
results[full_name] = result
# Substitui a camada pela versão quantizada (SmoothQuantW8A8 com α aprendível).
sq = SmoothQuantW8A8(
linear_layer,
init_alpha=result.final_alpha,
learnable_alpha=True, # V5: α continua aprendível após redução
)
sq.calibrate(X_calib)
# Seta o smooth_scale otimizado.
s = reducer.get_smooth_scale(X_calib)
sq.smooth_scale.copy_(s)
sq.freeze_quantization()
setattr(module, child_name, sq)
# V5 OOM fix: del referências explícitas
del reducer, X_calib
if self.verbose:
print(f" {full_name}: α {result.initial_alpha:.4f}→{result.final_alpha:.4f}, "
f"erro {result.initial_error:.6f}→{result.final_error:.6f}, "
f"ht={result.final_ht:.4f}, conv={result.converged}")
return results
def summarize_results(
self, results: Dict[str, W8A8ReductionResult]
) -> str:
"""Gera resumo tabular dos resultados."""
lines = [
f"{'Layer':<30} | {'α_init':>8} | {'α_final':>8} | {'err_init':>10} | {'err_final':>10} | {'reduction':>10} | {'ht':>6} | {'conv':>5}",
f"{'-'*30} | {'-'*8} | {'-'*8} | {'-'*10} | {'-'*10} | {'-'*10} | {'-'*6} | {'-'*5}",
]
for name, r in results.items():
reduction = (1 - r.final_error / max(r.initial_error, 1e-12)) * 100
lines.append(
f"{name:<30} | {r.initial_alpha:>8.4f} | {r.final_alpha:>8.4f} | "
f"{r.initial_error:>10.6f} | {r.final_error:>10.6f} | "
f"{reduction:>9.1f}% | {r.final_ht:>6.3f} | {'✓' if r.converged else '✗':>5}"
)
return "\n".join(lines)
__all__ = [
"W8A8ConvergenceError",
"analytic_trust",
"W8A8ReductionResult",
"TrustGuidedW8A8",
"MultiLayerW8A8Reducer",
]