La caché de claves y valores: cuánto cuesta cada token
30 min de lectura
Sesenta y dos tokens tras La Nela, lo que le cabe al mini-GPT en su ventana, le cuestan unas dos mil filas de cálculo. Cada token sale de un sorteo, como estableció la lección sobre el muestreo, y cada sorteo necesita una pasada del texto entero por la red, de la que sólo se aprovecha la última fila. De esas dos mil filas se usan 62; las demás ya se habían calculado antes, idénticas.
Esta lección pone números a ese derroche y a su remedio. ¿Cuánto más caro sale un token si todo se rehace? La respuesta es exacta y no depende del modelo: veces, en la posición . ¿Y qué cuesta no rehacerlo? Memoria, que crece con cada token y que, en un modelo grande con un texto largo, llega a ocupar más que el propio modelo. Ese número decide cuánto contexto cabe en una máquina.
Cada posición de una capa de atención produce tres vectores, como en el curso anterior: una consulta, con la que busca; una clave, con la que la encuentran las posiciones que vienen detrás; y un valor, lo que les entrega cuando la eligen. La consulta sólo le sirve a su propia posición, y una vez. La clave y el valor los usan todas las posiciones posteriores, y siempre son los mismos: salen de la fila de su posición, y la máscara hace que esa fila no vea nada de lo que hay a su derecha. Un token más al final no les cambia una cifra.
Así que se guardan. En el explorable, el mini-GPT llena su ventana; debajo del texto está lo guardado, una columna por posición y una fila por matriz (claves y valores de sus dos capas), y debajo, las multiplicaciones de cada token. Haz tres cosas. Con caché, avanza token a token o pulsa Animar: cada token calcula su columna y lee las demás, y la barra verde casi no se mueve. Cambia a sin caché: cada token recalcula todas las columnas, y las barras grises forman una rampa. Y lleva hasta 64: sin caché, ese token cuesta 64 veces más; sumados, los 64 tokens cuestan 33 veces más.
Las claves y los valores de una posición no cambian
Fijemos una capa (uno de los bloques del mini-GPT, que su código llama capas) y una de sus cabezas. Tras leer posiciones, la cabeza ha calculado las matrices , y de la lección del curso anterior sobre el producto interno escalado, de forma , con filas , y . (Esta , con su posición, es la consulta de siempre, no la distribución del muestreo.) La máscara pone en las puntuaciones de la fila contra las columnas , y no hay ninguna: la posición es la última. La fila de la salida es, entonces,
sin máscara: la posición nueva mira a todas las anteriores y a sí misma. Pide su consulta y las filas de y de . La consulta y la última fila de cada matriz salen de la fila de la entrada de la capa. Las otras son el asunto.
La lección sobre el modelo de lenguaje causal citó del curso anterior que, con la máscara, la fila de lo que sale de una capa depende sólo de las filas a de lo que entró, y que eso sobrevive a apilar capas. En la entrada de cada capa, por tanto, la fila es función de y de nada más, y y son esa fila, normalizada y multiplicada por una matriz. Con el argumento diciendo qué texto se le dio a la red, como allí,
La clave y el valor de la posición se calcularon cuando era la última, y son los mismos, cifra a cifra, en cualquier texto que sólo crezca por la derecha.
Guardarlos es la caché de claves y valores (key–value cache, o KV cache): y , las matrices y de la capa , que tras posiciones tienen números cada una, entre las . Con ella, el token de la posición son cuatro cosas:
- la fila de la entrada: el embedding del token más el vector de su posición;
- en cada capa, , y a partir de esa fila, y , al final de la caché de cada cabeza;
- la ecuación de arriba, con la caché entera como y , y el resto de la capa (juntar las cabezas, el perceptrón por posiciones) sobre esa sola fila;
- al salir de la última capa, los logits , con los que se sortea el token .
Lo que sale son los logits de adelante sobre el texto entero. No aproximadamente: los mismos.
El prompt no se sortea. Sus tokens ya existen, y por la misma máscara sus filas se calculan todas de una pasada, como en el entrenamiento, dejando la caché con una fila por token del prompt. Esa primera pasada se llama prefill; después viene un token cada vez.
La caché vale mientras nadie cambie de posición
La fila lleva dentro su posición. El mini-GPT suma a cada embedding un vector aprendido por posición, de una tabla con filas: la opción que la lección sobre la codificación posicional mencionó y apartó, porque la tabla sólo sabe de las posiciones que se entrenaron. Mientras el texto crece por la derecha, cada token conserva la suya y todo lo guardado vale.
En el token 64 se llena la ventana de contexto, lo más que la red puede leer de una vez. (No es la
ventana de entrenamiento de la lección sobre el
entrenamiento, el trozo de 65 tokens del corpus del que la red aprendía.) Para escribir el 65, generar suelta el primer token y cada uno baja una posición. Su fila
cambia, con ella cambian sus claves y sus valores en todas las capas, y la caché entera deja de
valer. Con la caché, , las posiciones que el texto tiene, no pasa de . Qué hacer
con un texto que no cabe es un problema que el curso retoma más adelante.
Lo que cuesta un token, con caché y sin ella
Contemos multiplicaciones, como el curso anterior contó el precio de la atención en su lección sobre el adiós a la recurrencia: multiplicar una matriz por un vector cuesta una multiplicación por casilla de la matriz. Lo demás (normalizar, el softmax, la ReLU) cuesta poco por número y queda fuera.
Una fila que atraviesa la red pasa, en cada capa, por las tres proyecciones de consulta, clave y valor, por la matriz que junta las cabezas y por las dos del perceptrón por posiciones; al final, por la capa de salida. Llamemos a lo que cuesta:
Es una multiplicación por cada peso de una matriz. En el mini-GPT, : todos sus parámetros menos los que nunca multiplican una fila (la tabla de posiciones, que se consulta, los sesgos y los de las normalizaciones).
A eso se suma la atención. La consulta de una posición contra claves son productos escalares de términos, y mezclar valores, otras multiplicaciones; entre las cabezas, por capa. Con la caché, el token de la posición es una fila por la red más esa atención en cada capa:
Sin caché es adelante sobre las filas: filas por las matrices y, en cada cabeza, la
rejilla entera de puntuaciones y su mezcla, que el código calcula completa y enmascara
después,
exactamente veces más, en cada posición y para cualquier modelo. Las dos partes se multiplican por : las matrices, porque se rehacen filas; la atención, porque se rehacen filas de puntuaciones.
Sumemos sobre los tokens de un texto, con y :
La atención crece como con la caché y como sin ella, y las matrices como y como : lo que los artículos escriben frente a . Con la ventana del mini-GPT llena, , son multiplicaciones frente a , 33 veces menos.
Qué parte pesa más depende del modelo. La atención de un token iguala a sus matrices cuando , que en el mini-GPT es , ocho veces su . Por eso, con la caché, su coste por token es casi plano ( es un 12 % más que ). En un modelo grande, con textos de decenas de miles de tokens, la atención sí manda.
Lo que se paga a cambio: memoria
La caché guarda, por cada posición y cada capa, una clave y un valor de números entre todas las cabezas. Un texto de tokens ocupa
El mini-GPT con la ventana llena guarda , 128 KiB a 8 bytes por número. Un modelo abierto de unos 7 000 millones de parámetros, con y , guarda números por token: con 4 096 tokens a 2 bytes por número, 2 GiB, y con 32 768, 16 GiB, más que sus propios pesos, unos 13 GiB. Y eso, por cada conversación abierta a la vez. Muchos modelos actuales la reducen haciendo que varias cabezas compartan claves y valores, una variante que queda fuera del curso.
La caché, escrita sobre el mini-GPT
Las dos celdas corren sobre el mini-GPT, el checkpoint del bloque, sin entrenar nada. La primera
escribe dos funciones. llenar(ids) es el prefill: pasa el prompt por adelante una vez y se
queda con las claves y los valores que la ida ya calcula y guarda para la vuelta del entrenamiento. avanzar(x, cache) es la lista de cuatro cosas de arriba para un
token. Con las dos llena la ventana tras La Nela, sorteando con top-p, y después compara los
logits de ocho posiciones con los que da adelante sobre el texto entero.
import numpy as np
from pyodide.http import open_url
exec(open_url("/courses/llm-agents/bpe.py").read()) # codificar, decodificar
exec(open_url("/courses/llm-agents/minigpt.py").read()) # MiniGPT, softmax, layer_norm, muestrear
F = [tuple(par) for par in json.load(open_url("/courses/llm-agents/bpe-merges.json"))]
modelo = MiniGPT.cargar(open_url("/courses/llm-agents/minigpt.json").read())
p, L, h, d = modelo.p, modelo.cfg["n_capas"], modelo.cfg["h"], modelo.cfg["d_model"]
d_k = d // h
def llenar(ids):
"""El prompt, de una pasada: los logits de su última posición y la caché de cada capa."""
Z, (_, capas, _, _) = modelo.adelante(np.array([ids]))
cache = [(K[0], V[0]) for (_, _, _, K, V, _, _), _ in capas] # (h, t, d_k) cada una
return Z[0, -1], cache
def avanzar(x, cache):
"""Los logits tras el token x, que ocupa la posición siguiente a la caché; la caché crece."""
t = cache[0][0].shape[1] # posiciones ya guardadas
H = p["E"][x] + p["P"][t] # (d,): la fila nueva, sola
for l in range(L):
Hn, _ = layer_norm(H, p[f"{l}.ln1_g"], p[f"{l}.ln1_b"])
q, k, v = (u.reshape(h, 1, d_k) for u in np.split(Hn @ p[f"{l}.Wqkv"], 3))
K = np.concatenate([cache[l][0], k], axis=1) # (h, t + 1, d_k)
V = np.concatenate([cache[l][1], v], axis=1)
cache[l] = (K, V)
a = softmax(q @ K.transpose(0, 2, 1) / np.sqrt(d_k)) # (h, 1, t + 1): todo es pasado
H = H + (a @ V).reshape(d) @ p[f"{l}.Wo"] # las cabezas, juntas
H = H + modelo.ffn(H[None, None], l)[0][0, 0]
return layer_norm(H, p["lnf_g"], p["lnf_b"])[0] @ p["E"].T
rng = np.random.default_rng(0)
ids = codificar("La Nela", F)
z, cache = llenar(ids)
logits = [z] # logits[j]: los de la posición j + 2
while len(ids) < 64: # hasta llenar la ventana
ids.append(muestrear(z, 1.0, None, 0.9, rng))
z = avanzar(ids[-1], cache)
logits.append(z)
print(repr(decodificar(ids, F)))
dif = max(np.abs(logits[t - 2] - modelo.adelante(np.array([ids[:t]]))[0][0, -1]).max() for t in range(8, 65, 8))
print("mayor diferencia con adelante sobre el texto entero, en 8 posiciones: %.1e" % dif)
print("caché de la capa 1:", cache[0][0].shape, " números entre las dos capas:", sum(K.size + V.size for K, V in cache))
La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador y se reutiliza en todas las lecciones.
El texto es el del explorable, y los logits de las dos maneras difieren en : son
las mismas sumas hechas en otro orden, y ningún sorteo cambia por ese redondeo. La caché de la
primera capa acaba con forma (4, 64, 16), , y entre las dos capas guarda los
números de la cuenta de arriba.
Ahora cambia el 64 del while por 65. avanzar falla en la posición 65 con un IndexError: la
tabla de posiciones, P, tiene 64 filas y no hay una más. Es la sección sobre las posiciones, en
forma de error.
La segunda pone la cuenta al lado del reloj en cinco posiciones: las multiplicaciones de y de , y lo que tarda el token de las dos maneras. La última línea mira la posición 64 fila a fila.
import time
c_fila = L * (4 * d * d + 2 * d * modelo.cfg["d_ff"]) + d * modelo.cfg["n_v"]
def c(t):
"""Multiplicaciones del token de la posición t con la caché."""
return c_fila + 2 * L * t * d
def ms(f, veces=5):
"""La mediana de varias medidas de f(), en milisegundos: un cronómetro así tiembla."""
f() # la primera, de calentamiento
medidas = []
for _ in range(veces):
t0 = time.perf_counter()
f()
medidas.append(1000 * (time.perf_counter() - t0))
return np.median(medidas)
miles = lambda n: format(n, ",").replace(",", " ")
print(" t multiplicaciones: con caché sin caché | medido: con caché sin caché")
for t in [8, 16, 32, 48, 64]:
antes = [(K[:, :t - 1], V[:, :t - 1]) for K, V in cache] # la caché con t - 1 posiciones
con = ms(lambda: avanzar(ids[t - 1], list(antes)))
sin = ms(lambda: modelo.adelante(np.array([ids[:t]])))
print("%3d %29s %11s | %16.1f ms %8.1f ms" % (t, miles(c(t)), miles(t * c(t)), con, sin))
print("en t = 64: %.2f ms por fila en la pasada de 64, %.2f ms la fila sola" % (sin / 64, con))
La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador y se reutiliza en todas las lecciones.
Tus milisegundos serán otros; la forma, no. Sin caché, el tiempo sube en rampa, de unos 12 ms en la posición 8 a unos 65 en la 64, como . Con ella se queda entre 5 y 9 ms en cualquier posición, como , que en ese tramo sólo crece un 12 %. Pero en la posición 64 la cuenta dice 64 veces y el reloj, unas diez.
Lo que falta es un coste que la cuenta no ve. Cada pasada por la red lanza decenas de operaciones de NumPy, y cada una paga una parte fija, sea del tamaño que sea. Una fila sola la paga entera; 64 filas en una pasada la reparten, y por eso salen a 1 ms por fila y la fila sola a unos 6. No es cosa del navegador. En una tarjeta gráfica (graphics processing unit, GPU), la parte fija es sobre todo traer de la memoria todos los pesos del modelo, que sirven igual para una fila que para cien. Por eso un modelo lee el prompt, cientos de filas por pasada, mucho más deprisa de lo que escribe la respuesta, que va de fila en fila: cada token necesita el anterior ya sorteado.
Comprueba tu intuición
Cinco preguntas y un reto: la atención de un token contra la caché, fuera de la red.
Cuando el texto crece un token por la derecha, las claves y los valores de las posiciones anteriores no cambian, y por eso se pueden guardar. ¿Qué lo garantiza?
Con la caché, el mini-GPT pasa de escribir el token de la posición al de la . Marca lo que crece.
Marca todas las opciones correctas. Se corrige todo o nada: no hay puntuación parcial.
Un modelo de capas y guarda su caché a 2 bytes por número. ¿Cuántos GiB ocupa la caché de una sola conversación de tokens? (1 GiB son bytes.)
Se acepta un margen de ±0.05.
El texto del mini-GPT llega a 64 tokens, su , y quieres el 65 sin tirar la caché. ¿Qué pasa?
En la segunda celda, la red saca las 64 filas de una pasada a cerca de un milisegundo cada una, y una fila sola, con la caché, le cuesta varios. ¿Por qué el prompt se lee mucho más deprisa, por token, de lo que se escribe la respuesta?
Escribe atiende(q, k, v, K, V): la atención de un token nuevo en una capa con caché.
q, k y v son la consulta, la clave y el valor de la posición nueva, uno por
cabeza: forma (h, d_k). K y V son la caché de la capa, (h, t, d_k), con las
t posiciones anteriores (puede que ninguna). Añade k y v al final de la caché y
devuelve (o, K, V): la salida de cada cabeza en la posición nueva, (h, d_k), y la
caché con t + 1 filas por cabeza.
La primera comprobación descarga el intérprete de Python (~15 MB); después queda en la caché del navegador. Este desafío se resuelve mejor con un teclado físico: en el móvil puedes leerlo y volver luego.
El mini-GPT escribe ahora a una fila por token, y esta lección dice cuánto cuesta cada una. Lo que ninguna de sus cuentas dice es si lo que escribe vale algo. En la primera celda, estrelas, que no existe, y personas, que sí, salen igual de baratos, y la caché hace que los dos lleguen antes.
Hasta aquí, los veredictos sobre el mini-GPT han sido dos: una pérdida de por token en páginas de Marianela que no vio, y un recuento de palabras que existen en la novela. ¿Es mucho ? ¿Frente a qué, si el mismo modelo con otro tokenizador daría otro número? Convertir esa pérdida en una medida que se pueda comparar, y ponerle delante el modelo más pobre posible como vara de medir, es la lección siguiente, sobre la perplejidad.
Para profundizar3 fuentes · 3 papers
De dónde sale lo de esta lección, y dónde seguir si quieres más. Nada de aquí hace falta para continuar el curso.
- Fast Transformer Decoding: One Write-Head is All You Need
Las cuentas de esta lección para generar con caché, en pocas páginas, y la primera idea para encogerla: que todas las cabezas compartan una sola clave y un solo valor.
- GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints
El término medio que usan muchos modelos abiertos: grupos de cabezas que comparten claves y valores. Divide la caché de esta lección por el tamaño del grupo, y pierde poco.
- Efficiently Scaling Transformer Inference
Lo que la segunda celda sólo asoma, a escala de un centro de datos: por qué leer el prompt va limitado por las multiplicaciones y generar por la memoria, y cuánto ocupa la caché.