Quantization Aware Training (QAT)#

TL;DR#

Le QAT (Quantization Aware Training) consiste à simuler les effets de la quantification pendant l'entraînement afin que le modèle apprenne à compenser les erreurs d'arrondi. C'est la méthode qui produit les modèles quantifiés les plus précis, surtout à bas bitrate (≤ 4 bits), mais aussi la plus coûteuse.


Pour un néophyte#

Imagine que tu dois apprendre à écrire avec un stylo qui a une pointe très épaisse. Deux stratégies :

  1. PTQ : Tu apprends à écrire normalement avec un stylo fin, puis on te donne un stylo épais. Tes lettres seront un peu déformées.
  2. QAT : Tu t'entraînes directement avec le stylo épais. Ton cerveau compense et tes lettres restent belles.

Le QAT, c'est ça : on entraîne le modèle en lui montrant ce que ça donne d'être quantifié, pour qu'il s'adapte.


Concept fondamental : Fake Quantization#

Le cœur du QAT est la fake quantization (quantification simulée). Pendant le forward pass, les poids et activations sont simulés comme s'ils étaient quantifiés (arrondis aux niveaux discrets INT4/INT8), mais restent stockés en FP16/FP32.

flowchart TD A["Poids FP32"] --> B["Fake Quantize
Arrondi + clip"] B --> C["Poids INT8 simulé"] C --> D["Calcul Forward"] D --> E["Loss computation"] E --> F["Backward pass
(gradients en FP32)"] B -.->|"STE"| F

Comment ça marche concrètement#

  1. Insertion de nœuds de fake quantization dans le graphe de calcul (avant chaque couche linéaire/conv)
  2. Forward : les valeurs sont arrondies et clampées → le modèle "voit" l'erreur de quantification
  3. Backward : les gradients sont calculés en pleine précision (FP32)
  4. Mise à jour : les poids FP32 sont mis à jour pour minimiser la loss en tenant compte de la quantification simulée
Fake Quantize (pour un tenseur x) :

  x_quant = round(x / scale) × scale

  puis clampé entre [min_quant, max_quant]

  Résultat : valeur "comme si" elle était en INT8, mais en float

Straight-Through Estimator (STE)#

Le problème : la fonction round() n'est pas différentiable. Son gradient est nul presque partout (et infini aux demi-entiers).

flowchart LR subgraph R["round(x) — non-différentiable"] direction LR X["x"] -->|"step"| Y["round(x)"] end

⚠️ Le gradient de round() est 0 partout → impossible de backpropagater !

Solution : le Straight-Through Estimator (STE). On "triche" :
- En forward : on utilise round() (la vraie fonction)
- En backward : on fait comme si round() était l'identité, c'est-à-dire que le gradient passe tel quel : ∂loss/∂x ≈ ∂loss/∂x_quant

flowchart LR subgraph FWD["Forward"] direction LR A1["x"] --> A2["x_quant = round(x)
vraie quantification"] end subgraph BWD["Backward"] direction LR B1["∂L/∂x_quant"] --> B2["∂L/∂x = ∂L/∂x_quant
gradient passe-through"] end

"On fait semblant que round() est l'identité pour les gradients"

C'est une approximation grossière mais étonnamment efficace. Le modèle apprend à placer ses poids dans des zones où l'erreur de quantification est minimale.


Processus complet du QAT#

flowchart TD A["1. MODÈLE PRÉ-ENTRAÎNÉ (FP16/FP32)"] --> B["2. INSERTION DES NŒUDS
DE FAKE QUANTIZATION"] B --> B1["• Observer les distributions
• Calculer scale factors et zero-points"] B1 --> C["3. FINE-TUNING AVEC FAKE QUANT"] C --> C1["• Forward : valeurs simulées quantifiées
• Backward : gradients via STE en FP32
• N époques"] C1 --> D["4. EXPORT DU MODÈLE QUANTIFIÉ"] D --> D1["• Poids réellement arrondis en INT4/INT8
• Prêt pour l'inférence accélérée"]

LLM-QAT : la variante Data-Free#

Le défi du QAT pour les LLMs : les données d'entraînement originales sont souvent inaccessibles (modèles propriétaires, données massives, raisons de confidentialité).

LLM-QAT (Meta/FAIR, 2023) résout ce problème avec une approche data-free :

flowchart TD A["1. GÉNÉRATION DE DONNÉES
à partir du modèle"] --> A1["• Le modèle FP16 génère du texte
(self-generation)
• Proxy pour la calibration"] A1 --> B["2. DISTILLATION DE CONNAISSANCE"] B --> B1["• Teacher = modèle original FP16
• Student = modèle QAT quantifié
• Loss = KL divergence"] B1 --> C["3. QUANTIFICATION W4A4KV4"] C --> C1["• Poids 4-bit + Activations 4-bit
+ KV cache 4-bit
• KV cache 4-bit crucial pour le throughput"]

Innovation clé : LLM-QAT quantifie non seulement les poids et activations, mais aussi le KV cache, ce qui réduit drastiquement la mémoire lors de la génération de longues séquences.


Comparaison QAT vs PTQ#

Critère QAT PTQ
Précision (≤ 4 bits) ✅ Meilleure — compense les erreurs ❌ Dégradation notable
Précision (8 bits) ≈ Équivalente ✅ Suffisante
Coût computationnel ❌ Élevé (entraînement nécessaire) ✅ Faible (minutes/heures)
Temps total Heures → jours (selon taille modèle) Minutes → heures
Données requises Données d'entraînement (ou génération data-free) Dataset de calibration (128-1024 échantillons)
Simplicité ❌ Complexe (STE, fake quant, hyperparams) ✅ Simple et plug-and-play
Récupération d'erreur ✅ Le modèle apprend à compenser ❌ Pas de récupération
Cas d'usage idéal Production avec contraintes de bits agressives Prototypage, INT8, déploiement rapide
Outils PyTorch native, LLM-QAT (Meta) bitsandbytes, AutoGPTQ, AutoAWQ

Quand choisir QAT vs PTQ ?#

flowchart TD A["Quelle précision cible ?"] --> B["≤ 4 bits"] A --> C["≥ 8 bits"] B --> D["Budget entraînement ?"] D -->|OUI| E["✅ QAT"] D -->|NON| F["PTQ + accepter
la dégradation"] C --> G["✅ PTQ suffit
(GPTQ, AWQ, bnb)"]

Bits cibles#

  • W4A4 (poids 4-bit, activations 4-bit) avec KV cache 4-bit — configuration agressive de LLM-QAT
  • W8A8 (poids 8-bit, activations 8-bit) — configuration conservatrice

Avantages#

  • ✅ Meilleure précision à bas bitrate que toutes les méthodes PTQ
  • ✅ Le modèle apprend activement à compenser les erreurs d'arrondi
  • ✅ Méthode data-free (LLM-QAT) : aucune donnée d'entraînement nécessaire
  • ✅ Quantifie le KV cache en plus des poids/activations
  • ✅ Préserve la distribution de sortie via distillation

Inconvénients#

  • ❌ Coût computationnel élevé (entraînement complet ou fine-tuning)
  • ❌ Complexité de mise en œuvre (STE, fake quantization, hyperparamètres)
  • ❌ Temps de traitement important pour les grands modèles
  • ❌ Nécessite des GPU avec suffisamment de mémoire pour l'entraînement
  • ❌ Le STE reste une approximation — convergence parfois instable

Exemple pratique#

Avec PyTorch natif (QAT pour un modèle standard)#

import torch
import torch.ao.quantization as quant

# 1. Préparer le modèle pour le QAT
model_fp32 = MyModel()
model_fp32.eval()

# Configuration QAT
model_fp32.qconfig = quant.get_default_qat_qconfig('fbgemm')

# 2. Insertion des nœuds de fake quantization
model_train = quant.prepare_qat(model_fp32.train())

# 3. Fine-tuning avec fake quantization
optimizer = torch.optim.Adam(model_train.parameters(), lr=1e-5)
for epoch in range(num_epochs):
    for batch in train_dataloader:
        outputs = model_train(batch)
        loss = criterion(outputs, batch.labels)
        loss.backward()  # STE géré automatiquement par PyTorch
        optimizer.step()
        optimizer.zero_grad()

# 4. Conversion vers le modèle réellement quantifié
model_int8 = model_train.to(torch.device('cpu'))
model_int8 = quant.convert(model_int8)
# → Prêt pour l'inférence INT8

Avec LLM-QAT (Meta/FAIR) — data-free#

# Clone du repo officiel
git clone https://github.com/facebookresearch/LLM-QAT
cd LLM-QAT

# Étape 1 : génération de données à partir du modèle
python generate_data.py \
    --model-path /path/to/llama-7b \
    --num-samples 100000 \
    --output-dir ./calibration_data

# Étape 2 : calibration pour déterminer les range des activations
python calibrate.py \
    --model-path /path/to/llama-7b \
    --data-path ./calibration_data \
    --output ./calibration_results

# Étape 3 : fine-tuning QAT avec distillation data-free
python train.py \
    --model-path /path/to/llama-7b \
    --calibration ./calibration_results \
    --bits-weight 4 \
    --bits-activation 4 \
    --bits-kv-cache 4 \
    --epochs 5 \
    --output ./llama-7b-w4a4kv4

Avec NVIDIA TensorRT Model Optimizer#

from modelopt.torch.quantization import quantize
from modelopt.torch.quantization.qat import QATConfig

# Configuration QAT pour TensorRT
config = QATConfig(
    weight_bits=8,
    activation_bits=8,
    calibration_method="entropy",
)

model_qat = quantize(model_fp32, config)
# Fine-tuning standard puis export ONNX/TensorRT

Papier arXiv source#

LLM-QAT: Data-Free Quantization Aware Training for Large Language Models
- 📄 arXiv : 2305.17888
- 👤 Auteurs : Zechun Liu, Barlas Oguz, Changsheng Zhao, Ernie Chang, Pierre Stock, Yashar Mehdad, Yangyang Shi, Raghuraman Krishnamoorthi, Vikas Chandra
- 🏢 Meta/FAIR
- 📅 29 mai 2023


Références#

ia llm quantification qat fine-tuning training int4 int8 fake-quantization