La recherche en IA te passionne ?
Les papers et avancées qui comptent, expliqués simplement, chaque soir. Gratuit.
Inclus dès l'inscription : notre sélection des meilleurs guides & comparatifs IA.
Choisis ton rythme
Gratuit · Pas de spam · Désabonnement en 1 clic
Transformer : comprendre l'entraînement et l'inférence
Utilisation d'un modèle Transformer : de l'entraînement à l'inférence
Si vous avez implémenté un modèle transformer dans PyTorch, vous pouvez utiliser le même code pour l'entraînement et l'inférence, mais de manières très différentes. Pendant l'entraînement, vous traitez généralement un lot de séquences de tokens de longueur fixe et mettez à jour les poids du modèle. Pendant l'inférence, les poids sont fixes et le modèle génère de nouveaux tokens un à un.
Cette différence change presque tout en matière de performance. L'entraînement est dominé par de grandes multiplications de matrices et le passage arrière. L'inférence est dominée par des passes avant répétées, le mouvement de la mémoire et la nécessité de garder les clés et valeurs d'attention précédentes disponibles pour le token suivant.
Dans ce chapitre, vous apprendrez à propos de :
- La boucle de génération autoregressive
- La différence entre pré-remplissage et décodage
- Pourquoi la mise en cache des clés et valeurs est nécessaire
- Comment implémenter un cache KV simple
- Comment raisonner sur la mémoire utilisée par le cache
Boucle de génération autoregressive
Un modèle transformer uniquement décodeur prédit le prochain token à partir des tokens qui le précèdent. L'exigence stricte d'utiliser uniquement les tokens précédents est imposée par le mécanisme d'attention causale. Si les tokens d'entrée sont :
The cat sat on the
Le modèle renvoie une distribution de probabilité sur le vocabulaire pour le prochain token. Un token suivant probable peut être "mat", mais le modèle ne renvoie pas directement un mot. Il renvoie des logits, qui sont des scores non normalisés pour chaque token dans le vocabulaire.
La boucle de génération est donc simple :
- Tokeniser le prompt.
- Exécuter le modèle pour obtenir les logits pour le prochain token.
- Choisir un token parmi les logits.
- Ajouter ce token à l'entrée.
- Répéter jusqu'à ce qu'une règle d'arrêt soit atteinte.
Ceci est appelé génération autoregressive car chaque nouveau token dépend des tokens générés précédemment. Le modèle ne peut pas générer le dixième token de sortie avant de connaître les neuf premiers tokens de sortie.
Une très petite boucle de décodage glouton peut être écrite comme suit :
import torch
@torch.no_grad()
def greedy_decode(model, input_ids, max_new_tokens):
output_ids = input_ids.clone()
for _ in range(max_new_tokens):
logits = model(output_ids)
next_token_logits = logits[:, -1, :]
next_token = next_token_logits.argmax(dim=-1, keepdim=True)
output_ids = torch.cat([output_ids, next_token], dim=1)
return output_ids
Dans le code ci-dessus, model est un modèle PyTorch, max_new_tokens est un entier positif, et toutes les autres variables sont des tenseurs PyTorch. La boucle for itère max_new_tokens fois, et à chaque itération, elle renvoie l'ensemble de la séquence au modèle pour obtenir les logits pour le prochain token. La fonction argmax() sélectionne le token ayant le score le plus élevé. La fonction cat() est utilisée pour concaténer le nouveau token à la séquence de sortie, qui sera utilisée lors de l'itération suivante jusqu'à ce que la règle d'arrêt soit atteinte.
Ce code est facile à comprendre, mais il est inefficace. À chaque itération, il renvoie l'ensemble de la séquence au modèle. Si le prompt a 1 000 tokens et que vous générez 100 nouveaux tokens, le modèle recompute à plusieurs reprises les états cachés pour les mêmes tokens de prompt. Le modèle traite O(N²) tokens dans cette fonction, pour un prompt de longueur N.
La complexité temporelle réelle du code est encore pire. Sans mise en cache, chaque passe avant recompute l'attention pour tous les tokens dans la séquence croissante. Si la longueur de la séquence est N, l'auto-attention a une complexité de calcul des scores de O(N²). Pour la génération, cela signifie que vous répétez une grande quantité de travail. (Précisément, si la longueur de la séquence de sortie est N=P+G avec la longueur du prompt P et le nombre de tokens générés G, la complexité de calcul devrait être O(P²G + PG² + G³) naïvement. Avec la mise en cache, nous pouvons la réduire à O(P² + PG).)
Les systèmes d'inférence atténuent cela en divisant la génération en deux phases : pré-remplissage et décodage.
Pré-remplissage et décodage
La génération commence généralement par un prompt. Le prompt est connu avant le début de la génération. Le modèle peut traiter tous les tokens du prompt en une seule passe avant. Cela s'appelle la phase de pré-remplissage.
Pendant le pré-remplissage, le modèle calcule les états cachés pour tous les tokens du prompt et produit des logits pour le prochain token. Il calcule également les clés et valeurs pour toutes les couches d'attention. Ces clés et valeurs peuvent être sauvegardées car elles seront nécessaires pour chaque token futur.
Après la sélection du premier nouveau token, la génération entre dans la phase de décodage. Dans cette phase, le modèle reçoit uniquement le token le plus récent. Il calcule la requête, la clé et la valeur pour ce token, ajoute la nouvelle clé et valeur au cache, et assiste la nouvelle requête sur toutes les clés et valeurs mises en cache.
Cela change le coût d'une étape de décodage. Au lieu de recomputer l'attention pour l'ensemble de la séquence, le modèle calcule l'attention pour une seule nouvelle requête par rapport à toutes les clés précédentes. Le coût d'attention par token passe de O(N²) à O(N) pour une séquence de longueur N. L'étape de pré-remplissage est toujours O(N²), mais elle n'est effectuée qu'une seule fois pour le prompt.
Cette distinction est suffisamment importante pour que les systèmes de service mesurent généralement le pré-remplissage et le décodage séparément :
- Le pré-remplissage affecte le temps jusqu'au premier token. Un pré-remplissage lent augmente le temps jusqu'au premier token.
- Le décodage affecte la vitesse des tokens de sortie en streaming. Un décodage lent réduit le taux auquel les tokens de sortie sont diffusés.
Un prompt court avec une longue réponse met l'accent sur le décodage. Un long prompt avec une courte réponse met l'accent sur le pré-remplissage. Une application de chat avec un long historique de conversation met l'accent sur les deux.
Un cache KV simple
Le cache KV est l'endroit où le modèle stocke les clés et valeurs d'attention produites par les tokens précédents. Pour voir comment cela fonctionne, vous n'avez pas besoin d'un grand modèle. Le code suivant construit un petit modèle de type transformer avec un cache.
Ce modèle n'est pas destiné à produire un texte utile. Son but est de montrer comment le cache est créé pendant le pré-remplissage et étendu pendant le décodage.
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class SelfAttention(nn.Module):
def __init__(self, hidden_size, num_heads):
super().__init__()
assert hidden_size % num_heads == 0
self.num_heads = num_heads
self.head_dim = hidden_size // num_heads
self.qkv = nn.Linear(hidden_size, 3 * hidden_size)
self.out = nn.Linear(hidden_size, hidden_size)
def forward(self, x, past_kv=None):
batch_size, seq_len, hidden_size = x.shape
qkv = self.qkv(x)
qkv = qkv.view(batch_size, seq_len, 3, self.num_heads, self.head_dim)
qkv = qkv.permute(2, 0, 3, 1, 4)
q, k, v = qkv[0], qkv[1], qkv[2]
if past_kv is not None:
past_k, past_v = past_kv
k = torch.cat([past_k, k], dim=2)
v = torch.cat([past_v, v], dim=2)
total_len = k.size(2)
past_len = total_len - seq_len
scores = q @ k.transpose(-2, -1)
scores = scores / math.sqrt(self.head_dim)
causal_mask = torch.ones(seq_len, total_len, device=x.device, dtype=torch.bool)
causal_mask = torch.tril(causal_mask, diagonal=past_len)
scores = scores.masked_fill(~causal_mask, float("-inf"))
attn = F.softmax(scores, dim=-1)
y = attn @ v
y = y.transpose(1, 2).contiguous().view(batch_size, seq_len, hidden_size)
return self.out(y), (k, v)
class Block(nn.Module):
def __init__(self, hidden_size, num_heads):
super().__init__()
self.attn_norm = nn.LayerNorm(hidden_size)
self.attn = SelfAttention(hidden_size, num_heads)
self.ffn_norm = nn.LayerNorm(hidden_size)
self.ffn = nn.Sequential(
nn.Linear(hidden_size, 4 * hidden_size),
nn.GELU(),
nn.Linear(4 * hidden_size, hidden_size),
)
def forward(self, x, past_kv=None):
attn_out, new_kv = self.attn(self.attn_norm(x), past_kv=past_kv)
x = x + attn_out
x = x + self.ffn(self.ffn_norm(x))
return x, new_kv
class TinyCausalLM(nn.Module):
def __init__(self, vocab_size=128, hidden_size=64, num_heads=4, num_layers=2):
super().__init__()
self.token_emb = nn.Embedding(vocab_size, hidden_size)
self.blocks = nn.ModuleList([
Block(hidden_size, num_heads) for _ in range(num_layers)
])
self.norm = nn.LayerNorm(hidden_size)
self.lm_head = nn.Linear(hidden_size, vocab_size, bias=False)
def forward(self, input_ids, past_kv=None):
x = self.token_emb(input_ids)
new_cache = []
if past_kv is None:
past_kv = [None] * len(self.blocks)
for block, layer_past in zip(self.blocks, past_kv):
x, layer_cache = block(x, past_kv=layer_past)
new_cache.append(layer_cache)
logits = self.lm_head(self.norm(x))
return logits, new_cache


