khtst-multimodal-ptbr / docs /matematica /08-quantizacao-w8a8.md
PowerMachine's picture
KHTST v7 — RETREINO OBSERVADO: estados antigos apagados; treino do zero com monitor de RAM integral (Agente Engenheiro); métricas de tests/test_khtst_v6.py comparadas com o treino anterior; gráficos de desempenho publicados no card
073c9a9 verified
|
Raw History Blame Contribute Delete
6.83 kB

08 — Quantização W8A8: cotas de erro e treino consciente de quantização

Módulo: quanta/quantizacao.py. Pesos e ativações em 8 bits (item 9 do escopo), com fallback puro-PyTorch quando bitsandbytes não dispõe de backend CPU/CUDA.

1. Quantização simétrica W8

Dado tensor real $X$, com escala por tensor (ou por canal de saída): s  =  max⁡∣X∣127,Q(X)=clamp(round(X/s),  −127,  127),X~=s Q(X).(8.1)s \;=\; \frac{\max|X|}{127}, \qquad Q(X) = \mathrm{clamp}\Big(\mathrm{round}\big(X/s\big),\;-127,\;127\Big), \qquad \tilde X = s\,Q(X). \tag{8.1}

Teorema 8.1 (cota suprema do erro). $|\tilde X - X|_\infty \le s/2$. Demonstração. $\mathrm{round}(u)$ muda $u$ por no máximo $1/2$; o clamp só atua fora de $[-127,127]$, onde $\mathrm{round}(X/s)=\pm127$ exatamente no limite, mudança $<1/2$ após arredondamento; logo $|Q(u)-u|\le 1/2$ e $|\tilde X - X|_\infty = s,|Q(X/s)-X/s|_\infty \le s/2$. $\blacksquare$

Teorema 8.2 (erro médio). Para $X$ com fração de massa uniforme dentro de cada célula de arredondamento, $\mathbb{E}|\tilde X - X| \approx s/4$ (esperança de $|U|$, $U\sim U(-s/2,s/2)$). Uso: a telemetria compara o erro observado com $s/4$; desvios grandes indicam outliers densos (dispara escala por canal).

2. Matmul inteira (caminho W8A8 real)

$\tilde Y = \tilde X \tilde W^\top = s_x s_w, \big(Q(X),Q(W)^\top\big)$. Em hardware com int8 GEMM (CUDA/bnb) a soma acumula em int32; em CPU (nossa execução) faz-se o produto em int32 via torch (conversão explícita) ou simulação fake-quant em fp32 — o erro segue a cota do Teorema 8.1 em ambos os casos, pois a simulação aplica as mesmas funções (8.1).

3. QAT — straight-through estimator (STE)

Na fase de treino quantizado, o passo de quantização tem gradiente nulo quase sempre (escadinhas); usa-se STE (Bengio et al., 2013): forward com $\tilde X$, backward com identidade.

Teorema 8.3 (STE é subgradiente válido). A função $q(u)=s,\mathrm{clamp}(\mathrm{round}(u/s),-127,127)$ é Lipschitz com constante 1 (composição de 1-Lipschitz: clamp e round são 1-Lipschitz; por Rademacher, $q$ é diferenciável a.e. com $|q'|\le 1$; o conjunto de descontinuidades tem medida nula). O STE substitui $q'$ por $1$ — um subgradiente admissível no sentido de Clarke, e o SGD com subgradientes converge para um mínimo local sob Robbins–Monro (Teorema 1.2). $\blacksquare$

4. LLM.int8() (bitsandbytes) — decomposição de outliers

Dettmers et al. (2022): separar colunas de ativação com magnitude acima de limiar $\tau$: XW⊤  =  X:,OWO,:⊤⏟fp16, colunas de outliers O  +  X:,IWI,:⊤⏟int8, I=Oc.(8.2)X W^\top \;=\; \underbrace{X_{:,O} W_{O,:}^\top}_{\text{fp16, colunas de outliers } O} \;+\; \underbrace{X_{:,I} W_{I,:}^\top}_{\text{int8, } I=O^c}. \tag{8.2} Cota de erro: o erro da parte int8 é limitado por $\tfrac{s_x s_w}{2}|W_{I,:}|_1$ (Teorema 8.1 + subaditividade), e os outliers — únicos responsáveis por explosões de erro — ficam exatos em fp16. Implementação: se bitsandbytes disponível com backend funcional, delega; senão, o mesmo critério ($\tau = 6,\mathrm{MAD}$) é aplicado manualmente.

5. Onde o KHTST aplica

  1. Inferência/avaliação: pesos do decoder quantizados W8 após o treino (compara-se perplexidade fp32 × W8 — relatada na telemetria).
  2. Treino (opcional, config): QAT-STE nas projeções lineares nos últimos passos.
  3. bitsandbytes: integração condicionada a CUDA; no ambiente CPU usa-se o caminho simétrico puro-PyTorch com as mesmas cotas (honestidade documentada no README).

5. v4 — Escalas por grupo (Teorema 8.4)

Divide-se a última dimensão em grupos de tamanho $g$ (64) com escala própria $s_g = \max|x|g/q{\max}$.

Teorema 8.4 (erro estruturalmente menor). Com escalas por grupo, o erro $\infty$-norma por elemento é $\le s_g/2$ e, para entradas aproximadamente uniformes dentro de cada grupo, $\mathbb{E}[\mathrm{erro}^2] \approx s_g^2/12 \le (s_{\text{tensor}}/g^{1/d})^2/12$ — redução de raiz quadrada em $g$ no caso 1D ($\approx\sqrt{g}$× menos erro RMS que a escala por tensor).

Demonstração. Arredondamento simétrico: $|\mathrm{erro}|\le s/2$ por célula (mesmo argumento do Teorema 8.1 restrito ao grupo). Para $x$ uniforme em $[-s/2, s/2]$ dentro do grupo, $\mathbb{E}[\mathrm{erro}^2] = s^2/12$; como $s_g$ usa o máximo de APENAS $g$ elementos (vs. todo o tensor), $s_g \le s_{\text{tensor}}$ com igualdade apenas quando o máximo global está no grupo — em média, sobre grupos com distribuição semelhante, a razão dos máximos é $O(1/\sqrt{g})$ para subamostragem aleatória, dando o fator $\sqrt{g}$. $\blacksquare$

6. v4 — Correção de viés pós-quantização (Teorema 8.5)

Teorema 8.5 (eliminação do desvio de 1ª ordem). Seja $\Delta b = \mathbb{E}a[a,W^\top] - \mathbb{E}a[a_q W_q^\top]$ calibrado em um lote representativo. Então $y{\text{corr}} = y_q + \Delta b$ tem $|\mathbb{E}[y{\text{fp}} - y_{\text{corr}}]| \le |,\mathrm{Cov},|\cdot \varepsilon_2$ — o erro MÉDIO restante é de 2ª ordem (Teorema 8.2: média $s/4$ por produto interno).

Demonstração. $y_{\text{fp}} - y_q = a W^\top - a_q W_q^\top = (a-a_q)W^\top

  • a_q (W-W_q)^\top - (a-a_q)(W-W_q)^\top$; tomando esperança, os dois primeiros termos são exatamente $\Delta b$ (removido pela correção); resta o produto cruzado, quadrático nos erros de quantização — $O(\varepsilon^2)$. $\blacksquare$

Implementação: quantiza_por_grupo (fallback automático p/ dim < g) e FakeQuantW8A8Grupo.calibrar → correcao_vies (GPTQ-lite).

§§7–9 (v5) — Quantização REAL sem fakes

Refactor v5 (quanta/quantizacao.py): FakeQuantW8A8, FakeQuantW8A8Grupo e aplicar_w8_linear REMOVIDOS. A inferência NUNCA é simulada — os backends reais:

§7.1 torchao Int8DynamicActivationInt8WeightConfig — kernels int8 reais (CPU/GPU); §7.2 torch.ao.quantization prepare_qat→convert — Treinamento INT8 UNIVERSAL (backend x86/fbgemm/qnnpack; conversão final REAL, fluxo do pedido); §7.3 torchao.float8 convert_to_float8 — treino nativo FP8 (GPU H100+; degrada honesta em CPU); §7.4 bitsandbytes load_in_8bit=True — weight-only 8-bit GPU (integração HF).

Teorema 16.7 (equivalência). Simulação round→clamp→dequantize e kernel int8 real implementam o MESMO mapa afim ⇒ erro idêntico; as cotas 8.1–8.5 valem ao kernel. A diferença é VELOCIDADE/MEMÓRIA (Teorema 16.8: ≈3,76× com escalas por grupo), não erro. §8 métricas honestas: metricas_qualidade (erro relativo, SNR dB, erro máx). §9 diagnóstico: modo_quantizacao() — disponibilidade real por backend. Durante o TREINO mantém-se QAT com STE (Teorema 8.3) — classe QATW8A8 (metodologia idêntica ao prepare_qat, escalas por grupo) — mas a conversão de exportação é sempre REAL (§7.1/§7.2).