AI

Accélérer les LLM : Décodage spéculatif et optimisation du cache KV pour une inférence sous 100 ms

Introduction

Dans le paysage en évolution rapide des grands modèles de langage (LLM), la latence d'inférence est souvent le goulot d'étranglement qui empêche les applications en temps réel, telles que les assistants de code interactifs ou les chatbots à faible latence, d'atteindre leur plein potentiel. Bien que la quantification des modèles réduise l'empreinte mémoire, elle a souvent du mal à fournir les accélérations drastiques nécessaires pour des temps de réponse inférieurs à 100 ms. C'est ici que la combinaison du décodage spéculatif et de l'optimisation du cache KV brille. En prédisant intelligemment les séquences de jetons et en gérant efficacement l'état de la mémoire, les développeurs peuvent obtenir des accélérations linéaires sans sacrifier la qualité de la sortie du modèle.

Le cache KV : le héros méconnu de la performance

Pour comprendre comment nous atteignons une inférence sous 100 ms, nous devons d'abord aborder le cache Clé-Valeur (KV). Lors de la génération de texte, le modèle Transformer porte son attention sur les jetons précédents. Au lieu de recalculer l'attention pour toute la séquence à chaque étape, le modèle met en cache les vecteurs Clé et Valeur des jetons précédents. Pour les applications à contexte long, ce cache peut devenir un goulot d'étranglement mémoire. S'il n'est pas géré efficacement, la surcharge liée à l'allocation et à la gestion de ce cache peut annuler les avantages du traitement parallèle.

Stratégie d'optimisation clé : Mettre en œuvre le regroupement continu (également connu sous le nom de regroupement conscient du planificateur). Contrairement au regroupement statique, le regroupement continu permet d'insérer de nouvelles requêtes dans le lot dès qu'une requête précédente se termine, maximisant ainsi l'utilisation du GPU et maintenant le cache KV compact et pertinent.

# Pseudo-code illustrant la gestion efficace du cache KV
def generate_with_kv_cache(model, prompt, max_length):
    # Initialiser le cache KV avec une mémoire pré-allouée
    kv_cache = model.init_cache(max_length)
    
    # Encoder le prompt et calculer les états KV initiaux
    inputs = tokenizer(prompt, return_tensors="pt")
    outputs = model(input_ids=inputs["input_ids"], past_key_values=kv_cache)
    
    next_token = outputs.logits[:, -1, :].argmax(dim=-1)
    
    for _ in range(max_length):
        # Mettre à jour le cache avec uniquement les paires KV du nouveau jeton
        # Cela évite de recalculer toute la matrice d'attention
        kv_cache = update_kv_cache(kv_cache, outputs)
        
        # Générer le jeton suivant
        outputs = model(input_ids=next_token, past_key_values=kv_cache)
        next_token = outputs.logits[:, -1, :].argmax(dim=-1)
        
        if next_token == tokenizer.eos_token_id:
            break
            
    return tokenizer.decode(outputs)

Décodage spéculatif : paralléliser le séquentiel

Le décodage autoregressif traditionnel est intrinsèquement séquentiel : le jeton $T_n$ dépend de $T_{n-1}$. Cela crée un chemin critique qui limite la vitesse. Le décodage spéculatif brise cette barrière en utilisant un modèle "brouillon" plus petit et plus rapide pour proposer plusieurs jetons, qui sont ensuite vérifiés en parallèle par le modèle "cible" plus grand et autoritaire. Le processus fonctionne comme suit : 1. Le modèle brouillon génère $N$ jetons candidats. 2. Le modèle cible traite le prompt ainsi que ces candidats simultanément. 3. Le modèle cible vérifie les candidats par rapport à sa distribution de probabilité. 4. Tout jeton non conforme entraîne un retour en arrière, et la génération reprend à partir du dernier jeton vérifié. Lorsque le modèle brouillon est précis, le modèle cible valide plusieurs jetons en un seul passage avant, ce qui conduit à des accélérations linéaires.

Exemple d'implémentation pratique

L'utilisation de bibliothèques modernes comme Hugging Face Transformers et vLLM rend la mise en œuvre du décodage spéculatif de plus en plus accessible. Voici une implémentation conceptuelle utilisant un modèle brouillon (par exemple, un modèle distillé de 7 milliards de paramètres) pour accélérer un modèle cible (par exemple, un modèle de 70 milliards de paramètres).
from transformers import AutoModelForCausalLM, AutoTokenizer
import torch

# Charger le modèle brouillon léger et le modèle cible lourd
draft_model = AutoModelForCausalLM.from_pretrained("draft-model-7b")
target_model = AutoModelForCausalLM.from_pretrained("target-model-70b")

draft_tokenizer = AutoTokenizer.from_pretrained("draft-model-7b")

def speculative_decode(prompt, draft_model, target_model, num_drafts=4):
    # 1. Phase de brouillon
    draft_inputs = draft_tokenizer(prompt, return_tensors="pt")
    draft_outputs = draft_model.generate(**draft_inputs, max_new_tokens=num_drafts)
    draft_tokens = draft_outputs[0]
    
    # 2. Phase de vérification
    # Le modèle cible traite le prompt + les jetons brouillon en une seule fois
    target_inputs = target_tokenizer(prompt + draft_tokenizer.decode(draft_tokens), 
                                     return_tensors="pt")
    
    with torch.no_grad():
        target_outputs = target_model(**target_inputs)
    
    # 3. Logique d'acceptation/rejet
    # Comparer les probabilités du brouillon avec les logits du cible
    # Si accepté, ajouter à la séquence. Si rejeté, tronquer.
    accepted_tokens = verify_acceptance(draft_tokens, target_outputs)
    
    return accepted_tokens

Combinaison des techniques pour atteindre l'objectif sous 100 ms

Pour atteindre constamment des temps d'inférence inférieurs à 100 ms, vous ne pouvez pas vous fier à une seule technique. Vous devez les combiner :
  • Accélération matérielle : Utilisez des noyaux CUDA optimisés pour le cache KV (comme FlashAttention) pour réduire l'utilisation de la bande passante mémoire.
  • Distillation de modèle : Entraînez un modèle brouillon plus petit spécifiquement pour imiter la distribution de sortie du modèle plus grand, augmentant ainsi le taux d'acceptation dans le décodage spéculatif.
  • Efficacité du regroupement : Comme mentionné, utilisez le regroupement continu pour garantir que le GPU traite toujours des données, minimisant ainsi les temps d'inactivité.

Conclusion

Atteindre une inférence LLM sous 100 ms n'est plus un exercice théorique, mais un défi d'ingénierie pratique. En tirant parti de l'efficacité mémoire de l'optimisation du cache KV et de la puissance de traitement parallèle du décodage spéculatif, les développeurs peuvent déployer des modèles puissants qui répondent avec une vitesse proche de celle des humains. À mesure que les écosystèmes matériels et logiciels continuent de maturer, ces techniques deviendront des pratiques standard pour toute application IA de niveau production. Commencez par profiler les goulets d'étranglement de latence actuels, mettez en œuvre un cache KV efficace et expérimentez avec des modèles brouillons légers pour débloquer le véritable potentiel de vos LLM.
Share: