CNN-BiGRU / cnn_bigru /utils /quantization.py
PowerMachine's picture
v3.0: reorganiza arquivos sob cnn_bigru/ (preserva árvore de pastas)
49b8205 verified
Raw
History Blame Contribute Delete
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
# ============================================================================
@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",
]