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_error_reduction.py from PowerMachine/BiGRU_T_version: direct link, hf CLI and curl.
- Browser
- Download file 18.2 kB
-
https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/quantization/w8a8_error_reduction.py
- Command line
-
hf download hf://PowerMachine/BiGRU_T_version/src/bigru_t/quantization/w8a8_error_reduction.py
-
curl -L -o w8a8_error_reduction.py https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/quantization/w8a8_error_reduction.py
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 | |
| # ============================================================================ | |
| 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)) | |
| 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", | |
| ] | |