File size: 6,361 Bytes
cbf308e 3275441 cbf308e 3275441 cbf308e 3275441 cbf308e 3275441 cbf308e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 | """quantized_linear.py — QuantizedLinear + apply_w8a8 (V5 — bug fix crítico).
═══════════════════════════════════════════════════════════════════════════════
V5 — BUG FIX CRÍTICO: apply_w8a8 não substituía módulos
═══════════════════════════════════════════════════════════════════════════════
BUG (V1-V4):
Em V1-V4, `apply_w8a8` era usado via `model.apply(apply_w8a8)`. Porém,
`nn.Module.apply(fn)` chama `fn(module)` e DESCARTA o valor retornado.
A função retornava `QuantizedLinear(...)` mas o módulo original não era
substituído — apenas a função era chamada (sem efeito).
Resultado: W8A8 NUNCA era aplicado em V1-V4. O modelo usava FP32 puro
em todas as camadas Linear, desperdiçando ~4x mais memória que o necessário
e sem obter o benefício do ruído de quantização (Lema 3).
FIX (V5):
Substituir `model.apply(apply_w8a8)` por `apply_w8a8_to_model(model)` que
percorre recursivamente `model._modules` e substitui `nn.Linear` por
`QuantizedLinear` (ou `SmoothQuantW8A8` se preferir V5).
Para manter compatibilidade com V1-V4, `apply_w8a8(module)` ainda existe
mas agora é chamado recursivamente por `apply_w8a8_to_model`.
Implementa o Lema 3 (Cancelamento de ruído de quantização):
A quantização W8A8 (pesos e ativações em 8 bits) é simulada no forward
via fake quantization (round + clamp + dequant). O backward usa Straight-
Through Estimator (STE) — o gradiente passa direto pela operação de round.
"""
from __future__ import annotations
import torch
import torch.nn as nn
import torch.nn.functional as F
def quantize_tensor(x: torch.Tensor, num_bits: int = 8) -> torch.Tensor:
"""Quantização simétrica uniforme (fake quant).
Args:
x: tensor a quantizar
num_bits: nº de bits (default 8 → range [-128, 127])
Returns:
x_deq: tensor dequantizado (mesmo shape, mesmo dtype)
"""
if x is None:
return None
qmin = -(2 ** (num_bits - 1))
qmax = 2 ** (num_bits - 1) - 1
# Scale por tensor (não por canal) — simplificação
scale = x.abs().max() / qmax
if scale < 1e-8:
scale = torch.tensor(1.0, dtype=x.dtype, device=x.device)
# Round + clamp (STE backward — grad passa direto)
x_q = torch.round(x / scale).clamp(qmin, qmax)
x_deq = x_q * scale
return x_deq
class QuantizedLinear(nn.Linear):
"""Camada Linear com quantização W8A8 (fake quant durante treino).
No forward:
q_weight = quantize_tensor(self.weight, 8) # W8
q_input = quantize_tensor(input, 8) # A8
output = F.linear(q_input, q_weight, self.bias)
O gradiente flui através do STE (round é não-diferenciável, mas o
autograd do PyTorch propaga o gradiente via x_deq = x_q * scale onde
x_q depende de round(x/scale) que tem gradiente zero — o STE substitui
esse gradiente por 1).
Args:
in_features, out_features, bias: iguais ao nn.Linear
num_bits: nº de bits (default 8)
"""
def __init__(self, in_features: int, out_features: int, bias: bool = True, num_bits: int = 8):
super().__init__(in_features, out_features, bias)
self.num_bits = num_bits
def forward(self, input: torch.Tensor) -> torch.Tensor:
# Quantização dos pesos (W8)
q_weight = quantize_tensor(self.weight, self.num_bits)
# Quantização da ativação de entrada (A8)
q_input = quantize_tensor(input, self.num_bits)
output = F.linear(q_input, q_weight, self.bias)
return output
def apply_w8a8(module: nn.Module) -> nn.Module:
"""Substitui uma nn.Linear por QuantizedLinear (mantém pesos).
V5: Agora é uma função que efetivamente substitui o módulo quando
chamada por apply_w8a8_to_model. Para compatibilidade com V1-V4,
ainda pode ser chamada diretamente em um único módulo.
NOTA: Esta função NÃO deve ser usada via `model.apply(apply_w8a8)` pois
`nn.Module.apply` descarta o valor retornado. Use
`apply_w8a8_to_model(model)` em vez disso.
"""
if isinstance(module, nn.Linear) and not isinstance(module, QuantizedLinear):
new_module = QuantizedLinear(module.in_features, module.out_features, module.bias is not None)
# Copiar pesos do módulo original
with torch.no_grad():
new_module.weight.copy_(module.weight)
if module.bias is not None:
new_module.bias.copy_(module.bias)
return new_module
return module
def apply_w8a8_to_model(model: nn.Module) -> int:
"""Substitui recursivamente todas as nn.Linear por QuantizedLinear.
V5: Implementação correta — percorre `model._modules` recursivamente
e substitui in-place. Retorna o número de camadas substituídas.
Diferentemente de `model.apply(apply_w8a8)` (que não funciona porque
`apply` descarta o retorno), esta função modifica o modelo in-place.
Args:
model: modelo a ter as Linears substituídas
Returns:
n_substituidas: número de camadas Linear substituídas
Uso:
from bigru_t.quantization import apply_w8a8_to_model
n = apply_w8a8_to_model(model)
print(f"{n} camadas Linear substituídas por QuantizedLinear")
"""
n_substituidas = 0
# Lista de (parent_module, child_name) para substituir
# Não podemos substituir durante a iteração, então coletamos primeiro
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, QuantizedLinear):
to_replace.append((module, child_name))
for parent, child_name in to_replace:
old_linear = getattr(parent, child_name)
new_linear = apply_w8a8(old_linear)
setattr(parent, child_name, new_linear)
n_substituidas += 1
return n_substituidas
__all__ = [
"quantize_tensor",
"QuantizedLinear",
"apply_w8a8",
"apply_w8a8_to_model",
]
|