Outlier Suppression+ / SpAtten#

Outlier Suppression+ (OS+) est un framework qui élimine les outliers d'activation par shifting et scaling équivalents par canal — des transformations mathématiques qui naltèrent pas la sortie du modèle. SpAtten est un co-design algorithm-architecture qui sparsifie l'attention en éliminant progressivement les tokens et têtes peu importants. Ensemble, ils repoussent les limites de la quantification des transformers.


Le problème des outliers dans les activations#

Pour un néophyte#

Les activations sont les valeurs qui circulent à l'intérieur du modèle pendant l'inférence. Dans les transformers, certaines activations développent des valeurs extrêmes (outliers) — des pics qui sont 20× à 100× plus grands que la moyenne.

flowchart LR subgraph NORMAL["Niveau normal"] G["Guitare"] --> N1["||||||||||||"] B["Basse"] --> N2["||||||||||||"] D["Batterie"] --> N3["||||||||||||"] T["Trompette"] --> N4["||||||||||||"] end OUT["⚡ DIDJERIDOO
████████████████████████
← OUTLIER !!!"]

Si tu enregistres en « basse qualité » (quantification), le didjéridoo sature TOUT l'enregistrement. On n'entend plus les autres instruments.

Solution OS+ : atténuer le didjéridoo AVANT l'enregistrement, puis le ré-amplifier APRÈS (transformation équivalente).

Le problème technique détaillé#

Distribution des activations : la majorité des valeurs est concentrée entre -1 et +1, mais des outliers systématiques apparaissent à 5, 10, 20, 50, voire 100. Ces 0.1% des valeurs dominent la dynamique.

Quantification INT8 : tout doit tenir entre -128 et +127. Les outliers à 100 forcent un scale = 100/127 ≈ 0.79. Toutes les valeurs « normales » (entre -1 et +1) sont quantifiées à 0 ou ±1 → perte massive d'information !

Sans traitement des outliers : Scale = max|X| / 127 = 100/127 ≈ 0.79

Valeur Quantifiée Erreur
0.03 0 0.03 ❌
0.12 0 0.12 ❌
0.50 1 0.50 ❌
100.0 127 0.21 ✓

Les valeurs normales sont DÉTRUITES. L'outlier est bien représenté mais tout le reste est perdu.


Outlier Suppression+ : la solution#

Les deux opérations équivalentes#

OS+ applique deux transformations mathématiquement équivalentes (elles ne changent pas la sortie du modèle) pour éliminer les outliers :

flowchart TD subgraph SHIFT["1. Channel-wise Shifting (alignement du centre)"] SH["Pour chaque canal j : X'_j = X_j - μ_j
où μ_j = moyenne du canal j
→ Élimine l'asymétrie, centre chaque canal autour de zéro
La soustraction est compensée en ajoutant μ_j au biais de la couche suivante"] end subgraph SCALE["2. Channel-wise Scaling (équilibrage)"] SC["Pour chaque canal j : X''_j = X'_j / s_j
où s_j = écart-type (ou max) du canal j
→ Équilibre la dynamique entre canaux
Les canaux avec outliers sont divisés par s_j grand
La division est compensée en multipliant les poids de la couche suivante par s_j"] end

Diagramme de la distribution avant/après suppression#

flowchart LR subgraph AVANT["AVANT OS+"] direction TB BA1["Canal 1: ||||||||||||||||||||||||| (normal)"] BA2["Canal 2: ||||||||||||||||||||||||| (normal)"] BA3["Canal 3: ||||||||||||||||||||||||||| ████ (outlier !)"] BA4["Canal 4: ||||||||||||||||||||||||| (normal)"] BA5["Dynamique globale: [-2, +100]
← dominée par Canal 3"] BA6["Scale = 100/127 = 0.79
→ Canaux 1, 2, 4: quasi-tout à 0 ❌
→ Canal 3: bien représenté ✓"] end subgraph APRES["APRÈS OS+ (shift + scale par canal)"] direction TB AP1["Canal 1: |||||||||||||||||||||| (centré, normalisé)"] AP2["Canal 2: |||||||||||||||||||||| (centré, normalisé)"] AP3["Canal 3: |||||||||||||||||||||| (outlier ATTÉNUÉ !)"] AP4["Canal 4: |||||||||||||||||||||| (centré, normalisé)"] AP5["Dynamique globale: [-4, +4] ← équilibrée !"] AP6["Scale = 4/127 = 0.031
→ Tous les canaux bien représentés ✓
→ Résolution fine partout ✓"] end

AVANT (distribution brute avec outliers) : un seul canal outlier atteint l'amplitude 100, tandis que tous les autres canaux sont concentrés autour de 0. La distribution est totalement déséquilibrée.

APRÈS (shifting + scaling par canal) : la distribution devient uniforme — tous les canaux ont une amplitude similaire entre -4 et +4. → Distribution UNIFORME → quantification efficace !

Pourquoi c'est "équivalent" (sans perte)#

Couche linéaire : Y = XW + b

OS+ shifting : X' = X - μ (par canal). Y = (X - μ)W + b = XW + (b - μW) = XW + b'b' = b - μW.
→ La sortie Y est IDENTIQUE, seul le biais change. C'est une transformation ÉQUIVALENTE (lossless).

OS+ scaling : X'' = X' / s (par canal). Y = (X'/s)W + b' = X'(W/s) + b'W' = W/s.
→ La sortie Y est IDENTIQUE, seuls les poids changent. Transformation ÉQUIVALENTE.

Le modèle MATHÉMATIQUE ne change pas, mais la DISTRIBUTION des activations devient quantification-friendly.


SpAtten : attention sparsifiée#

SpAtten est un co-design algorithm-architecture qui réduit le coût du mécanisme d'attention en éliminant les tokens et têtes peu importants.

Les trois techniques de SpAtten#

flowchart TD subgraph T1["1. Cascade Token Pruning"] T1A["Élimination progressive des tokens peu importants
Critère: score d'attention cumulé bas → token inutile"] T1B["Couche 1: [T1 T2 T3 T4 T5 T6 T7 T8]
Couche 3: [T1 T2 — T4 T5 — T7 T8]
Couche 6: [T1 — T4 T5 — — T8]
Couche 12: [T1 — T4 — — — — T8] ← très sparse"] end subgraph T2["2. Head Pruning"] T2A["Élimination des têtes d'attention redondantes
Head 1: gardée (importante)
Head 2: prunée (peu utile)
Head 3: gardée (importante)
Head 4: prunée (redondante avec Head 1)
Head 5: gardée (importante)"] end subgraph T3["3. Quantification adaptative"] T3A["Précision variable selon l'importance
Token important → INT8 (haute précision)
Token moyen → INT4
Token faible → INT2 (ou pruné)"] end

Diagramme : attention dense vs sparsifiée#

flowchart LR subgraph DENSE["Attention DENSE (standard)"] direction TB DA["Q · Kᵀ = Scores d'attention (N × N)
Tous les tokens interagissent
Coût: O(N²)"] end subgraph SPARSE["Attention SPARSIFIÉE (SpAtten)"] direction TB SA["Tokens prunés (cases vides)
Coût réduit: O(N × k), k << N
Moins de calculs, moins de mémoire, même qualité"] end

Niveau néophyte : l'analogie complète#

Imagine que tu lis un livre de 500 pages. L'attention dense (standard), c'est lire chaque mot avec la même concentration. SpAtten, c'est :
1. Skim les passages peu importants (token pruning)
2. Sauter les chapitres redondants (head pruning)
3. Lire attentivement les passages clés et diagonalement le reste (quantification adaptative)

Le résultat : tu comprends l'histoire aussi bien, mais en 3× moins de temps.


OS+ vs les autres méthodes de gestion des outliers#

mindmap root((Méthodes de gestion
des outliers)) LLM.int8() Mixed-precision Sépare physiquement outliers du reste INT8 + FP16 SmoothQuant Migration fixe α Migrer difficulté activations → poids INT8 + INT8 OS+ Shifting + scaling équivalents par canal Apprenable INT8/INT6/INT4 (tout en INT) Complémentaire: peut être combiné OmniQuant LWC + LET apprenables W2-W6, A4-A16

Caractéristiques#

Propriété OS+ SpAtten
Cible Activations (outliers) Attention (tokens + heads)
Approche Shifting + scaling équivalents Pruning + quantification adaptative
Bits cibles INT8, INT6, INT4 Variable (adaptatif)
Équivalent ✅ Oui (lossless mathématiquement) Partiellement (pruning = perte contrôlée)
Modèles validés BERT, OPT, BLOOM, BLOOMZ, LLaMA Modèles transformer généraux
Plug-and-play ✅ Oui Co-design hardware/software
Backprop requise Non (transformation offline) Calibration nécessaire

Résultats#

Outlier Suppression+#

xychart-beta title "BERT-base SQuAD F1 — INT8 (plus haut = mieux)" x-axis ["FP32", "OS+ INT8", "Baseline INT8"] y-axis "F1 Score" 85 --> 89 bar [88.5, 88.4, 86.2]
xychart-beta title "BERT-base SQuAD F1 — INT4 (plus haut = mieux)" x-axis ["OS+ INT4", "Baseline INT4"] y-axis "F1 Score" 55 --> 80 bar [76.3, 60.8]

OS+ INT4 : +15.5% d'accuracy par rapport à la baseline INT4 !

Tableau récapitulatif#

Modèle Configuration Résultat OS+ Baseline (sans OS+) Gain
BERT-base INT8 (SQuAD) 88.4% F1 86.2% F1 +2.2%
BERT-base INT4 (SQuAD) 76.3% F1 60.8% F1 ⭐ +15.5%
OPT-1.3B INT8 (PPL) ~FP16 dégradé near-FP
BLOOM-7B INT8 (PPL) ~FP16 dégradé near-FP
BLOOMZ-7B INT6 (PPL) ~FP16 cassé ⭐ viable
LLaMA-7B INT8 (PPL) ~FP16 dégradé near-FP

Résultat marquant : OS+ permet une quantification INT4 de BERT avec +15.5% d'accuracy par rapport à la baseline — rendant le 4-bit enfin viable pour BERT.


Exemple pratique#

Outlier Suppression+#

git clone https://github.com/ModelTC/Outlier_Suppression_Plus.git
cd Outlier_Suppression_Plus
pip install -e .

Quantification BERT INT8 avec OS+#

python run_quant.py \
    --model bert-base-uncased \
    --task squad \
    --bits 8 \
    --suppression plus \
    --calibration_data squad \
    --nsamples 128 \
    --output_dir ./bert-int8-osplus

Quantification BERT INT4 avec OS+#

python run_quant.py \
    --model bert-base-uncased \
    --task squad \
    --bits 4 \
    --suppression plus \
    --calibration_data squad \
    --nsamples 256 \
    --output_dir ./bert-int4-osplus

Inférence#

from outlier_suppression import OSPlusModel
from transformers import AutoTokenizer

model = OSPlusModel.from_quantized("./bert-int8-osplus")
tokenizer = AutoTokenizer.from_pretrained("./bert-int8-osplus")

# Q&A
question = "Qu'est-ce que l'outlier suppression ?"
context = "L'outlier suppression est une technique qui élimine les valeurs extrêmes..."

inputs = tokenizer(question, context, return_tensors="pt")
outputs = model(**inputs)
answer = tokenizer.decode(outputs.start_logits.argmax())
print(f"Réponse : {answer}")

SpAtten#

git clone https://github.com/mit-han-lab/spatten.git
cd spatten
pip install -e .
from spatten import SpAttenConfig, SpAttenModel
from transformers import AutoTokenizer

# Configuration avec pruning
config = SpAttenConfig(
    token_pruning=True,           # cascade token pruning
    head_pruning=True,            # head pruning
    adaptive_quant=True,          # quantification adaptative
    token_prune_ratio=0.3,        # 30% de tokens prunés par couche
    head_prune_ratio=0.2,         # 20% de têtes prunées
)

model = SpAttenModel.from_pretrained("bert-base-uncased", config=config)
tokenizer = AutoTokenizer.from_pretrained("bert-base-uncased")

# Inférence avec attention sparsifiée
inputs = tokenizer("Sparsifier l'attention réduit le coût computationnel.", return_tensors="pt")
outputs = model(**inputs)

OS+ en combinaison avec d'autres méthodes#

flowchart LR M["Modèle FP16"] --> OS["OS+
(shift + scale équivalents)"] --> Q["Quantif. INT8/INT4
GEMM pure"] Q --> COMB["Peut être combiné avec :
• SmoothQuant (migration supplémentaire)
• GPTQ/AWQ (quantification des poids)
• OmniQuant (paramètres apprenables)"]

OS+ prépare les activations, les autres méthodes gèrent les poids. Effet MULTIPLICATIF.


Avantages et inconvénients#

✅ Avantages#

  • Near-floating-point en INT8 et INT6 (perte négligeable)
  • INT4 enfin viable pour BERT (+15.5% vs baseline)
  • Transformation équivalente : pas de changement mathématique
  • Plug-and-play : se combine avec d'autres méthodes
  • Validé sur de nombreux modèles : BERT, OPT, BLOOM, BLOOMZ, LLaMA
  • SpAtten : réduction drastique du coût d'attention

⚠️ Inconvénients#

  • OS+ : principalement validé sur des modèles plus petits (BERT, OPT)
  • OS+ : moins éprouvé sur les très grands LLMs modernes (70B+)
  • SpAtten : co-design hardware/software complexe
  • SpAtten : le pruning introduit une perte contrôlée (pas lossless)
  • Support limité dans les frameworks mainstream de serving
  • Écosystème moins mature que GPTQ/AWQ/bitsandbytes

Références#

ia llm quantification outlier-suppression spatten activation-quantization shifting scaling pruning sparse-attention int8 int4