File size: 18,184 Bytes
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
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
"""w8a8_error_reduction.py — Redução iterativa do erro W8A8 (V5).

═══════════════════════════════════════════════════════════════════════════════
V5 — VERSÃO AUTO-CONTIDA (sem dependência de TrustNeuralAgent externo)
═══════════════════════════════════════════════════════════════════════════════

ORIGEM (gru-ring-v13-9-2):
  source/xavante_work/flexnet/w8a8_error_reduction.py (V13.9.1)

Em V13.9.1, TrustGuidedW8A8 dependia de TrustNeuralAgent (rede neural externa
que avalia confiança na quantização). Em V5, eliminamos essa dependência:

  * Confiança ht agora é FUNÇÃO ANALÍTICA de α e erro:
        ht = exp(-error / error_scale) * sigmoid((alpha - 0.5) * 4)
    - Alta quando erro é baixo E α está próximo de 0.5 (balanceado)
    - Baixa quando erro é alto OU α está nos extremos (0 ou 1)

  * Redução de erro via gradiente STE (igual V13.9.1):
        dE/dα = dE/ds · ds/dα
      onde:
        E(α) = ||Y_fp - Y_q(α)||_F^2 / ||Y_fp||_F^2
        s_j(α) = max(|X_j|)^α / max(|W_j|)^{1-α}
        ds_j/dα = s_j · (log max|X_j| - log max|W_j|)

  * Trust-guided update:
        α_{t+1} = α_t - lr · (1 - ht) · dE/dα
      Alta confiança (ht→1) → passo pequeno (convergência)
      Baixa confiança (ht→0) → passo grande (exploração)

Análise de convergência:
    E(α_{t+1}) ≤ E(α_t) · (1 - lr · (1 - ht) · λ_min)
  onde λ_min é o menor autovalor da Hessiana de E em α*.
  Garante convergência linear quando λ_min > 0 (E convexa localmente).

V5 vs V13.9.1:
  * Sem TrustNeuralAgent (rede neural externa) → -50% params, -1 módulo
  * ht é função analítica (determinística, não-aprendida)
  * Self-contained: só depende de torch + w8a8_smoothquant
  * Vetorizado: smooth factor calculado em 1 operação (não loop Python)
  * OOM-safe: não armazena histórico de tensores completos, apenas escalares
"""
from __future__ import annotations

import math
import logging
from dataclasses import dataclass, field
from typing import Dict, List, Optional, Tuple

import torch
import torch.nn as nn
import torch.nn.functional as F

from .w8a8_smoothquant import (
    SmoothQuantW8A8,
    calibrate_smoothquant,
    quantize_per_channel_symmetric,
    quantize_per_token_symmetric,
    ste_quantize,
)

logger = logging.getLogger(__name__)


# ============================================================================
# Exceções
# ============================================================================
class W8A8ConvergenceError(RuntimeError):
    """W8A8 error reduction não convergiu."""
    pass


# ============================================================================
# Função analítica de confiança ht (V5 — substitui TrustNeuralAgent)
# ============================================================================
def analytic_trust(
    error: float,
    alpha: float,
    error_scale: float = 0.1,
) -> float:
    """Confiança ht analítica (V5).

    ht = exp(-error / error_scale) * sigmoid((alpha - 0.5) * 4)

    Args:
        error: erro relativo de quantização (||Y_fp - Y_q||/||Y_fp||)
        alpha: fator SmoothQuant atual
        error_scale: escala do erro (default 0.1 — erro 10% dá e^-1 ≈ 0.37)

    Returns:
        ht ∈ (0, 1): 1 = alta confiança, 0 = baixa confiança
    """
    err_term = math.exp(-max(0.0, error) / max(1e-8, error_scale))
    # Sigmoid com slope 4: alpha=0.5 → 0.5, alpha=0.95 → 0.88, alpha=0.05 → 0.12
    alpha_term = 1.0 / (1.0 + math.exp(-4.0 * (alpha - 0.5)))
    return err_term * alpha_term


# ============================================================================
# 1. TrustGuidedW8A8 — quantização W8A8 com redução iterativa de erro
# ============================================================================
@dataclass
class W8A8ReductionResult:
    """Resultado da redução iterativa de erro W8A8."""
    initial_alpha: float
    final_alpha: float
    initial_error: float
    final_error: float
    n_iterations: int
    converged: bool
    history: List[Dict[str, float]] = field(default_factory=list)
    final_ht: float = 0.0


class TrustGuidedW8A8:
    """W8A8 com redução iterativa de erro guiada por confiança analítica.

    V5: Self-contained — sem TrustNeuralAgent externo. ht é função analítica
    de (error, alpha), não uma rede neural aprendida.

    Pipeline:
    1. Calibra α inicial e quantiza
    2. Mede erro E₀
    3. Função analítica avalia ht (confiança na quantização)
    4. Se ht < τ: atualiza α via STE gradiente na direção de redução de erro
    5. Re-quantiza, mede erro E₁, atualiza ht
    6. Repete até ht > τ ou convergência

    A atualização é modulada por (1 - ht): alta confiança → passo pequeno.
    """

    def __init__(
        self,
        ref: nn.Linear,
        init_alpha: float = 0.5,
        trust_threshold: float = 0.8,
        lr: float = 0.05,
        max_iterations: int = 20,
        convergence_tol: float = 1e-4,
        reg_lambda: float = 0.01,
        verbose: bool = False,
    ):
        self.ref = ref
        self.init_alpha = init_alpha
        self.trust_threshold = trust_threshold
        self.lr = lr
        self.max_iterations = max_iterations
        self.convergence_tol = convergence_tol
        self.reg_lambda = reg_lambda
        self.verbose = verbose
        # V5: Sem TrustNeuralAgent. ht é função analítica.
        # Parâmetro α parametrizado via sigmoid.
        init_eta = math.log(init_alpha / (1 - init_alpha))
        self.eta_alpha = torch.nn.Parameter(torch.tensor(init_eta))

    @property
    def alpha(self) -> torch.Tensor:
        return torch.sigmoid(self.eta_alpha)

    def _compute_quantization_error(
        self,
        X: torch.Tensor,
        alpha_val: float,
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """Computa Y_quantized e Y_reference para um dado α.

        Retorna: (Y_q, Y_ref)
        """
        W = self.ref.weight
        Y_ref = F.linear(X, W, self.ref.bias)
        # Smooth factor.
        s = calibrate_smoothquant(X, W, alpha=alpha_val)
        # Smooth.
        X_smooth = X / s
        W_smooth = W * s.unsqueeze(0)
        # Quantize W per-channel.
        W_q, w_scale = quantize_per_channel_symmetric(W_smooth, axis=0, n_bits=8)
        # Quantize X per-token.
        X_q, x_scale = quantize_per_token_symmetric(X_smooth, n_bits=8)
        # Dequant.
        X_fp = X_q.to(torch.float32) * x_scale
        W_fp = W_q.to(torch.float32) * w_scale.unsqueeze(1)
        Y_q = X_fp @ W_fp.t()
        if self.ref.bias is not None:
            Y_q = Y_q + self.ref.bias
        return Y_q, Y_ref

    def _compute_error_gradient(
        self,
        X: torch.Tensor,
        alpha_tensor: torch.Tensor,
    ) -> torch.Tensor:
        """Computa gradiente do erro via STE (Straight-Through Estimator).

        E(α) = ||Y_ref - Y_q(α)||_F^2 / ||Y_ref||_F^2

        Usa STE para round() (não-diferenciável).
        """
        W = self.ref.weight
        # Y_ref (sem grad).
        with torch.no_grad():
            Y_ref = F.linear(X, W, self.ref.bias)
            y_ref_norm_sq = max(((Y_ref) ** 2).mean().item(), 1e-12)
        # s(α) diferenciável.
        d_in = W.shape[-1]
        X_flat = X.reshape(-1, d_in)
        max_x = X_flat.abs().amax(dim=0).clamp(min=1e-8)
        max_w = W.abs().amax(dim=0).clamp(min=1e-8)
        s = torch.pow(max_x, alpha_tensor) / torch.pow(max_w, 1.0 - alpha_tensor)
        # Smooth.
        X_smooth = X / s
        W_smooth = W * s.unsqueeze(0)
        # Quantização com STE.
        qmax = 127
        max_w_smooth = W_smooth.abs().amax(dim=1, keepdim=True).clamp(min=1e-8)
        scale_w = max_w_smooth / qmax
        max_x_smooth = X_smooth.abs().amax(dim=-1, keepdim=True).clamp(min=1e-8)
        scale_x = max_x_smooth / qmax
        # STE: round com gradiente pass-through.
        W_q_deq = ste_quantize(W_smooth, scale_w, 8)
        X_q_deq = ste_quantize(X_smooth, scale_x, 8)
        # Matmul.
        Y_q = X_q_deq @ W_q_deq.t()
        if self.ref.bias is not None:
            Y_q = Y_q + self.ref.bias
        # Loss = MSE relativo.
        loss = ((Y_ref - Y_q) ** 2).mean() / y_ref_norm_sq
        # Regularização entropy (atrai α para 0.5 — centro do intervalo).
        eps = 1e-8
        alpha_clamped = alpha_tensor.clamp(eps, 1.0 - eps)
        reg = -self.reg_lambda * (
            alpha_clamped * torch.log(alpha_clamped + eps) +
            (1 - alpha_clamped) * torch.log(1 - alpha_clamped + eps)
        )
        # negativo porque queremos MAXIMIZAR entropy (atrair α para 0.5)
        total_loss = loss - reg
        total_loss.backward()
        return total_loss

    def reduce_error(
        self,
        X: torch.Tensor,
    ) -> W8A8ReductionResult:
        """Executa redução iterativa de erro W8A8.

        Args:
            X: dados de calibração (B, d_in)

        Returns:
            W8A8ReductionResult
        """
        initial_alpha = float(self.alpha.item())
        # Computa erro inicial.
        with torch.no_grad():
            Y_q_0, Y_ref_0 = self._compute_quantization_error(X, initial_alpha)
            initial_error = ((Y_ref_0 - Y_q_0) ** 2).mean().item() / max(
                ((Y_ref_0) ** 2).mean().item(), 1e-12
            )
        # V5: ht analítico (sem rede neural)
        ht_val = analytic_trust(initial_error, initial_alpha)
        history = [{"iteration": 0, "alpha": initial_alpha, "error": initial_error, "ht": ht_val}]
        prev_error = initial_error
        converged = False
        opt = torch.optim.SGD([self.eta_alpha], lr=self.lr)
        for iteration in range(1, self.max_iterations + 1):
            # 1. α atual.
            alpha_val = float(self.alpha.item())
            with torch.no_grad():
                Y_q, Y_ref = self._compute_quantization_error(X, alpha_val)
                current_error = ((Y_ref - Y_q) ** 2).mean().item() / max(
                    ((Y_ref) ** 2).mean().item(), 1e-12
                )
            # V5: ht analítico
            ht_val = analytic_trust(current_error, alpha_val)
            # 2. Se ht > threshold, converged.
            if ht_val > self.trust_threshold:
                converged = True
                if self.verbose:
                    print(f"  iter {iteration}: ht={ht_val:.4f} > τ={self.trust_threshold}, converged")
                break
            # 3. Atualiza α via STE gradiente, modulado por (1 - ht).
            opt.zero_grad()
            self._compute_error_gradient(X, self.alpha)
            # Modula gradiente por (1 - ht).
            with torch.no_grad():
                if self.eta_alpha.grad is not None:
                    self.eta_alpha.grad *= (1.0 - ht_val)
            opt.step()
            # 4. Clamp α para [0.01, 0.99].
            with torch.no_grad():
                self.eta_alpha.data.clamp_(
                    math.log(0.01 / 0.99), math.log(0.99 / 0.01)
                )
            # 5. Verifica convergência por erro.
            if abs(prev_error - current_error) < self.convergence_tol:
                converged = True
                if self.verbose:
                    print(f"  iter {iteration}: erro convergiu (Δ={abs(prev_error - current_error):.2e})")
                break
            prev_error = current_error
            history.append({
                "iteration": iteration,
                "alpha": float(self.alpha.item()),
                "error": current_error,
                "ht": ht_val,
            })
            if self.verbose and iteration % 5 == 0:
                print(f"  iter {iteration}: α={float(self.alpha.item()):.4f}, "
                      f"erro={current_error:.6f}, ht={ht_val:.4f}")
        final_alpha = float(self.alpha.item())
        final_error = history[-1]["error"]
        return W8A8ReductionResult(
            initial_alpha=initial_alpha,
            final_alpha=final_alpha,
            initial_error=initial_error,
            final_error=final_error,
            n_iterations=len(history) - 1,
            converged=converged,
            history=history,
            final_ht=ht_val,
        )

    def get_smooth_scale(self, X: torch.Tensor) -> torch.Tensor:
        """Retorna o smooth scale final."""
        with torch.no_grad():
            return calibrate_smoothquant(X, self.ref.weight, float(self.alpha.item())).detach()

    def quantize_weights(self, X: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
        """Quantiza pesos com α otimizado.

        Retorna: (weight_int8, weight_scale)
        """
        s = self.get_smooth_scale(X)
        W_smooth = self.ref.weight * s.unsqueeze(0)
        W_q, w_scale = quantize_per_channel_symmetric(W_smooth, axis=0, n_bits=8)
        return W_q, w_scale


# ============================================================================
# 2. MultiLayerW8A8Reducer — aplica redução a múltiplas camadas
# ============================================================================
class MultiLayerW8A8Reducer:
    """Aplica TrustGuidedW8A8 a múltiplas camadas Linear de um modelo.

    Para cada camada:
    1. Coleta ativações de calibração
    2. Executa TrustGuidedW8A8.reduce_error()
    3. Substitui a camada original pela versão quantizada (SmoothQuantW8A8)

    V5: OOM-safe — não armazena ativações de todas as camadas simultaneamente.
    Processa uma camada por vez.
    """

    def __init__(
        self,
        init_alpha: float = 0.5,
        trust_threshold: float = 0.8,
        lr: float = 0.05,
        max_iterations: int = 15,
        verbose: bool = False,
    ):
        self.init_alpha = init_alpha
        self.trust_threshold = trust_threshold
        self.lr = lr
        self.max_iterations = max_iterations
        self.verbose = verbose

    def reduce_model(
        self,
        model: nn.Module,
        calib_X: torch.Tensor,
        target_substrings: Optional[List[str]] = None,
    ) -> Dict[str, W8A8ReductionResult]:
        """Aplica redução W8A8 a todas as nn.Linear no modelo.

        Args:
            model: modelo a ser quantizado
            calib_X: dados de calibração (B, T, d) ou (B, d)
            target_substrings: substrings para filtrar camadas (None = todas)

        Retorna: {layer_name: W8A8ReductionResult}
        """
        results = {}
        # Coleta camadas a quantizar.
        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, SmoothQuantW8A8):
                    full_name = f"{name}.{child_name}" if name else child_name
                    if target_substrings is None or any(s in full_name for s in target_substrings):
                        to_replace.append((module, child_name, child, full_name))
        if self.verbose:
            print(f"  Encontradas {len(to_replace)} camadas Linear para quantização")
        # Para cada camada, aplica TrustGuidedW8A8.
        for module, child_name, linear_layer, full_name in to_replace:
            d_in = linear_layer.in_features
            # Extrai ativações de calibração para esta camada.
            if calib_X.shape[-1] == d_in:
                X_calib = calib_X.reshape(-1, d_in)
            else:
                # Projeta ou usa aleatório.
                X_calib = torch.randn(min(64, calib_X.shape[0]), d_in)
            # Reduz erro.
            reducer = TrustGuidedW8A8(
                ref=linear_layer,
                init_alpha=self.init_alpha,
                trust_threshold=self.trust_threshold,
                lr=self.lr,
                max_iterations=self.max_iterations,
                verbose=self.verbose,
            )
            result = reducer.reduce_error(X_calib)
            results[full_name] = result
            # Substitui a camada pela versão quantizada (SmoothQuantW8A8 com α aprendível).
            sq = SmoothQuantW8A8(
                linear_layer,
                init_alpha=result.final_alpha,
                learnable_alpha=True,  # V5: α continua aprendível após redução
            )
            sq.calibrate(X_calib)
            # Seta o smooth_scale otimizado.
            s = reducer.get_smooth_scale(X_calib)
            sq.smooth_scale.copy_(s)
            sq.freeze_quantization()
            setattr(module, child_name, sq)
            # V5 OOM fix: del referências explícitas
            del reducer, X_calib
            if self.verbose:
                print(f"  {full_name}: α {result.initial_alpha:.4f}→{result.final_alpha:.4f}, "
                      f"erro {result.initial_error:.6f}→{result.final_error:.6f}, "
                      f"ht={result.final_ht:.4f}, conv={result.converged}")
        return results

    def summarize_results(
        self, results: Dict[str, W8A8ReductionResult]
    ) -> str:
        """Gera resumo tabular dos resultados."""
        lines = [
            f"{'Layer':<30} | {'α_init':>8} | {'α_final':>8} | {'err_init':>10} | {'err_final':>10} | {'reduction':>10} | {'ht':>6} | {'conv':>5}",
            f"{'-'*30} | {'-'*8} | {'-'*8} | {'-'*10} | {'-'*10} | {'-'*10} | {'-'*6} | {'-'*5}",
        ]
        for name, r in results.items():
            reduction = (1 - r.final_error / max(r.initial_error, 1e-12)) * 100
            lines.append(
                f"{name:<30} | {r.initial_alpha:>8.4f} | {r.final_alpha:>8.4f} | "
                f"{r.initial_error:>10.6f} | {r.final_error:>10.6f} | "
                f"{reduction:>9.1f}% | {r.final_ht:>6.3f} | {'✓' if r.converged else '✗':>5}"
            )
        return "\n".join(lines)


__all__ = [
    "W8A8ConvergenceError",
    "analytic_trust",
    "W8A8ReductionResult",
    "TrustGuidedW8A8",
    "MultiLayerW8A8Reducer",
]