| """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__) |
|
|
|
|
| |
| |
| |
|
|
| @dataclass |
| class W8A8Config: |
| """Configuração da quantização W8A8 com SmoothQuant.""" |
| |
| alpha: float = 0.5 |
| |
| n_calibration_batches: int = 5 |
| |
| scale_type: str = "per_channel" |
| |
| dynamic_weight: bool = False |
| |
| quantize_embeddings: bool = False |
| |
| quantize_conv1d: bool = True |
| |
| quantize_layernorm: bool = False |
| |
| outlier_clip_std: Optional[float] = 4.0 |
| |
| device: str = "cpu" |
| |
| simulate: bool = False |
|
|
|
|
| |
| |
| |
|
|
| 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): |
| |
| if not isinstance(input, tuple) or len(input) == 0: |
| return |
| x = input[0] |
| if not isinstance(x, torch.Tensor): |
| return |
| |
| with torch.no_grad(): |
| if x.dim() >= 2: |
| |
| x_flat = x.reshape(-1, x.size(-1)) |
| max_x = x_flat.abs().max(dim=0).values |
| else: |
| max_x = x.abs() |
| |
| 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 |
|
|
| |
| hooks = [] |
| layer_modules = [] |
| for name, module in model.named_modules(): |
| if isinstance(module, (nn.Linear, nn.Conv1d)): |
| |
| 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)) |
|
|
| |
| 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: |
| |
| 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: |
| |
| 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: |
| |
| 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: |
| |
| for hook in hooks: |
| hook.remove() |
| if was_training: |
| model.train() |
|
|
| |
| for name, module in layer_modules: |
| if name not in self.stats: |
| continue |
| w = module.weight.data |
| if isinstance(module, nn.Linear): |
| |
| max_w = w.abs().amax(dim=0) |
| elif isinstance(module, nn.Conv1d): |
| |
| |
| max_w = w.abs().amax(dim=(0, 2)) |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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: |
| |
| scale = torch.ones_like(max_x) |
| else: |
| max_w = max_w.float().to(max_x.device) |
| |
| |
| scale = (max_x.clamp(min=eps).pow(alpha) / |
| max_w.clamp(min=eps).pow(1 - alpha)) |
| |
| 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) |
| |
| 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 |
| qmin = -qmax |
|
|
| if scale is None: |
| |
| max_abs = tensor.abs().max() |
| scale = max_abs.clamp(min=1e-8) / qmax |
| |
| q = torch.round(tensor / scale).clamp(qmin, qmax) |
| return q * scale |
| else: |
| |
| scale = scale.to(tensor.device) |
| |
| if scale.dim() == 1: |
| shape = [1] * tensor.dim() |
| shape[axis] = scale.size(0) |
| scale_b = scale.view(shape) |
| else: |
| |
| 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) |
|
|
| |
| if dataloader is not None: |
| stats = self.calibrator.calibrate(model, dataloader) |
| scales = self.compute_scales(stats) |
| else: |
| stats = {} |
| scales = {} |
|
|
| |
| 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 |
|
|
| |
| scale = scales.get(name) |
| w = module.weight.data |
| if scale is not None: |
| |
| if isinstance(module, nn.Linear): |
| |
| w_smoothed = w * scale.to(w.device).unsqueeze(0) |
| elif isinstance(module, nn.Conv1d): |
| |
| w_smoothed = w * scale.to(w.device).view(1, -1, 1) |
| else: |
| w_smoothed = w |
| else: |
| w_smoothed = w |
|
|
| |
| if config.scale_type == "per_channel": |
| if isinstance(module, nn.Linear): |
| |
| 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: |
| |
| 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, |
| ) |
|
|
| |
| model._is_quantized_w8a8 = True |
| model._quantization_config = config |
|
|
| return model |
|
|
| @staticmethod |
| def is_quantized(model: nn.Module) -> bool: |
| """Verifica se um modelo foi quantizado.""" |
| return getattr(model, "_is_quantized_w8a8", False) |
|
|
|
|
| |
| |
| |
|
|
| 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) |
|
|
|
|
| |
| |
| |
|
|
| 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 |
| quantizable_bytes = 0 |
|
|
| for name, param in model.named_parameters(): |
| n = param.numel() |
| if "embedding" in name.lower(): |
| |
| fp32_bytes += n * 4 |
| elif "bias" in name.lower(): |
| |
| fp32_bytes += n * 4 |
| elif param.dim() >= 2: |
| |
| quantizable_bytes += n |
| else: |
| |
| fp32_bytes += n * 4 |
|
|
| |
| |
| if is_quantized: |
| int8_bytes = quantizable_bytes * 1 |
| else: |
| int8_bytes = quantizable_bytes * 4 |
|
|
| fp32_mb = fp32_bytes / (1024 ** 2) |
| int8_mb = int8_bytes / (1024 ** 2) |
| total_mb = fp32_mb + int8_mb |
| |
| 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, |
| } |
|
|
|
|
| |
| |
| |
|
|
| def _self_test(): |
| """Teste rápido da quantização W8A8.""" |
| torch.manual_seed(42) |
|
|
| |
| 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") |
|
|
| |
| quantized = quantize_model_w8a8(model, dataloader=None, alpha=0.5) |
| print(f"Quantizado: {SmoothQuantizer.is_quantized(quantized)}") |
|
|
| |
| x = torch.randn(2, 32) |
| out = quantized(x) |
| print(f"Output shape: {out.shape}") |
|
|
| |
| 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", |
| ] |
|
|