La caché de claves y valores: cuánto cuesta cada token

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: tt veces, en la posición tt. ¿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 tt hasta 64: sin caché, ese token cuesta 64 veces más; sumados, los 64 tokens cuestan 33 veces más.

El mini-GPT escribe 64 tokens tras «La Nela». Con la caché, cada token calcula la columna de su posición y lee las anteriores; sin ella, las recalcula todas. Las barras son las multiplicaciones de cada token, de las dos maneras.

Las claves y los valores de una posición no cambian

Fijemos una capa (uno de los L=2L = 2 bloques del mini-GPT, que su código llama capas) y una de sus cabezas. Tras leer tt posiciones, la cabeza ha calculado las matrices Q\mathbf{Q}, K\mathbf{K} y V\mathbf{V} de la lección del curso anterior sobre el producto interno escalado, de forma t×dkt \times d_k, con filas qi⊤\mathbf{q}_i^{\top}, ki⊤\mathbf{k}_i^{\top} y vi⊤\mathbf{v}_i^{\top}. (Esta qt\mathbf{q}_t, con su posición, es la consulta de siempre, no la distribución q\mathbf{q} del muestreo.) La máscara pone −∞-\infty en las puntuaciones de la fila tt contra las columnas j>tj > t, y no hay ninguna: la posición tt es la última. La fila tt de la salida es, entonces,

ot⊤=softmax(qt⊤K⊤dk)V,\mathbf{o}_t^{\top} = \text{softmax}\left(\frac{\mathbf{q}_t^{\top}\mathbf{K}^{\top}}{\sqrt{d_k}}\right)\mathbf{V},

sin máscara: la posición nueva mira a todas las anteriores y a sí misma. Pide su consulta y las tt filas de K\mathbf{K} y de V\mathbf{V}. La consulta y la última fila de cada matriz salen de la fila tt de la entrada de la capa. Las otras t−1t - 1 son el asunto.

La lección sobre el modelo de lenguaje causal citó del curso anterior que, con la máscara, la fila ii de lo que sale de una capa depende sólo de las filas 11 a ii de lo que entró, y que eso sobrevive a apilar capas. En la entrada de cada capa, por tanto, la fila ii es función de x≤ix_{\le i} y de nada más, y ki\mathbf{k}_i y vi\mathbf{v}_i son esa fila, normalizada y multiplicada por una matriz. Con el argumento diciendo qué texto se le dio a la red, como allí,

ki(x1:t)=ki(x≤i),vi(x1:t)=vi(x≤i)para todo i≤t.\mathbf{k}_i(x_{1:t}) = \mathbf{k}_i(x_{\le i}), \qquad \mathbf{v}_i(x_{1:t}) = \mathbf{v}_i(x_{\le i}) \qquad \text{para todo } i \le t.

La clave y el valor de la posición ii se calcularon cuando ii 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): K(l)\mathcal{K}^{(l)} y V(l)\mathcal{V}^{(l)}, las hh matrices K\mathbf{K} y V\mathbf{V} de la capa ll, que tras tt posiciones tienen t×dkt \times d_k números cada una, t×dmodelt \times d_{\text{model}} entre las hh. Con ella, el token de la posición tt son cuatro cosas:

  1. la fila tt de la entrada: el embedding del token más el vector de su posición;
  2. en cada capa, qt\mathbf{q}_t, kt\mathbf{k}_t y vt\mathbf{v}_t a partir de esa fila, y kt⊤\mathbf{k}_t^{\top}, vt⊤\mathbf{v}_t^{\top} al final de la caché de cada cabeza;
  3. la ecuación de arriba, con la caché entera como K\mathbf{K} y V\mathbf{V}, y el resto de la capa (juntar las cabezas, el perceptrón por posiciones) sobre esa sola fila;
  4. al salir de la última capa, los logits zt\mathbf{z}_t, con los que se sortea el token t+1t + 1.

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 ii lleva dentro su posición. El mini-GPT suma a cada embedding un vector aprendido por posición, de una tabla con Tctx=64T_{\text{ctx}} = 64 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é, TT, las posiciones que el texto tiene, no pasa de TctxT_{\text{ctx}}. 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 cfilac_{\text{fila}} a lo que cuesta:

cfila=L(4dmodel2+2dmodel⋅dff)+dmodel⋅∣V∣.c_{\text{fila}} = L\left(4d_{\text{model}}^{2} + 2d_{\text{model}} \cdot d_{\text{ff}}\right) + d_{\text{model}} \cdot \lvert V \rvert.

Es una multiplicación por cada peso de una matriz. En el mini-GPT, 2(4⋅642+2⋅64⋅256)+64⋅512=131 0722\left(4 \cdot 64^{2} + 2 \cdot 64 \cdot 256\right) + 64 \cdot 512 = 131\,072: todos sus 136 448136\,448 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 tt claves son tt productos escalares de dkd_k términos, y mezclar tt valores, otras t⋅dkt \cdot d_k multiplicaciones; entre las hh cabezas, 2t⋅dmodel2t \cdot d_{\text{model}} por capa. Con la caché, el token de la posición tt es una fila por la red más esa atención en cada capa:

c(t)=cfila+2Lt⋅dmodel.c(t) = c_{\text{fila}} + 2Lt \cdot d_{\text{model}}.

Sin caché es adelante sobre las tt filas: tt filas por las matrices y, en cada cabeza, la rejilla entera de t×tt \times t puntuaciones y su mezcla, que el código calcula completa y enmascara después,

t cfila+2Lt2⋅dmodel=t c(t),t\,c_{\text{fila}} + 2Lt^{2} \cdot d_{\text{model}} = t\,c(t),

exactamente tt veces más, en cada posición y para cualquier modelo. Las dos partes se multiplican por tt: las matrices, porque se rehacen tt filas; la atención, porque se rehacen tt filas de puntuaciones.

Sumemos sobre los TT tokens de un texto, con ∑tt=T(T+1)/2\sum_t t = T(T+1)/2 y ∑tt2=T(T+1)(2T+1)/6\sum_t t^{2} = T(T+1)(2T+1)/6:

con cacheˊ:∑t=1Tc(t)=T cfila+LT(T+1)⋅dmodel,sin cacheˊ:∑t=1Tt c(t)=12T(T+1) cfila+13LT(T+1)(2T+1)⋅dmodel.\begin{aligned} \text{con caché:} \quad \sum_{t=1}^{T} c(t) &= T\,c_{\text{fila}} + LT(T+1) \cdot d_{\text{model}}, \\ \text{sin caché:} \quad \sum_{t=1}^{T} t\,c(t) &= \tfrac{1}{2}T(T+1)\,c_{\text{fila}} + \tfrac{1}{3}LT(T+1)(2T+1) \cdot d_{\text{model}}. \end{aligned}

La atención crece como T2T^{2} con la caché y como T3T^{3} sin ella, y las matrices como TT y como T2T^{2}: lo que los artículos escriben O(T2d)O(T^{2}d) frente a O(T3d)O(T^{3}d). Con la ventana del mini-GPT llena, T=64T = 64, son 8 921 0888\,921\,088 multiplicaciones frente a 295 526 400295\,526\,400, 33 veces menos.

Qué parte pesa más depende del modelo. La atención de un token iguala a sus matrices cuando 2Lt⋅dmodel=cfila2Lt \cdot d_{\text{model}} = c_{\text{fila}}, que en el mini-GPT es t=512t = 512, ocho veces su TctxT_{\text{ctx}}. Por eso, con la caché, su coste por token es casi plano (c(64)c(64) es un 12 % más que c(1)c(1)). 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 dmodeld_{\text{model}} números entre todas las cabezas. Un texto de TT tokens ocupa

2LT⋅dmodelnuˊmeros.2LT \cdot d_{\text{model}} \quad \text{números}.

El mini-GPT con la ventana llena guarda 2⋅2⋅64⋅64=16 3842 \cdot 2 \cdot 64 \cdot 64 = 16\,384, 128 KiB a 8 bytes por número. Un modelo abierto de unos 7 000 millones de parámetros, con L=32L = 32 y dmodel=4 096d_{\text{model}} = 4\,096, guarda 262 144262\,144 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 json
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))
numpy

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 4×10−154 \times 10^{-15}: 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), h×T×dkh \times T \times d_k, y entre las dos capas guarda los 16 38416\,384 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 c(t)c(t) y de t c(t)t\,c(t), y lo que tarda el token de las dos maneras. La última línea mira la posición 64 fila a fila.

# Necesita la celda anterior: modelo, ids, cache, avanzar, L, d, np.
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))
numpy

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 t c(t)t\,c(t). Con ella se queda entre 5 y 9 ms en cualquier posición, como c(t)c(t), 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 t=10t = 10 al de la t=60t = 60. Marca lo que crece.

Marca todas las opciones correctas. Se corrige todo o nada: no hay puntuación parcial.

Un modelo de L=32L = 32 capas y dmodel=4 096d_{\text{model}} = 4\,096 guarda su caché a 2 bytes por número. ¿Cuántos GiB ocupa la caché de una sola conversación de 8 1928\,192 tokens? (1 GiB son 2302^{30} bytes.)

GiB

Se acepta un margen de ±0.05.

El texto del mini-GPT llega a 64 tokens, su TctxT_{\text{ctx}}, 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 3.23.2 por token en páginas de Marianela que no vio, y un recuento de palabras que existen en la novela. ¿Es mucho 3.23.2? ¿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.

¿Te ha sido útil?
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
    paperNoam Shazeer, 2019arXiv:1911.02150EN

    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
    paperAinslie, Lee-Thorp, de Jong, Zemlyanskiy, Lebrón y Sanghai, 2023EMNLP 2023 · arXiv:2305.13245EN

    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
    paperPope, Douglas, Chowdhery, Devlin, Bradbury, Levskaya, Heek, Xiao, Agrawal y Dean, 2022MLSys 2023 · arXiv:2211.05102EN

    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é.