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 :
- PTQ : Tu apprends à écrire normalement avec un stylo fin, puis on te donne un stylo épais. Tes lettres seront un peu déformées.
- 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.
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#
- Insertion de nœuds de fake quantization dans le graphe de calcul (avant chaque couche linéaire/conv)
- Forward : les valeurs sont arrondies et clampées → le modèle "voit" l'erreur de quantification
- Backward : les gradients sont calculés en pleine précision (FP32)
- 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).
⚠️ 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
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#
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 :
à 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 ?#
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#
- 📄 LLM-QAT (arXiv:2305.17888) — Papier original LLM-QAT data-free
- 🔧 Repo GitHub LLM-QAT — Implémentation officielle Meta/FAIR
- 📄 Quantization and Training of Neural Networks (Bengio et al.) — STE originel
- 📦 PyTorch Quantization — Documentation officielle QAT
- 🔗 Voir aussi : [[ptq]] — Post-Training Quantization (comparaison directe)
- 🔗 Voir aussi : [[llm-int8]] — LLM.int8() (alternative PTQ INT8)
- 🔗 Voir aussi : [[index]] — Vue d'ensemble des méthodes de quantification