En el artículo anterior, vimos cómo un modelo de lenguaje convierte logits en probabilidades y toma muestras del siguiente token. ¿Pero de dónde vienen estos logits?
En este tutorial, adoptamos un enfoque práctico para comprender el proceso de generación:
Cómo la fase de prellenado procesa todo el mensaje en un solo paso paralelo Cómo la fase de decodificación genera tokens uno a la vez utilizando el contexto previamente calculado Cómo la caché KV elimina el cálculo redundante para hacer que la decodificación sea eficiente
Al final, comprenderá la mecánica de dos fases detrás de la inferencia LLM y por qué la caché KV es esencial para generar respuestas largas a escala.
Empecemos.
De la indicación a la predicción: comprensión del prellenado, la decodificación y la caché KV en los LLM
Foto de Neda Astani. Algunos derechos reservados.
Descripción general
Este artículo se divide en tres partes; ellos son:
Cómo funciona la atención durante el llenado previo de la fase de decodificación de la caché KV de inferencia LLM: cómo hacer que la decodificación sea más eficiente
Cómo funciona la atención durante el llenado previo
Considere el mensaje:
El clima de hoy es tan…
Como seres humanos, podemos inferir que el siguiente símbolo debería ser un adjetivo, porque la última palabra "entonces" es una configuración. También sabemos que probablemente describe el clima, por lo que palabras como "agradable" o "cálido" son más probables que algo no relacionado como "delicioso".
Los transformadores llegan a la misma conclusión mediante la atención. Durante el llenado previo, el modelo procesa todo el mensaje en un solo paso hacia adelante. Cada token se ocupa de sí mismo y de todos los tokens anteriores, creando una representación contextual que captura las relaciones a lo largo de la secuencia completa.
El mecanismo detrás de esto es la fórmula de atención del producto escalado:
$$
text{Atención}(Q, K, V) = mathrm{softmax}left(frac{QK^top}{sqrt{d_k}}right)V
$$
Analizaremos esto concretamente a continuación.
Para que el cálculo de la atención sea rastreable, asignamos a cada token un valor escalar que representa la información que contiene:
Fichas de posición Valores 1 Hoy 10 2 tiempo 20 3 es 1 4 entonces 5
Palabras como "es" y "así" tienen menos peso semántico que "hoy" o "tiempo" y, como veremos, la atención refleja esto de forma natural.
Cabezas de atención
En los transformadores reales, los pesos de atención son valores continuos aprendidos durante el entrenamiento a través del producto escalar $Q$ y $K$. El comportamiento de las cabezas de atención es aprendido y normalmente imposible de describir. Ninguna cabeza está programada para “atender posiciones iguales”. Las cuatro reglas siguientes son una ilustración simplificada para hacer que el mecanismo de atención sea más intuitivo, mientras que la agregación ponderada sobre $V$ es la misma.
Estas son las reglas en nuestro ejemplo de juguete:
Atender a los tokens en posiciones de números pares Atender al último token Atender al primer token Atender a cada token
Para simplificar en este ejemplo, las salidas de estos cabezales se combinan (promedian).
Repasemos el proceso de precarga:
Hoy
Fichas pares → ninguna Última ficha → Hoy → 10 Primera ficha → Hoy → 10 Todas las fichas → Hoy → 10
clima
Fichas pares → clima → 20 Última ficha → clima → 20 Primera ficha → Hoy → 10 Todas las fichas → promedio (Hoy, clima) → 15
es
Fichas pares → clima → 20 Última ficha → es → 1 Primera ficha → Hoy → 10 Todas las fichas → promedio (Hoy, clima, es) → 10.33
entonces
Fichas pares → promedio(tiempo, entonces) → 12.5 Última ficha → entonces → 5 Primera ficha → Hoy → 10 Todas las fichas → promedio(Hoy, tiempo, es, entonces) → 9
Atención paralela
Si el mensaje contuviera 100.000 tokens, calcular la atención paso a paso sería extremadamente lento. Afortunadamente, la atención se puede expresar como operaciones tensoriales, lo que permite calcular todas las posiciones en paralelo.
Esta es la idea clave de la fase de prellenado en la inferencia LLM: cuando proporciona un mensaje, contiene varios tokens y se pueden procesar en paralelo. Este procesamiento paralelo ayuda a acelerar el tiempo de respuesta del primer token generado.
Para evitar que los tokens vean tokens futuros, aplicamos una máscara causal, de modo que solo puedan atender a sí mismos y a los tokens anteriores.
import torch
tokens = [“Today”, “weather”, “is”, “so”]
n = len(tokens)
d_k = 64
V = torch.tensor([[10.], [20.], [1.], [5.]], dtype=torch.float32)
positions = torch.arange(1, n + 1).float() # 1-based: [1, 2, 3, 4]
idx = torch.arange(n)
causal_mask = idx.unsqueeze(1) >= idx.unsqueeze(0)
print(causal_mask)
import torch
tokens = [ “Today” , “weather” , “is” , “so” ]
n = len ( tokens )
d_k = 64
V = torch . tensor ( [ [ 10. ] , [ 20. ] , [ 1. ] , [ 5. ] ] , dtype = torch . float32 )
positions = torch . arange ( 1 , n + 1 ) . float ( ) # 1-based: [1, 2, 3, 4]
idx = torch . arange ( n )
causal_mask = idx . unsqueeze ( 1 ) >= idx . unsqueeze ( 0 )
print ( causal_mask )
Producción:
tensor([[ True, False, False, False],
[ True, True, False, False],
[ True, True, True, False],
[ True, True, True, True]])
tensor([[ True, False, False, False],
[ True, True, False, False],
[ True, True, True, False],
[ True, True, True, True]])
Ahora podemos empezar a escribir las “reglas” para las 4 cabezas de atención.
En lugar de calcular puntuaciones de los vectores $Q$ y $K$ aprendidos, los elaboramos a mano directamente para que coincidan con nuestras cuatro reglas de atención. Cada cabeza produce una matriz de puntuación de forma (n, n), con una puntuación por par de claves de consulta, que se enmascara y se pasa a través de softmax para producir pesos de atención:
def selector(condition, size):
“””Return a (size, d_k) tensor of +1/-1 depending on condition.”””
val = torch.where(condition, torch.ones(
size), -torch.ones(size)) # (size,)
# (size, d_k)
return val.unsqueeze(1).expand(size, d_k).contiguous()
# Shared query: every row asks for a property, and K encodes which tokens match it.
Q = torch.ones(n, d_k)
# Head 1: select even positions
# K says whether each token is at an even position.
K1 = selector(positions % 2 == 0, n)
scores1 = (Q @ K1.T) / (d_k ** 0.5)
# Head 2: select the last token
# K says whether each token is the last one.
K2 = selector(positions == n, n)
scores2 = (Q @ K2.T) / (d_k ** 0.5)
# Head 3: select the first token
# K says whether each token is the first one.
K3 = selector(positions == 1, n)
scores3 = (Q @ K3.T) / (d_k ** 0.5)
# Head 4: select all visible tokens uniformly
# K says all the tokens
K4 = selector(positions == positions, n)
scores4 = (Q @ K4.T) / (d_k ** 0.5)
# Stack all head score matrices: shape (4, n, n)
scores = torch.stack([scores1, scores2, scores3, scores4], dim=0)
# Apply causal mask so position i can only attend to positions <= i
scores = scores.masked_fill(~causal_mask.unsqueeze(0), -1e9)
# Convert logits to attention weights
weights = torch.softmax(scores, dim=-1)
# Optional safeguard for fully masked rows
all_masked = (scores <= -1e4).all(dim=-1, keepdim=True)
weights = torch.where(all_masked, torch.zeros_like(weights), weights)
# Compute contexts: (heads, n, n) @ (n, 1) -> (heads, n, 1)
contexts = (weights @ V).squeeze(-1)
print(“Contexts by attention head (rows) x token position (columns):n”, contexts)
context4 = contexts[:, -1]
print(“nContext for final prompt position:n”, context4)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
def selector ( condition , size ) :
“” “Return a (size, d_k) tensor of +1/-1 depending on condition.” “”
val = torch . where ( condition , torch . ones (
size ) , – torch . ones ( size ) ) # (size,)
# (size, d_k)
return val . unsqueeze ( 1 ) . expand ( size , d_k ) . contiguous ( )
# Shared query: every row asks for a property, and K encodes which tokens match it.
Q = torch . ones ( n , d_k )
# Head 1: select even positions
# K says whether each token is at an even position.
K1 = selector ( positions % 2 == 0 , n )
scores1 = ( Q @ K1 . T ) / ( d_k * * 0.5 )
# Head 2: select the last token
# K says whether each token is the last one.
K2 = selector ( positions == n , n )
scores2 = ( Q @ K2 . T ) / ( d_k * * 0.5 )
# Head 3: select the first token
# K says whether each token is the first one.
K3 = selector ( positions == 1 , n )
puntuaciones3 = ( Q @ K3 . T ) / ( maldita sea * * 0,5 )
# Cabeza 4: seleccione todos los tokens visibles de manera uniforme
# K dice todas las fichas
K4 = selector ( posiciones == posiciones , norte )
puntuaciones4 = ( Q @ K4 . T ) / ( maldita sea * * 0,5 )
# Apilar todas las matrices de puntuación de cabeza: forma (4, n, n)
montones = antorcha . pila ( [ puntuaciones1 , puntuaciones2 , puntuaciones3 , puntuaciones4 ] , tenue = 0 )
# Aplicar máscara causal para que la posición i solo pueda atender a posiciones <= i
montones = puntuaciones . relleno_enmascarado ( ~ máscara_causal . unsqueeze ( 0 ) , – 1e9 )
# Convertir logits en pesos de atención
pesas = antorcha . softmax ( puntuaciones , tenue = – 1 )
# Protección opcional para filas completamente enmascaradas
todos_enmascarados = ( puntuaciones <= – 1e4 ) . todo ( tenue = – 1 , mantenerdim = Verdadero )
pesas = antorcha . donde ( todos_enmascarados , antorcha . zeros_like ( pesos ) , pesos )
# Calcular contextos: (cabezas, n, n) @ (n, 1) -> (cabezas, n, 1)
contextos = ( pesos @ V ) . apretar ( – 1 )
print ( "Contextos por encabezado de atención (filas) x posición del token (columnas):n" , contextos )
contexto4 = contextos [ : , – 1 ]
print ( "nContexto para la posición final del mensaje:n" , contexto4 )
Producción:
Contexts by attention heads (rows) x token position (columns):
tensor([[10.0000, 20.0000, 20.0000, 12.5000],
[10.0000, 15.0000, 10.3333, 5.0000],
[10.0000, 10.0000, 10.0000, 10.0000],
[10.0000, 15.0000, 10.3333, 9.0000]])
Context for final prompt position:
tensor([12.5000, 5.0000, 10.0000, 9.0000])
Contexts by attention heads (rows) x token position (columns):
tensor([[10.0000, 20.0000, 20.0000, 12.5000],
[10.0000, 15.0000, 10.3333, 5.0000],
[10.0000, 10.0000, 10.0000, 10.0000],
[10.0000, 15.0000, 10.3333, 9.0000]])
Context for final prompt position:
tensor([12.5000, 5.0000, 10.0000, 9.0000])
El resultado de este paso se denomina vector de contexto, que representa un resumen ponderado de todos los tokens anteriores.
De contextos a logits
Cada cabeza de atención ha aprendido a captar diferentes patrones en la entrada. Juntos, los cuatro valores de contexto [12.5, 5.0, 10.0, 9.0] forman un resumen de lo que representa “El clima de hoy es tan…”. Luego se proyectará en una matriz, en la que cada columna codifica qué tan fuerte está asociado un vocabulario determinado con la señal de cada cabeza de atención, para dar una puntuación logit por palabra.
…
logits = context @ W_vocab
. . .
logits = context @ W_vocab
Para nuestro ejemplo, digamos que tenemos "agradable", "cálido" y "delicioso" en el vocabulario:
…
vocab = [“nice”, “warm”, “delicious”]
# Each column corresponds to a vocab word
# Each row corresponds to one attention head feature
W_vocab = torch.tensor([
[0.8, 0.6, 0.1], # head 1 weights → nice, warm, delicious
[0.5, 0.4, 0.2], # head 2 weights
[0.1, 0.2, 0.5], # head 3 weights
[0.2, 0.3, 0.1], # head 4 weights
]) # shape: (4, 3)
logits = context4 @ W_vocab # (4,) @ (4, 3) → (3,)
for word, logit in zip(vocab, logits):
print(f”{word:10s} {logit.item():.3f}”)
“`
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
. . .
vocab = [ “nice” , “warm” , “delicious” ]
# Each column corresponds to a vocab word
# Each row corresponds to one attention head feature
W_vocab = torch . tensor ( [
[ 0.8 , 0.6 , 0.1 ] , # head 1 weights → nice, warm, delicious
[ 0.5 , 0.4 , 0.2 ] , # head 2 weights
[ 0.1 , 0.2 , 0.5 ] , # head 3 weights
[ 0.2 , 0.3 , 0.1 ] , # head 4 weights
] ) # shape: (4, 3)
logits = context4 @ W _ vocab # (4,) @ (4, 3) → (3,)
for word , logit in zip ( vocab , logits ) :
print ( f “{word:10s} {logit.item():.3f}” )
` ` `
Entonces, los logits para "agradable" y "cálido" son mucho más altos que "delicioso".
nice 15.300
warm 14.200
delicious 8.150
nice 15.300
warm 14.200
delicious 8.150
La fase de decodificación de la inferencia LLM
Ahora supongamos que el modelo genera el siguiente token: "agradable". La tarea ahora es generar el siguiente token con el mensaje extendido:
El clima de hoy es tan agradable…
Las primeras cuatro palabras del mensaje extendido son las mismas que las del mensaje original. Y ahora tenemos la quinta palabra del mensaje.
Durante la decodificación, no volvemos a calcular la atención de todos los tokens anteriores ya que el resultado sería el mismo. En cambio, calculamos la atención solo para el nuevo token para ahorrar tiempo y recursos informáticos. Esto produce una única fila de atención nueva.
new_token = “nice”
tokens = tokens + [new_token]
new_value = torch.tensor([[7.0]]) # value of “nice” is 7
V = torch.cat([V, new_value], dim=0)
n = len(tokens)
idx = torch.arange(n)
pos = torch.arange(1, n + 1).float() # [1, 2, 3, 4, 5]
print(“New tokens: “, tokens)
print(“New Values: “, V)
new_token = “nice”
tokens = tokens + [ new_token ]
new_value = torch . tensor ( [ [ 7.0 ] ] ) # value of “nice” is 7
V = torch . cat ( [ V , new_value ] , dim = 0 )
n = len ( tokens )
idx = torch . arange ( n )
posición = antorcha . organizar ( 1 , norte + 1 ) . flotar ( ) # [1, 2, 3, 4, 5]
print ( "Nuevos tokens: " , fichas )
imprimir ( "Nuevos valores: " , V )
Producción:
New tokens: [‘Today’, ‘weather’, ‘is’, ‘so’, ‘nice’]
New Values: tensor([[10.],
[20.],
[ 1.],
[ 5.],
[ 7.]])
New tokens: [‘Today’, ‘weather’, ‘is’, ‘so’, ‘nice’]
New Values: tensor([[10.],
[20.],
[ 1.],
[ 5.],
[ 7.]])
Ahora aplicamos las 4 cabezas de atención y calculamos el nuevo vector de contexto:
# Rebuild all K matrices for the next token (n=5)
# We will introduce KV-cache later
K1_new = selector(pos % 2 == 0, n) # even positions → +1
K2_new = selector(pos == n, n) # last token → +1
K3_new = selector(pos == 1, n) # first token → +1
K4_new = selector(pos == pos, n) # all tokens → +1
# During decode, only compute Q for the NEW token (one row)
Q_new = torch.ones(1, d_k)
scores1_new = (Q_new @ K1_new.T) / (d_k ** 0.5) # (1, 5)
scores2_new = (Q_new @ K2_new.T) / (d_k ** 0.5) # (1, 5)
scores3_new = (Q_new @ K3_new.T) / (d_k ** 0.5) # (1, 5)
scores4_new = (Q_new @ K4_new.T) / (d_k ** 0.5) # (1, 5)
# Stack: shape (4, 1, 5)
new_scores = torch.stack(
[scores1_new, scores2_new, scores3_new, scores4_new], dim=0)
# No causal mask needed — new token can see all previous tokens by definition
new_weights = torch.softmax(new_scores, dim=-1) # (4, 1, 5)
context5 = (new_weights @ V).squeeze() # (4,)
print(“Visible tokens:”, tokens)
print(“Context for new token position:n”, context5)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
# Rebuild all K matrices for the next token (n=5)
# We will introduce KV-cache later
K1_new = selector ( pos % 2 == 0 , n ) # even positions → +1
K2_new = selector ( pos == n , n ) # last token → +1
K3_new = selector ( pos == 1 , n ) # first token → +1
K4_new = selector ( pos == pos , n ) # all tokens → +1
# During decode, only compute Q for the NEW token (one row)
Q_new = torch . ones ( 1 , d_k )
scores1_new = ( Q _ new @ K1_new . T ) / ( d_k * * 0.5 ) # (1, 5)
scores2_new = ( Q _ new @ K2_new . T ) / ( d_k * * 0.5 ) # (1, 5)
scores3_new = ( Q _ new @ K3_new . T ) / ( d_k * * 0.5 ) # (1, 5)
puntuaciones4_new = ( Q_nuevo @ K4_nuevo . T ) / ( maldita sea * * 0,5 ) # (1, 5)
# Pila: forma (4, 1, 5)
nuevos_puntuaciones = antorcha . pila (
[ puntuaciones1_new , puntuaciones2_new , puntuaciones3_new , puntuaciones4_new ] , tenue = 0 )
# No se necesita máscara causal: el nuevo token puede ver todos los tokens anteriores por definición
nuevos_pesos = antorcha . softmax ( nuevas puntuaciones , tenue = – 1 ) # (4, 1, 5)
contexto5 = ( nuevos _ pesos @ V ) . estrujar ( ) # (4,)
imprimir ( "Tokens visibles:" , fichas )
print ( "Contexto para la nueva posición del token:n" , contexto5 )
Producción:
Visible tokens: [‘Today’, ‘weather’, ‘is’, ‘so’, ‘nice’]
Context for new token position:
tensor([12.5000, 7.0000, 10.0000, 8.6000])
Visible tokens: [‘Today’, ‘weather’, ‘is’, ‘so’, ‘nice’]
Context for new token position:
tensor([12.5000, 7.0000, 10.0000, 8.6000])
Sin embargo, a diferencia del prellenado, donde todo el mensaje se procesa en paralelo, la decodificación debe generar tokens uno a la vez (autorregresivamente) porque los tokens futuros aún no se han generado. Sin el almacenamiento en caché, cada paso de decodificación volvería a calcular las claves y los valores de todos los tokens anteriores desde cero, lo que haría que el trabajo total en todos los pasos de decodificación fuera $O(n^2)$ en longitud de secuencia. La caché de KV reduce esto a $O(n)$ calculando los $K$ y $V$ de cada token exactamente una vez.
Caché KV: cómo hacer que la decodificación sea más eficiente
Para que la codificación autorregresiva sea eficiente, podemos almacenar las claves ($K$) y los valores ($V$) para cada token por separado para cada cabeza de atención. En este ejemplo simplificado usaríamos solo un caché. Luego, durante la decodificación, cuando se genera un nuevo token, el modelo no vuelve a calcular las claves y los valores de todos los tokens anteriores. Calcula la consulta del nuevo token y atiende las claves y valores almacenados en caché de los tokens anteriores.
Si volvemos a mirar el código anterior, podemos ver que no es necesario volver a calcular $K$ para todo el tensor:
K1_new = selector(pos % 2 == 0, n) # even positions → +1
K1_new = selector ( pos % 2 == 0 , n ) # even positions → +1
En su lugar, podemos simplemente calcular K para la nueva posición y adjuntarla a la matriz K que ya hemos calculado y guardado en caché:
K1_new = selector(new_pos % 2 == 0, 1) # is pos 5 even? → -1
K1_cache = torch.cat([K1, K1_new], dim=0) # (4→5, d_k)
K1_new = selector ( new_pos % 2 == 0 , 1 ) # is pos 5 even? → -1
K1_cache = torch . cat ( [ K1 , K1_new ] , dim = 0 ) # (4→5, d_k)
Aquí está el código completo para la fase de decodificación usando caché KV:
# In decode we only compute the query for the NEW token (position 5).
new_pos = pos[-1:] # tensor([5.])
# Compute ONLY the new token’s key for each head
K1_new = selector(new_pos % 2 == 0, 1) # is pos 5 even? → -1
K2_new = selector(new_pos == n, 1) # is pos 5 last? → +1
K3_new = selector(new_pos == 1, 1) # is pos 5 first? → -1
K4_new = selector(new_pos == new_pos, 1) # always → +1
# Append new key to the cached prefill keys
K1_cache = torch.cat([K1, K1_new], dim=0) # (4→5, d_k)
K2[-1] = -torch.ones(d_k) # position 4 is no longer last
K2_cache = torch.cat([K2, K2_new], dim=0)
K3_cache = torch.cat([K3, K3_new], dim=0)
K4_cache = torch.cat([K4, K4_new], dim=0)
# Q is only for the new token
Q_dec = torch.ones(1, d_k)
scores1_dec = (Q_dec @ K1_cache.T) / (d_k ** 0.5)
scores2_dec = (Q_dec @ K2_cache.T) / (d_k ** 0.5)
scores3_dec = (Q_dec @ K3_cache.T) / (d_k ** 0.5)
scores4_dec = (Q_dec @ K4_cache.T) / (d_k ** 0.5)
# Stack → (4 heads × 1 query × n keys)
scores_dec = torch.stack([scores1_dec, scores2_dec, scores3_dec, scores4_dec], dim=0)
# Softmax over key dimension
weights_dec = torch.softmax(scores_dec, dim=-1)
# Edge case: all-masked rows → zero context (same guard as prefill)
all_masked_dec = (scores_dec <= -1e4).all(dim=-1, keepdim=True)
weights_dec = torch.where(all_masked_dec, torch.zeros_like(weights_dec), weights_dec)
# Context vectors: (4 × 1 × n) @ (n × 1) → (4 × 1 × 1) → squeeze → (4,)
contexts_dec = (weights_dec @ V).squeeze(-1).squeeze(-1)
print(“nDecode context for ‘nice’ (one value per head):n”, contexts_dec)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
# In decode we only compute the query for the NEW token (position 5).
new_pos = pos[–1:] # tensor([5.])
# Compute ONLY the new token’s key for each head
K1_new = selector(new_pos % 2 == 0, 1) # is pos 5 even? → -1
K2_new = selector(new_pos == n, 1) # is pos 5 last? → +1
K3_new = selector(new_pos == 1, 1) # is pos 5 first? → -1
K4_new = selector(new_pos == new_pos, 1) # always → +1
# Append new key to the cached prefill keys
K1_cache = torch.cat([K1, K1_new], dim=0) # (4→5, d_k)
K2[–1] = –torch.ones(d_k) # position 4 is no longer last
K2_cache = torch.cat([K2, K2_new], dim=0)
K3_cache = torch.cat([K3, K3_new], dim=0)
K4_cache = torch.cat([K4, K4_new], dim=0)
# Q is only for the new token
Q_dec = torch.ones(1, d_k)
scores1_dec = (Q_dec @ K1_cache.T) / (d_k ** 0.5)
puntuaciones2_dec = ( Q _ dic @ K2_caché . T ) / ( maldita sea * * 0,5 )
puntuaciones3_dec = ( Q _ dic @ K3_caché . T ) / ( maldita sea * * 0,5 )
puntuaciones4_dec = ( Q _ dic @ K4_caché . T ) / ( maldita sea * * 0,5 )
# Pila → (4 cabezas × 1 consulta × n claves)
puntuaciones_dec = antorcha . pila ( [ puntuaciones1_dec , puntuaciones2_dec , puntuaciones3_dec , puntuaciones4_dec ] , tenue = 0 )
# Softmax sobre dimensión clave
pesos_dec = antorcha . softmax ( puntuaciones_dec , tenue = – 1 )
# Caso extremo: filas completamente enmascaradas → contexto cero (la misma protección que el prerrelleno)
todo_enmascarado_dec = ( puntuaciones_dec <= – 1e4 ) . todo ( tenue = – 1 , mantenerdim = Verdadero )
pesos_dec = antorcha . donde ( all_masked_dec , antorcha . ceros_like ( pesos_dec ) , pesos_dec )
# Vectores de contexto: (4 × 1 × n) @ (n × 1) → (4 × 1 × 1) → apretar → (4,)
contextos_dec = ( pesos _ diciembre @ V ) . apretar ( -1 ) . apretar ( – 1 )
print ( "nContexto de decodificación para 'agradable' (un valor por encabezado):n" , contextos_dec )
Producción:
Decode context for ‘nice’ (one value per head):
tensor([12.5000, 6.0000, 10.0000, 8.6000])
Decode context for ‘nice’ (one value per head):
tensor([12.5000, 6.0000, 10.0000, 8.6000])
Observe que esto es idéntico al resultado que calculamos sin el caché. La caché KV no cambia lo que calcula el modelo, pero elimina cálculos redundantes.
La caché de KV se diferencia de la caché de otras aplicaciones en que el objeto almacenado no se reemplaza sino que se actualiza. Cada nuevo token agregado al mensaje agrega una nueva fila al tensor almacenado. Implementar un caché KV que pueda actualizar eficientemente el tensor es la clave para acelerar la inferencia de LLM.
Lecturas adicionales
A continuación se presentan algunos recursos que pueden resultarle útiles:
Resumen
En este artículo, recorrimos las dos fases de la inferencia LLM. Durante el llenado previo, el mensaje completo se procesa en un paso directo paralelo y las claves y valores de cada token se calculan y almacenan. Durante la decodificación, el modelo genera un token a la vez, utilizando solo la consulta del nuevo token contra las claves y valores almacenados en caché para evitar un nuevo cálculo redundante. Prefill calienta el caché KV y la decodificación lo actualiza. Un llenado previo más rápido significa que verá más rápido el primer token en la respuesta y una decodificación más rápida significa que verá más rápido el resto de la respuesta. Juntas, estas dos fases explican por qué los LLM pueden procesar solicitudes largas rápidamente pero generar salida token por token, y por qué la caché KV es esencial para que esa generación sea práctica a escala.