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",
]