Download cnn_bigru/utils/quantization.py from PowerMachine/CNN-BiGRU: direct link, hf CLI and curl.
- Browser
- Download file 21.9 kB
-
https://huggingface.co/PowerMachine/CNN-BiGRU/resolve/main/cnn_bigru/utils/quantization.py
- Command line
-
hf download hf://PowerMachine/CNN-BiGRU/cnn_bigru/utils/quantization.py
-
curl -L -o quantization.py https://huggingface.co/PowerMachine/CNN-BiGRU/resolve/main/cnn_bigru/utils/quantization.py
21.9 kB
| """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 | |
| # ============================================================================ | |
| 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 | |
| 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", | |
| ] | |