"""quantization.py — W8A8 Quantization via SmoothQuant para CNN-BiGRU. Implementa quantização weight-only-8-bit + activation-8-bit (W8A8) usando a técnica SmoothQuant (Xiao et al., 2023) para migrar a variância das ativações para os pesos, reduzindo a perda de precisão. ============================================================================== ANÁLISE MATEMÁTICA E LÓGICA — SmoothQuant + W8A8 ============================================================================== PROBLEMA: Em modelos LLM, as ativações têm outliers em alguns canais que tornam a quantização INT8 difícil. Se quantizarmos diretamente, os outliers saturam os outros canais, levando a grandes erros. SOLUÇÃO SmoothQuant: Seja Y = X * W, onde X ∈ R^{B×T×d_in} e W ∈ R^{d_in×d_out}. 1. Para cada canal de entrada i, computa o máximo absoluto: s_i = max|X_i| / max|W_i| (em batch, suavizado por alpha) Ou mais precisamente: s_i = (max|X_i|^alpha) / (max|W_i|^(1-alpha)) com alpha ∈ [0, 1] tipicamente 0.5. 2. Migra a variância: dividir X por s e multiplicar W por s: X' = X / s (ativações suavizadas — outliers reduzidos) W' = W * s (pesos absorvem a escala) Como Y = X * W = (X/s) * (s*W) = X' * W', a operação matricial é matematicamente equivalente. 3. Quantiza ambos para INT8 com escala por tensor ou por canal: X_q = round(X' / scale_x) * scale_x W_q = round(W' / scale_w) * scale_w 4. Y ≈ X_q * W_q (com erro de quantização reduzido) W8A8: - W: pesos em INT8 (8-bit weights) - A: ativações em INT8 (8-bit activations) - Redução de memória: ~4x (FP32 -> INT8) - Speedup: 2-4x em hardware com suporte INT8 (AMX, AVX512-VNNI) INTEGRAÇÃO COM CNN-BiGRU: - Aplicado após o treino (post-training quantization) - Aplicável a: Linear (atenção, FFN, classificador), Conv1d - Não aplicado a: Embedding (mantém FP32 para precisão) - Em CPU sem AMX, ainda economiza memória (sem speedup significativo) ============================================================================== USO ============================================================================== from cnn_bigru.utils.quantization import ( SmoothQuantizer, W8A8Config, quantize_model_w8a8 ) # Após treino: quantizer = SmoothQuantizer(W8A8Config(alpha=0.5, n_calibration_batches=5)) quantized_model = quantizer.quantize(model, calibration_dataloader) Autor: CNN-BiGRU Project """ from __future__ import annotations import logging import math from dataclasses import dataclass, field from typing import Dict, List, Optional, Tuple, Union import torch import torch.nn as nn import torch.nn.functional as F logger = logging.getLogger(__name__) # ============================================================================ # Configuração # ============================================================================ @dataclass class W8A8Config: """Configuração da quantização W8A8 com SmoothQuant.""" # Alpha de suavização (0=só pesos, 1=só ativações, 0.5=balanceado) alpha: float = 0.5 # Número de batches para calibração n_calibration_batches: int = 5 # Tipo de escala: "per_tensor" ou "per_channel" scale_type: str = "per_channel" # Manter权重 em FP32 para quantização dinâmica (default: False = INT8 estático) dynamic_weight: bool = False # Quantizar embeddings (default: False) quantize_embeddings: bool = False # Quantizar Conv1d (default: True) quantize_conv1d: bool = True # Quantizar LayerNorm (default: False — sensível) quantize_layernorm: bool = False # Clip threshold para outliers (em desvios-padrão; None = sem clip) outlier_clip_std: Optional[float] = 4.0 # Device para calibração device: str = "cpu" # Simular quantização (mantém FP32 mas aplica ruído de quantização) simulate: bool = False # ============================================================================ # SmoothQuant Calibrator # ============================================================================ class SmoothQuantCalibrator: """Coleta estatísticas de ativações e pesos para SmoothQuant. Per-corre o modelo com dados de calibração e registra max|X| e max|W| para cada camada Linear/Conv1d. """ def __init__(self, config: W8A8Config): self.config = config self.stats: Dict[str, Dict[str, torch.Tensor]] = {} def _hook_factory(self, name: str): """Cria um hook forward para coletar max|X|.""" def hook(module, input, output): # input é uma tupla; input[0] é o tensor principal if not isinstance(input, tuple) or len(input) == 0: return x = input[0] if not isinstance(x, torch.Tensor): return # Reduzir para [d_in] (assumindo última dim = features) with torch.no_grad(): if x.dim() >= 2: # max sobre todas as dims exceto a última x_flat = x.reshape(-1, x.size(-1)) max_x = x_flat.abs().max(dim=0).values # [d_in] else: max_x = x.abs() # [d_in] # Acumular max (não é média — pegamos o máximo global) if name not in self.stats: self.stats[name] = {"max_x": max_x.clone()} else: self.stats[name]["max_x"] = torch.maximum( self.stats[name]["max_x"], max_x ) return hook def calibrate( self, model: nn.Module, dataloader, n_batches: Optional[int] = None, ) -> Dict[str, Dict[str, torch.Tensor]]: """Coleta estatísticas via forward hooks. Args: model: modelo a calibrar dataloader: iterador de batches n_batches: número de batches (default: config.n_calibration_batches) Returns: dict {layer_name: {"max_x": [d_in], "max_w": [d_out]}} """ n_batches = n_batches or self.config.n_calibration_batches device = self.config.device # Registrar hooks em todas as camadas Linear e Conv1d hooks = [] layer_modules = [] for name, module in model.named_modules(): if isinstance(module, (nn.Linear, nn.Conv1d)): # Skip embeddings if "embedding" in name.lower() and not self.config.quantize_embeddings: continue if isinstance(module, nn.Conv1d) and not self.config.quantize_conv1d: continue hook = module.register_forward_hook(self._hook_factory(name)) hooks.append(hook) layer_modules.append((name, module)) # Forward pass em modo eval (sem gradientes) model.eval() model.to(device) was_training = model.training model.eval() try: with torch.no_grad(): for i, batch in enumerate(dataloader): if i >= n_batches: break try: # Tentar diferentes formatos de batch if isinstance(batch, dict): input_ids_a = batch.get("input_ids_a") input_ids_b = batch.get("input_ids_b") images = batch.get("images") audios = batch.get("audios") if input_ids_a is not None and input_ids_b is not None: # Tentar forward do modelo multimodal try: model( input_ids_a.to(device), input_ids_b.to(device), images=images.to(device) if images is not None else None, audios=audios.to(device) if audios is not None else None, mode="classify", ) except Exception: # Fallback: forward sem imagens/áudios model(input_ids_a.to(device), input_ids_b.to(device)) elif isinstance(batch, (list, tuple)) and len(batch) >= 2: model(batch[0].to(device), batch[1].to(device)) else: logger.debug(f"Batch format não reconhecido: {type(batch)}") except Exception as e: logger.debug(f"Calibração batch {i} falhou: {e}") continue finally: # Remover hooks for hook in hooks: hook.remove() if was_training: model.train() # Coletar max|W| para cada camada for name, module in layer_modules: if name not in self.stats: continue w = module.weight.data if isinstance(module, nn.Linear): # w: [d_out, d_in] -> max sobre d_out max_w = w.abs().amax(dim=0) # [d_in] elif isinstance(module, nn.Conv1d): # w: [out_channels, in_channels//groups, kernel_size] # max sobre out_channels e kernel_size max_w = w.abs().amax(dim=(0, 2)) # [in_channels] else: continue self.stats[name]["max_w"] = max_w.clone() logger.info( "SmoothQuant calibrado: %d camadas, %d batches", len(self.stats), n_batches, ) return self.stats # ============================================================================ # SmoothQuantizer # ============================================================================ class SmoothQuantizer: """Aplica quantização W8A8 com SmoothQuant a um modelo. Args: config: configuração W8A8 """ def __init__(self, config: W8A8Config): self.config = config self.calibrator = SmoothQuantCalibrator(config) def compute_scales( self, stats: Dict[str, Dict[str, torch.Tensor]], ) -> Dict[str, torch.Tensor]: """Computa fatores de escala s_i = (max_x^alpha) / (max_w^(1-alpha)). Args: stats: dict {layer_name: {"max_x": [d_in], "max_w": [d_in]}} Returns: dict {layer_name: scale [d_in]} """ scales = {} alpha = self.config.alpha eps = 1e-8 for name, s in stats.items(): max_x = s["max_x"].float() max_w = s.get("max_w") if max_w is None: # Sem max_w (não calibrado), usa só max_x scale = torch.ones_like(max_x) else: max_w = max_w.float().to(max_x.device) # s = (max_x^alpha) / (max_w^(1-alpha)) # Adiciona eps para evitar divisão por zero scale = (max_x.clamp(min=eps).pow(alpha) / max_w.clamp(min=eps).pow(1 - alpha)) # Clipa outliers se configurado if self.config.outlier_clip_std is not None: mean = scale.mean() std = scale.std() threshold = self.config.outlier_clip_std * std scale = scale.clamp(min=mean - threshold, max=mean + threshold) # Normaliza para média 1 (preserva escala global) scale = scale * (scale.numel() / scale.sum().clamp(min=eps)) scales[name] = scale return scales def quantize_tensor_per_channel( self, tensor: torch.Tensor, scale: Optional[torch.Tensor] = None, n_bits: int = 8, axis: int = -1, ) -> torch.Tensor: """Quantiza um tensor para INT8 (per-channel ou per-tensor). Args: tensor: tensor a quantizar scale: escala por canal (se None, calcula automaticamente) Se fornecido, deve ter shape broadcastable com `tensor`. Para per-row em [d_out, d_in], use scale shape [d_out, 1]. n_bits: número de bits (default 8) axis: eixo para per-channel (apenas quando scale=None e é 1D) Returns: tensor quantizado (dequantizado para FP32 para uso em forward) """ qmax = 2 ** (n_bits - 1) - 1 # 127 para INT8 simétrico qmin = -qmax if scale is None: # Per-tensor max_abs = tensor.abs().max() scale = max_abs.clamp(min=1e-8) / qmax # Quantize-dequantize q = torch.round(tensor / scale).clamp(qmin, qmax) return q * scale else: # Per-channel: scale deve ser broadcastable com tensor scale = scale.to(tensor.device) # Se scale é 1D, expandimos para o eixo if scale.dim() == 1: shape = [1] * tensor.dim() shape[axis] = scale.size(0) scale_b = scale.view(shape) else: # scale já tem shape broadcastable scale_b = scale q = torch.round(tensor / scale_b).clamp(qmin, qmax) return q * scale_b def quantize_model( self, model: nn.Module, dataloader=None, ) -> nn.Module: """Aplica quantização W8A8 ao modelo. Args: model: modelo a quantizar dataloader: dados de calibração (necessário para SmoothQuant estático) Returns: modelo quantizado (parâmetros substituídos por versões INT8 simuladas) """ config = self.config device = config.device model.to(device) # 1. Calibrar se dataloader fornecido if dataloader is not None: stats = self.calibrator.calibrate(model, dataloader) scales = self.compute_scales(stats) else: stats = {} scales = {} # 2. Aplicar SmoothQuant + quantização W8A8 a cada camada n_quantized = 0 n_skipped = 0 with torch.no_grad(): for name, module in model.named_modules(): if not isinstance(module, (nn.Linear, nn.Conv1d)): continue if "embedding" in name.lower() and not config.quantize_embeddings: n_skipped += 1 continue if isinstance(module, nn.Conv1d) and not config.quantize_conv1d: n_skipped += 1 continue # Aplicar SmoothQuant: W' = W * s (migrar escala dos pesos) scale = scales.get(name) w = module.weight.data if scale is not None: # Suavizar pesos: multiplicar pela escala if isinstance(module, nn.Linear): # w: [d_out, d_in], scale: [d_in] w_smoothed = w * scale.to(w.device).unsqueeze(0) elif isinstance(module, nn.Conv1d): # w: [out_ch, in_ch//groups, kernel_size], scale: [in_ch] w_smoothed = w * scale.to(w.device).view(1, -1, 1) else: w_smoothed = w else: w_smoothed = w # Quantizar pesos para INT8 (simulado — dequantiza de volta) if config.scale_type == "per_channel": if isinstance(module, nn.Linear): # Per-output-channel para pesos w_scale = w_smoothed.abs().amax(dim=1) / 127.0 w_scale = w_scale.clamp(min=1e-8) w_q = self.quantize_tensor_per_channel( w_smoothed, scale=w_scale.unsqueeze(1), axis=1, ) elif isinstance(module, nn.Conv1d): w_scale = w_smoothed.abs().amax(dim=(1, 2)) / 127.0 w_scale = w_scale.clamp(min=1e-8) w_q = self.quantize_tensor_per_channel( w_smoothed, scale=w_scale.view(-1, 1, 1), axis=0, ) else: # Per-tensor w_q = self.quantize_tensor_per_channel(w_smoothed) module.weight.data = w_q.to(module.weight.dtype) n_quantized += 1 logger.info( "SmoothQuant W8A8 aplicado: %d camadas quantizadas, %d ignoradas", n_quantized, n_skipped, ) # Marcar modelo como quantizado (para uso futuro) model._is_quantized_w8a8 = True # type: ignore model._quantization_config = config # type: ignore return model @staticmethod def is_quantized(model: nn.Module) -> bool: """Verifica se um modelo foi quantizado.""" return getattr(model, "_is_quantized_w8a8", False) # ============================================================================ # Helper: quantize modelo inteiro # ============================================================================ def quantize_model_w8a8( model: nn.Module, dataloader=None, alpha: float = 0.5, n_calibration_batches: int = 5, device: str = "cpu", **kwargs, ) -> nn.Module: """Atalho para quantizar um modelo W8A8. Args: model: modelo a quantizar dataloader: dados de calibração (None para quantização dinâmica) alpha: fator SmoothQuant (0.5 default) n_calibration_batches: batches para calibração device: device para calibração **kwargs: outros parâmetros de W8A8Config Returns: modelo quantizado """ config = W8A8Config( alpha=alpha, n_calibration_batches=n_calibration_batches, device=device, **kwargs, ) quantizer = SmoothQuantizer(config) return quantizer.quantize_model(model, dataloader=dataloader) # ============================================================================ # Helper: estimar redução de memória # ============================================================================ def estimate_memory_savings(model: nn.Module) -> Dict[str, float]: """Estima redução de memória após quantização W8A8. Como a quantização é SIMULADA (dequantiza de volta para FP32 para uso em forward), a memória real não muda. Esta função reporta a redução POTENCIAL se os pesos fossem armazenados como INT8 real. Args: model: modelo (preferencialmente quantizado) Returns: dict com tamanhos em MB """ is_quantized = getattr(model, "_is_quantized_w8a8", False) fp32_bytes = 0 # embeddings + biases (sempre FP32) quantizable_bytes = 0 # pesos que seriam INT8 for name, param in model.named_parameters(): n = param.numel() if "embedding" in name.lower(): # Embeddings mantêm FP32 fp32_bytes += n * 4 elif "bias" in name.lower(): # Biases mantêm FP32 (típico em W8A8) fp32_bytes += n * 4 elif param.dim() >= 2: # Pesos de matrizes (Linear, Conv) — quantizáveis quantizable_bytes += n else: # Outros 1D — mantêm FP32 fp32_bytes += n * 4 # Se quantizado: pesos seriam INT8 (1 byte cada) # Se não quantizado: pesos seriam FP32 (4 bytes cada) if is_quantized: int8_bytes = quantizable_bytes * 1 # INT8 else: int8_bytes = quantizable_bytes * 4 # FP32 fp32_mb = fp32_bytes / (1024 ** 2) int8_mb = int8_bytes / (1024 ** 2) total_mb = fp32_mb + int8_mb # Sem quantização, tudo seria FP32 no_quant_mb = (fp32_bytes + quantizable_bytes * 4) / (1024 ** 2) reduction = (1 - total_mb / no_quant_mb) * 100 if no_quant_mb > 0 else 0 return { "fp32_mb": fp32_mb, "int8_mb": int8_mb, "total_mb": total_mb, "no_quant_mb": no_quant_mb, "reduction_pct": reduction, "is_quantized": is_quantized, } # ============================================================================ # Self-test # ============================================================================ def _self_test(): """Teste rápido da quantização W8A8.""" torch.manual_seed(42) # Modelo simples para teste class TinyModel(nn.Module): def __init__(self): super().__init__() self.fc1 = nn.Linear(32, 64) self.fc2 = nn.Linear(64, 32) self.embedding = nn.Embedding(100, 32) def forward(self, x): return self.fc2(torch.relu(self.fc1(x))) model = TinyModel() print(f"Antes: {sum(p.numel() for p in model.parameters())} params") # Quantizar sem dataloader (dinâmico) quantized = quantize_model_w8a8(model, dataloader=None, alpha=0.5) print(f"Quantizado: {SmoothQuantizer.is_quantized(quantized)}") # Verificar forward ainda funciona x = torch.randn(2, 32) out = quantized(x) print(f"Output shape: {out.shape}") # Estimar savings savings = estimate_memory_savings(quantized) print(f"Memory savings: {savings}") if __name__ == "__main__": _self_test() __all__ = [ "W8A8Config", "SmoothQuantCalibrator", "SmoothQuantizer", "quantize_model_w8a8", "estimate_memory_savings", ]