Saltar al contenido
9/30Capítulo 9 de 30

Attention y el bloque transformer, derivados de una media

Partimos de la media como resumen barato del contexto, medimos su fallo y dejamos que la fórmula de attention nazca al repararla.

En esta página

Llegas aquí con un tokenizador del Capítulo 7, una tabla de embedding del Capítulo 8 y el objetivo que los acompaña: dados los tokens hasta ahora, asignar una probabilidad al siguiente.

Lo que falta es el medio. Para predecir el token tt, el modelo necesita un vector que resuma todo lo anterior, y nada de lo que has construido produce uno. El embedding del token t1t-1 no sirve: eso es un modelo bigrama, y no puede saber que la frase empezó con una pregunta. Una concatenación de todos los embeddings anteriores tampoco sirve: su número cambia en cada paso, y una matriz de pesos fija no puede aceptar una entrada de longitud variable.

Así que: un vector de tamaño fijo que resume un número variable de vectores. Ese es todo el problema, y attention es lo que obtienes al resolverlo de la forma más perezosa posible y después reparar las dos cosas que se rompen.

La respuesta que tenía el campo, y por qué no la estamos construyendo

Enlace a la sección: La respuesta que tenía el campo, y por qué no la estamos construyendo

De 1997 a aproximadamente 2017, el resumen era un estado recurrente: mantener un vector h\mathbf{h} y actualizarlo en cada token, ht=f(ht1,xt)\mathbf{h}_t = f(\mathbf{h}_{t-1}, \mathbf{x}_t). Tamaño fijo, entrada variable, exactamente la forma adecuada.

Fallaba de tres maneras, y la arquitectura de este capítulo responde a las tres. Hacer backpropagation a través de TT pasos multiplica TT jacobianos, así que el gradient se desvanece o explota: la enfermedad que el Capítulo 5 midió dentro de un único nodo tanh\tanh. La LSTM1 se diseñó precisamente contra eso y empujó el rango utilizable de decenas de pasos a cientos, sin cambiar el hecho de que la información del token 5 llega al token 500 solo si sobrevive a 495 actualizaciones secuenciales. Toda la fuente tenía que caber en un vector: en traducción sequence-to-sequence2, un codificador comprime la entrada en su estado final. Bahdanau, Cho y Bengio pusieron nombre a ese cuello de botella y lo arreglaron en 2014, tres años antes del transformer, dejando que el decodificador tomara una suma ponderada de todos los estados del codificador con pesos que calculaba él mismo.3 Todo lo que sigue es esa idea, aplicada por una secuencia a sí misma, con la recurrencia eliminada. Y la actualización es secuencial por construcción: ht\mathbf{h}_t necesita ht1\mathbf{h}_{t-1}, y una GPU con diez mil núcleos no puede hacer nada con eso. La arquitectura que ganó no es obviamente más lista; es aquella cuyo paso caro es una multiplicación de matrices.

El otro sesgo inductivo clásico, la convolución —deslizar un filtro pequeño por toda la entrada, de modo que una característica detectada en cualquier sitio se detecte en todas partes— tampoco se construye aquí; es casi exactamente lo correcto para imágenes y queda delegado a un curso de visión. Ni la recurrencia ni la convolución reaparecen después de esta página, por eso ninguna recibe un capítulo: el Capítulo 1 prometió que las omisiones se declararían, no se esconderían.

La función más obvia que toma un número variable de vectores y devuelve un vector es la media:

ct=1ti=1txi\mathbf{c}_t = \frac{1}{t}\sum_{i=1}^{t} \mathbf{x}_i

Cualquier número de entradas, tamaño de salida fijo, diferenciable, gratis. Una tabla de embedding más esta media más una capa lineal hacia el vocabulario es un modelo de lenguaje completo en quince líneas. También es terrible, y la forma en que es terrible es toda la derivación.

El corpus de abajo es un megabyte de Shakespeare, 1.115.394 caracteres, pasado por un tokenizador BPE a nivel de bytes del tipo construido en el Capítulo 7 con un vocabulario de 1024: 459.760 tokens a 2,43 caracteres cada uno, dividido 90/10. Todos los modelos tienen anchura 128, ven 128 tokens y entrenan durante 3000 pasos de AdamW a 10310^{-3} con un batch de 64. La perplejidad se mide en la partición reservada.4

modeloparámetrosperplejidad de validación
solo el token actual, sin contexto alguno263.16859,71
más la media uniforme de todo lo anterior263.168248,07
más embeddings de posición aprendidos279.552245,93
media uniforme sumada al token en vez de reemplazarlo263.16860,45

Lee la segunda fila dos veces. Promediar el contexto no ayuda un poco; hace que el modelo sea cuatro veces peor que ignorar el contexto por completo. Dos razones, ambas demostrables en vez de empíricas.

La media no puede ver el orden. La suma conmuta, así que barajar la context window deja el resumen intacto; no aproximadamente:

order.pyPYTHON
A = torch.tril(torch.ones(T, T))
A = A / A.sum(1, keepdim=True)          # rows of the averaging matrix
y = x[torch.randperm(T)]                # the same tokens, shuffled
print((A[-1] @ x - A[-1] @ y).abs().max().item())
TEXT
2.9802322387695312e-08

Ruido de coma flotante en una suma reordenada: los dos resúmenes son el mismo vector. Un modelo cuya única vista del contexto es una media no puede distinguir the dog bit the man de the man bit the dog. La tercera fila demuestra que esto no se arregla añadiendo posiciones a las entradas: un embedding de posición aprendido en cada token antes de promediar compró 2,14 puntos de 188. Las posiciones entran en la suma, y la suma las olvida.

Y la media ahoga el presente. En la posición 100, el token actual es una centésima parte del resumen. Eso tiene una reparación barata que ya conoces: conservar el token y sumarle el resumen: una conexión residual, del Capítulo 6, y la cuarta fila muestra lo que hace. Con la dilución reparada, la media uniforme no aporta nada en absoluto: 60,45 frente a una línea base de 59,71. Todos los tokens están ahí, ponderados por igual, y ponderar por igual equivale a no tener información.

El problema no es promediar. Son los pesos.

La media es una multiplicación de matrices, y la máscara es una softmax

Enlace a la sección: La media es una multiplicación de matrices, y la máscara es una softmax

Promediar sobre un prefijo creciente parece un bucle. Es una multiplicación por una matriz triangular inferior cuyas filas suman uno; y también es, exactamente, una softmax:

mechanics.pyPYTHON
loop = torch.stack([x[:t + 1].mean(0) for t in range(T)])   # the obvious version

A = torch.tril(torch.ones(T, T))
A = A / A.sum(1, keepdim=True)
mat = A @ x                                                  # the same thing

S = torch.zeros(T, T).masked_fill(torch.tril(torch.ones(T, T)) == 0, float("-inf"))
soft = F.softmax(S, dim=-1) @ x                              # and the same thing again
TEXT
loop vs matmul   max |diff| = 5.960464477539063e-08
loop vs softmax  max |diff| = 5.960464477539063e-08

the averaging matrix A (rows sum to 1, upper triangle is zero):
  1.000 0.000 0.000 0.000 0.000 0.000
  0.500 0.500 0.000 0.000 0.000 0.000
  0.333 0.333 0.333 0.000 0.000 0.000
  0.250 0.250 0.250 0.250 0.000 0.000
  0.200 0.200 0.200 0.200 0.200 0.000
  0.167 0.167 0.167 0.167 0.167 0.167

Tres componentes con nombre de un transformer están ya en pantalla. El triángulo es la máscara causal, impuesta por el objetivo: si la posición tt pudiera ver la posición t+1t{+}1, la respuesta estaría en la entrada: la fuga que el Capítulo 6 te dijo que auditaras, pero dentro de la arquitectura. La softmax es cómo se implementa la máscara: poner las entradas prohibidas a -\infty las envía exactamente a cero y normaliza lo que queda, así que enmascarar y normalizar son una sola operación. (Usa -\infty, no -1e9: es el valor que el enmascaramiento significa, sobrevive a una conversión a float16 como -\infty y te evita decidir si la constante que elegiste es lo bastante grande para el rango en el que estás, que es la caja de coma flotante del Capítulo 2 haciendo una pregunta que no tienes que responder.) Y los scores son el parámetro libre. La media uniforme es lo que obtienes cuando todos los scores permitidos son el mismo número; pon ahí cualquier número y la softmax los convierte en pesos válidos.

El resto de este capítulo es una pregunta: ¿de dónde salen esos números?

No pueden ser parámetros simples. Una matriz T×TT \times T aprendida sería idéntica para cada frase: podría codificar “mira cuatro tokens atrás”, pero nunca “mira el sustantivo al que se refiere este pronombre”. El peso que enlaza la posición tt con la posición ii debe depender de lo que hay en ambas posiciones, porque la relevancia es una relación, no una propiedad: la palabra it no es intrínsecamente relevante; es relevante para algo.

La función más barata de dos vectores que devuelve un número es el producto escalar del Capítulo 1. Puntúa la posición ii para la posición tt como xtxi\mathbf{x}_t \cdot \mathbf{x}_i y el mecanismo funciona, mal, de dos maneras que obligan a todo lo demás. El producto escalar de un vector consigo mismo es su norma al cuadrado, así que cada token atendería sobre todo a sí mismo. Y la relación sería simétrica: si it atiende mucho a animal, entonces animal atiende mucho a it, lo cual es falso en el lenguaje, donde un adjetivo necesita su sustantivo mucho más de lo que el sustantivo necesita el adjetivo.

Así que dale a cada token dos roles, como dos mapas lineales aprendidos de él: lo que esta posición está buscando, qt=Wqxt\mathbf{q}_t = W_q\mathbf{x}_t, la query; y lo que ofrece para que lo encuentren, ki=Wkxi\mathbf{k}_i = W_k\mathbf{x}_i, la key. Puntúa qtki\mathbf{q}_t \cdot \mathbf{k}_i y la simetría desaparece, porque WqWkW_q \neq W_k: un token puede anunciar una cosa y buscar otra.

Aún queda una cosa mal. La suma ponderada era sobre los propios xi\mathbf{x}_i, lo que obliga a que lo que se copia sea lo que se empareja. Emparejar quiere las características que identifican un token; copiar quiere las características útiles aguas abajo. Así que aprende un tercer mapa, vi=Wvxi\mathbf{v}_i = W_v\mathbf{x}_i, el value, y suma esos.

La fórmula es ahora contabilidad:

Attention(Q,K,V)=softmax ⁣(QKdk+M)V\mathrm{Attention}(Q, K, V) = \mathrm{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}} + M\right)V

con MM como máscara causal, cero en y por debajo de la diagonal y -\infty por encima. En código son treinta líneas, veinte de ellas formas:

attention.pyPYTHON
class Head(nn.Module):
    """One head of causal self-attention."""

    def __init__(self, d_model, d_head, block):
        super().__init__()
        self.q = nn.Linear(d_model, d_head, bias=False)      
        self.k = nn.Linear(d_model, d_head, bias=False)      
        self.v = nn.Linear(d_model, d_head, bias=False)      
        self.d_head = d_head
        self.register_buffer("mask", torch.tril(torch.ones(block, block)).bool())

    def forward(self, x):
        T = x.shape[1]
        q, k, v = self.q(x), self.k(x), self.v(x)
        s = q @ k.transpose(-2, -1) / math.sqrt(self.d_head)          
        s = s.masked_fill(~self.mask[:T, :T], float("-inf"))          
        w = F.softmax(s, dim=-1)                                      
        return w @ v                                                  

Puntuar, enmascarar, normalizar, mezclar. Todo lo demás es una proyección.

La división por la raíz cuadrada, y de qué protege

Enlace a la sección: La división por la raíz cuadrada, y de qué protege

Casi toda explicación de dk\sqrt{d_k} dice “para evitar que la softmax se sature”, lo cual es cierto y no explica nada. El argumento son dos líneas de la varianza del Capítulo 2. Si las entradas de q\mathbf{q} y k\mathbf{k} son independientes, con media cero y varianza uno, cada producto qjkjq_j k_j tiene varianza uno, y las varianzas de cosas independientes se suman:

Var(qk)=j=1dkVar(qjkj)=dk\mathrm{Var}(\mathbf{q}\cdot\mathbf{k}) = \sum_{j=1}^{d_k}\mathrm{Var}(q_j k_j) = d_k

Así que los scores tienen desviación típica dk\sqrt{d_k}. Medido sobre veinte mil pares aleatorios:

TEXT
     d     Var(q.k)         std   sqrt(d)
     4        3.975       1.994     2.000
    16       16.071       4.009     4.000
    64       64.249       8.016     8.000
   256      253.065      15.908    16.000
  1024     1015.562      31.868    32.000

Por qué importa: la softmax es sensible a la escala de una forma en que una capa lineal no lo es. Duplicar la entrada de una capa lineal duplica su salida; multiplicar los scores por diez antes de una softmax convierte una mezcla suave en una elección dura. Una fila de 64 scores, con y sin la división:

dkd_kpeso mayor, sin dividirentropíatokens efectivospeso mayor, divididoentropíatokens efectivos
40,2052,94419,00,0813,75842,9
160,4381,6925,40,0753,84946,9
640,4890,8742,40,0853,67339,4
2560,99990,00071,00,1433,54734,7
10241,00000,00001,00,1323,64438,3

“Tokens efectivos” es el exponencial de la entropía: sobre cuántas posiciones promedia realmente la fila. Sin dividir, en dk=256d_k = 256, una head recién inicializada atiende exactamente a un token de 64, elegido solo por el sorteo aleatorio.

Eso es malo hacia delante y peor hacia atrás, con una forma que el Capítulo 5 ya midió en una tanh\tanh. Una softmax comprometida con una entrada tiene casi derivada cero: la diagonal de su jacobiano es wi(1wi)w_i(1-w_i), cero en ambos extremos. Sobre dos mil filas aleatorias:

dkd_kiwi(1wi)\sum_i w_i(1-w_i) sin dividirdivididofilas saturadas (peso mayor por encima de 0,99)
40,84270,95680,2 % → 0,0 %
640,29400,960917,9 % → 0,0 %
2560,14060,960949,1 % → 0,0 %
10240,06810,961170,4 % → 0,0 %

En dk=1024d_k = 1024, siete filas de cada diez están congeladas antes de empezar el entrenamiento, y una head que empieza congelada no puede aprender dónde mirar. Dividida, la cantidad se mantiene plana en 0,96 en cualquier anchura y nada se satura.

Ahora la parte que nadie publica: ¿cambia la perplejidad final? Elimina la división y entrena, con cuatro anchuras de head:

anchura de headsin dividirdividido por dk\sqrt{d_k}dividido por dkd_k
cuatro heads, dk=32d_k = 3237,2938,0737,89
una head, dk=128d_k = 12848,5146,1045,99
una head, dk=256d_k = 25665,3747,53
una head, dk=512d_k = 51267,0649,15
una head, dk=1024d_k = 102476,6959,17

Las dos primeras filas salen del presupuesto de 3000 pasos anterior; las tres últimas son una ejecución más corta —1500 pasos, batch de 32, una head, sin normalización antes de las proyecciones— con ambas variantes bajo ajustes idénticos.

En dk=32d_k = 32 la división no vale nada y la ejecución sin ella queda ligerísimamente por delante. Eso no es una licencia para quitarla, porque en 256 vale 18 puntos de perplejidad y en 1024 vale 17. El mecanismo se ve en los propios scores:

dkd_kdesviación típica del score al iniciotras 1500 pasos, sin dividirtras 1500 pasos, divididofilas saturadas, sin dividirdividido
25610,49121,672,1391,9 %0,8 %
51215,13836,852,6698,7 %1,3 %
102421,155147,463,4499,9 %16,5 %

La head sin dividir no se recupera. Se desboca: la desviación típica de sus scores pasa de 21 en la inicialización a 5147, la entropía de attention cae a cero, y el 99,9 % de las filas pone más de 0,99 de su peso en un único token. Una vez que una head es un selector duro, su gradient es casi cero y nada la trae de vuelta, así que el colapso es estable. La head dividida se queda en una desviación típica de score de 3,44 tras el mismo entrenamiento, que es una mezcla suave que todavía puede cambiarse.

Vaswani et al. dicen exactamente esto y nada más: sospechan que los productos “crecen en magnitud para valores grandes de dkd_k” y dividen.5 La palabra grandes soporta el peso, y las tablas dicen dónde empieza lo grande: nada en 32, todo en 256.

Más de una opinión, y los dos tercios de los que nadie habla

Enlace a la sección: Más de una opinión, y los dos tercios de los que nadie habla

Una head es una fila de softmax por posición, así que contiene una respuesta a “qué es relevante aquí”. Predecir la palabra después de the en the animal that crossed the wet street necesita a la vez el hueco sintáctico, el sujeto y el token anterior, y una única distribución de probabilidad no puede concentrarse en tres lugares. Así que ejecuta varias heads en paralelo, cada una de anchura dmodel/hd_{\text{model}}/h, concatena y mezcla con una matriz más WoW_o: has particionado la anchura, no la has aumentado.

Attention también hace exactamente una cosa: mueve información entre posiciones. Cada operación del código anterior es lineal a lo largo del eje de características, y el Capítulo 5 demostró qué es una pila de mapas lineales. Así que cada bloque también lleva una pequeña MLP aplicada a cada posición de forma independiente, que expande la anchura por cuatro y vuelve, con una GELU en medio. Merece la pena memorizar la división del trabajo: attention mezcla entre posiciones; la red feed-forward calcula dentro de una posición.

La escalera completa, cada fila añadiendo una pieza a la fila de arriba:

modeloparámetrosperplejidad de validación
media uniforme, sumada279.55260,45
una attention head, reemplazando el token328.70455,47
una attention head, sumada328.70446,10
cuatro heads en vez de una345.21643,21
más la red feed-forward476.92839,87
más LayerNorm: el bloque completo477.69638,07

Los pesos aprendidos vencen a los uniformes por 14 puntos de perplejidad, que es todo el argumento de este capítulo en una fila. Cuatro heads compran otros 3 por 16.512 parámetros extra. Y la misma head vale 9 puntos más sumada que reemplazando: attention trae información hacia dentro, no decide qué es una posición.

Ahora dónde se sientan realmente los parámetros, lo que sorprende a quienes solo han visto el diagrama:

anchuraheadsattentionfeed-forwardtotal por bloque
128465.664 (33,2 %)131.712 (66,6 %)197.888
768122.360.064 (33,3 %)4.722.432 (66,6 %)7.085.568
40963267.112.960 (33,3 %)134.238.208 (66,7 %)201.367.552

Dos tercios de cada bloque transformer son la red feed-forward, a cualquier escala, porque attention tiene cuatro matrices d×dd \times d y la MLP tiene el equivalente a ocho. Sea lo que sea que sepa un modelo, la mayoría de los parámetros que lo sostienen están en la MLP por posición.

Residuales y LayerNorm, heredados del Capítulo 6

Enlace a la sección: Residuales y LayerNorm, heredados del Capítulo 6

LayerNorm se construyó y midió en el Capítulo 6, y este capítulo la usa tal y como quedó allí; las conexiones residuales se nombraron y ablacionaron allí, y se construyen aquí. Las filas de “sumada, no reemplazando” de arriba son conexiones residuales, que valen 188 puntos de perplejidad para la media y 9 para una head. LayerNorm7 normaliza cada ejemplo a través de sus características, y el Capítulo 6 dio las razones por las que ella, y no BatchNorm, sobrevivió aquí: sin dependencia del batch, sin estadísticas acumuladas, idéntica en entrenamiento e inferencia, indiferente a la longitud de secuencia. Cada una de esas razones se convierte en requisito cuando generas un token cada vez para un usuario, que es donde acaba el Capítulo 13. Cuesta 768 parámetros y compra 1,8 puntos de perplejidad.

block.pyPYTHON
class Block(nn.Module):
    def forward(self, x):
        x = x + self.att(self.ln1(x))     
        x = x + self.ff(self.ln2(x))      
        return x

Mira dónde se coloca la normalización: en la entrada de cada subcapa, con la ruta residual de entrada a salida nunca normalizada. Eso es pre-norm. El artículo de 2017 hace lo contrario, x = LayerNorm(x + Att(x)): post-norm, que pone una LayerNorm en la propia ruta residual.

Xiong et al. explicaron la diferencia a través del gradient en la inicialización, que en una red post-norm está mal escalado con la profundidad: la razón por la que el transformer original necesitaba un warmup de learning rate para entrenar siquiera.8 Doce bloques, 1000 pasos, learning rate 3×1033 \times 10^{-3}:

TEXT
gradient norm per block at initialisation, before any step
  pre-norm    block 1 0.0498 ... block 12 0.0657   ratio last/first  1.32
  post-norm   block 1 0.0977 ... block 12 0.1613   ratio last/first  1.65

  pre-norm,  no warmup          perplexity   37.82
  pre-norm,  200-step warmup    perplexity   37.62
  post-norm, no warmup          perplexity  308.05
  post-norm, 200-step warmup    perplexity   37.88

Post-norm sin warmup es ocho veces peor, y post-norm con warmup iguala exactamente a pre-norm. Warmup no es aquí una buena práctica general; es un parche para una disposición concreta de la normalización, y mover la LayerNorm elimina su necesidad. Por eso prácticamente todos los modelos desde 2019 son pre-norm, y por eso el diagrama de 2017 debe leerse como historia, no como especificación.

Elimina los embeddings de posición y el modelo sigue entrenando; simplemente no puede saber dónde está nada, y eso es una simetría, no un fallo de entrenamiento. Nada en el score de attention menciona las propias tt o ii, así que permutar la entrada permuta la salida: self-attention es equivariante a permutaciones. Es la ceguera al orden de la media con un disfraz mejor: la máscara causal restaura parte del orden, ya que cada posición ve un prefijo distinto, pero dentro de un prefijo todos los ordenamientos son iguales.

Cuatro formas de inyectar posición, entrenadas en ventanas de 64 tokens y evaluadas en 64, 128 y 256, más allá de cualquier longitud que vieran:

posicionesperplejidad en 64en 128en 256
ninguna48,7952,6357,52
embeddings absolutos aprendidos38,63108,47181,94
sinusoides fijos42,9695,26152,25
RoPE44,1250,5284,84
ALiBi44,9543,5142,49

Embeddings absolutos aprendidos —un vector por posición, sumado al token— ganan en la longitud entrenada y luego caen por un precipicio, porque la posición 100 nunca estuvo en un batch y su embedding sigue siendo el vector aleatorio con el que empezó. Sinusoides, la elección original, se calculan en vez de aprenderse, a partir de senos y cosenos con frecuencias espaciadas geométricamente; el artículo de 2017 esperaba que eso extrapolara, y la tabla dice que no: la función está definida en la posición 200, pero el modelo nunca aprendió a leerla allí. RoPE9 no añade nada y en su lugar rota query y key por un ángulo proporcional a la posición, en cortes bidimensionales; como rotar ambos lados de un producto escalar por igual lo deja sin cambios, el score acaba dependiendo solo de tit - i, así que la posición se vuelve relativa gratis y no hay una tabla que se agote. Se degrada, pero se degrada. ALiBi10 es el resultado más simple y más extraño aquí: una penalización lineal sobre el score proporcional a la distancia, con una pendiente distinta por head. Su perplejidad mejora cuando la context window crece más allá de la longitud de entrenamiento, de 44,95 a 42,49, porque la penalización está definida a cualquier distancia y cada head sigue haciendo lo que se entrenó para hacer.

La lección sobrevive a la tabla: una arquitectura que no puede representar algo es un problema distinto de una que nunca aprendió ese rango, y el segundo es el que muerde. También es la maquinaria detrás de cada anuncio de “hemos ampliado el contexto a 128K”: casi siempre son reescalados de una codificación rotatoria, y por eso el Capítulo 16 dice que el límite de contexto se mueve en vez de desaparecer.

Dropout se hereda de la misma forma: aparece en los pesos de attention después de la softmax, en la salida de cada subcapa antes de la suma residual y en la suma de embedding, haciendo exactamente lo que describió el Capítulo 6. En grandes ejecuciones de pretraining suele ponerse a cero, porque un modelo que ve cada token una sola vez no está en posición de sobreajustar.

Dos tensores de la capa tienen forma n×nn \times n, donde nn es el número de tokens: los scores y los pesos después de la softmax. Todo lo demás —cada proyección, toda la MLP— es lineal en nn.

Una capa de attention, 512 de anchura, 8 heads, batch de uno, float32, en una GPU de portátil. Lee las dos columnas de milisegundos solo por sus proporciones: son tiempo real en una tarjeta de portátil de 8 GB que baja de 1.785 MHz a menos de 300 MHz cuando se calienta, así que una ejecución en frío del mismo código vuelve de siete a diez veces más rápido y una ocupada aún más lenta. Las columnas de megabytes son recuentos de bytes del asignador y no se mueven.

TEXT
  tokens   ms total    ms x4   ms projections   attn matrix MB    peak MB    MB x4
     128      2.246        -            1.324              0.5       14.6        -
     256      2.855     1.27            2.113              2.0       19.2     1.31
     512      5.761     2.02            3.105              8.0       34.4     1.79
    1024     16.414     2.85            4.008             32.0       89.1     2.59
    2048     51.573     3.14            9.989            128.0      296.1     3.32
    4096    225.432     4.37           20.176            512.0     1100.1     3.72
    8192    832.838     3.69           40.106           2048.0     4300.1     3.91
   16384   OUT OF MEMORY                                 8192.0

fitted exponent (log-log slope, last four rows):  time ~ n^1.91   memory ~ n^1.87

Las columnas x4 son la proporción respecto a la fila superior, y duplicar nn converge exactamente a 4 tanto para tiempo como para memoria: 3,91 en el último paso frente a un 4 teórico. La columna de proyecciones es el control: 4,0 ms en 1024 tokens frente a 40,1 ms en 8192, un factor de diez para un factor de ocho. Lineal, como se prometió.

Luego la última fila. Una capa de attention, una secuencia, sin modelo alrededor, se queda sin memoria en una GPU de 8 GB con 16.384 tokens: solo la matriz de scores serían 8 GB, al ser 8 heads por 16.384 por 16.384 por 4 bytes. No el modelo; un tensor intermedio en una capa.

Ese es el hecho físico debajo de tres capítulos posteriores. Es por lo que una context window tiene un límite, que el Capítulo 16 convierte en precio. Es por lo que existe FlashAttention, que calcula el mismo resultado en teselas sin almacenar nunca la matriz: una optimización de memoria antes que de velocidad.11 Y es la aritmética detrás del precio de un prompt largo, que el Capítulo 24 paga en un bucle de agent; un asunto separado del otro hallazgo de ese capítulo, que un modelo también usa peor un contexto largo, cosa que mide y se niega a achacar a esta fórmula.

Mostrar detalles

Las dos variantes que reducen la caché, nombradas aquí y pagadas en el Capítulo 13.

La generación cachea las keys y values de los tokens ya procesados: una key y un value por token, por head y por capa. Multi-query attention12 mantiene hh proyecciones de query, pero una sola proyección de key y value compartida por todas las heads, dividiendo esa cache por hh. Grouped-query attention13 interpola: las heads se agrupan, cada grupo comparte una key y un value, así que g=hg = h es attention ordinaria y g=1g = 1 es multi-query. Casi todos los modelos abiertos desde 2023 la usan con 4 u 8 grupos. Ninguna existe por calidad; ambas existen por el tamaño de esa cache, y el Capítulo 13 hace la aritmética que la convierte en “qué modelo cabe en tu GPU”.

El artículo de 2017 describe un encoder-decoder: una pila que lee la fuente con attention sin máscara, una segunda que genera el objetivo de forma causal, y un tercer tipo de attention en medio donde las queries del decodificador se encuentran con las keys del codificador. Eso es correcto para traducción, donde entrada y salida son dos secuencias.

Lo que ganó fue la mitad decoder-only: una pila, causal de principio a fin, entrada y salida en la misma secuencia; y la razón no es la elegancia. “Predecir el siguiente token” funciona sobre cualquier texto, así que el conjunto de entrenamiento es internet en vez de un corpus paralelo, y todo se convierte en esa única tarea: una traducción es un documento que contiene fuente y luego objetivo, una pregunta y su respuesta son un documento, una conversación con una tool call en medio es un documento. El Capítulo 11 trata de cómo se fabrica esa última. Los codificadores no desaparecieron: uno ve toda la entrada de una vez, que es lo que quieres cuando el trabajo es representar un texto en vez de continuarlo, y por eso los embeddings de retrieval del Capítulo 19 vienen de codificadores y no del modelo que chatea.

Con el bloque definido, el tamaño del modelo es aritmética. Por bloque, con anchura dd y una expansión por cuatro: 4d2+4d4d^2 + 4d para Wq,Wk,Wv,WoW_q, W_k, W_v, W_o con sesgos en las cuatro, como los tiene GPT-2; la tabla anterior deja fuera el sesgo en tres de ellas, de ahí 2304 menos por bloque en d=768d = 768; 8d2+5d8d^2 + 5d para la MLP; 4d4d para dos LayerNorms: 12d2+13d12d^2 + 13d, más una tabla de tokens de V×dV \times d y, para posiciones absolutas, nctx×dn_{\text{ctx}} \times d. Para la forma de GPT-2 small —d=768d = 768, 12 bloques, un vocabulario de 50.257, un contexto de 1024, la capa de salida compartiendo pesos con el embedding—:

TEXT
  token embeddings     50,257 x 768 = 38,597,376
  position embeddings   1,024 x 768 =    786,432
  one block                             7,087,872
  12 blocks                            85,054,464
  final LayerNorm         2 x 768 =        1,536
  total (weights tied)                124,439,808

Que es el tamaño publicado de ese modelo. La fórmula no es una aproximación; es el modelo. Observa también que casi un tercio de un modelo pequeño es la tabla de embedding, por eso el tamaño del vocabulario es una decisión arquitectónica y no de preprocesamiento: el intercambio que planteó el Capítulo 7.

La perplejidad es un número sobre un corpus. Lo que hace una head es otra pregunta, y un modelo entrenado en un megabyte de Shakespeare es el instrumento equivocado para ello: lo honesto sobre el mapa de attention de un modelo de 500.000 parámetros es decir que en su mayoría no es interpretable. Así que: un lenguaje donde la pregunta tiene una respuesta correcta.

La ilustración clásica es the animal did not cross the street because it was too tired, donde it es el animal, frente a …because it was too wet, donde una palabra mueve el referente a la calle. Son esquemas de Winograd14: pares de frases idénticas salvo por una palabra, donde esa palabra decide a qué se refiere un pronombre.

También son resolubles haciendo trampas, que es la parte que los tutoriales omiten. Si los dos candidatos son un animal y un lugar, tired y wet identifican el referente por categoría, y un modelo que solo sabe qué palabras están presentes acierta sin saber nada de orden. Medido en esa versión de la tarea, con pares animal/lugar reservados:

TEXT
uniform causal average           held-out referent accuracy 100.0 %
one transformer block            held-out referent accuracy  91.7 %

La bolsa de palabras vence al transformer. Cualquier demostración construida sobre esa frase no prueba nada sobre attention.

Así que cierra el agujero: toma ambos candidatos de un único conjunto de dieciséis sustantivos, cualquiera de los cuales puede aparecer en cualquiera de los dos huecos, y divide los adjetivos por rol en vez de por categoría: cuatro hacen que it sea quien cruza (tired, scared, slow, weak), cuatro hacen que sea lo cruzado (wet, wide, busy, steep).

TEXT
the {x} did not cross the {y} because it was too {adj} , so the {ref} waited .

Entrena como un predictor ordinario del siguiente token, puntúa una posición —la palabra después de so the— y construye el conjunto reservado a partir de pares de sustantivos cuyo orden invertido estuvo en entrenamiento, de modo que cualquier cosa que sepa qué dos sustantivos están presentes pero no cuál vino primero debe responder al revés.

modeloparámetrosreservadonombra el otro sustantivo
solo el token actual57965,2 %5,2 %
media causal uniforme579627,9 %50,0 %
una head de attention aprendida18.08435,4 %64,6 %
cuatro heads22.24475,0 %15,6 %
un bloque transformer55.71692,7 %4,2 %
dos bloques transformer105.508100,0 %0,0 %

El azar entre los dos sustantivos presentes es 50 %. La media uniforme cae en 27,9 % y responde con el sustantivo equivocado del par exactamente la mitad de las veces: la firma de algo que sabe qué palabras hay y nada sobre su orden, tal y como predijo la prueba de barajado tres secciones atrás.

Ahora el mapa: la attention en la posición que debe nombrar el referente, promediada sobre las cuatro heads de cada bloque, para las dos frases que difieren en una palabra. Una media uniforme pondría 0,067 en cada uno de los quince tokens visibles.

TEXT
the animal did not cross the street because it was too tired , so the animal waited .
  blk 1  the:0.00 animal:0.70 did:0.00 not:0.00 cross:0.00 the:0.00 street:0.06
         because:0.00 it:0.00 was:0.00 too:0.00 tired:0.00 ,:0.05 so:0.00 the:0.19
  blk 2  the:0.00 animal:0.00 did:0.00 not:0.00 cross:0.00 the:0.00 street:0.00
         because:0.00 it:0.00 was:0.00 too:0.00 tired:1.00 ,:0.00 so:0.00 the:0.00

the animal did not cross the street because it was too wet , so the street waited .
  blk 1  the:0.00 animal:0.70 did:0.00 not:0.00 cross:0.00 the:0.00 street:0.06
         because:0.00 it:0.00 was:0.00 too:0.00   wet:0.00 ,:0.05 so:0.00 the:0.19
  blk 2  the:0.00 animal:0.00 did:0.00 not:0.00 cross:0.03 the:0.00 street:0.49
         because:0.00 it:0.00 was:0.00 too:0.20   wet:0.03 ,:0.00 so:0.00 the:0.25

El bloque 1 es idéntico en ambas frases: 0,70 sobre el primer sustantivo, sea cual sea el adjetivo. Eso no es un fallo sino una prueba: en la primera capa, la query de una posición es función del propio token y del índice de esa posición, y the en la posición 14 es el mismo token en ambas frases. Una head de primera capa no puede condicionar sobre una palabra que aún no ha traído. Así que el bloque 1 hace lo único útil disponible y arrastra el primer sustantivo hacia delante.

El bloque 2 es donde las frases se separan, y la misma fila a través de los ocho adjetivos muestra la regla que encontró el modelo:

adjetivobloque 2 en animalen streeten el adjetivorespuesta
tired, scared, slow, weak0,0000,0001,000animal
wet, wide, busy, steep0,0000,4910,00–0,03street

Para un adjetivo de quien cruza, el segundo bloque gasta todo su peso en el adjetivo, porque la respuesta ya está en el flujo residual —el bloque 1 la puso ahí— y solo necesita confirmación. Para un adjetivo de lo cruzado, va a buscar el otro sustantivo. Eso es un circuito de dos saltos: una head mueve un candidato hacia delante, una head de una capa posterior lee un token que decide si conservarlo. La composición entre capas es el mecanismo, y por eso un bloque llegó al 92,7 % y dos llegaron al 100 %.

También es la forma del circuito mejor documentado en modelos reales. Induction heads —una head de token anterior que alimenta a una head en la capa siguiente que completa el patrón [A][B] … [A] → [B]— son lo que el trabajo de interpretabilidad de Anthropic identifica detrás de una gran parte del in-context learning, y se forman en un momento identificable durante el pretraining. Este capítulo no intenta ese análisis: queda delegado, con ambos artículos en las referencias, porque leer circuitos de un modelo real es un campo de investigación y no una sección.

Por último, la implementación. Las treinta líneas de arriba, con sus pesos copiados de los propios de PyTorch:

TEXT
ours vs nn.MultiheadAttention           max |diff| = 1.7881393432617188e-07
ours vs F.scaled_dot_product_attention  max |diff| = 1.7881393432617188e-07

1.8×1071.8 \times 10^{-7} en salidas cuya magnitud media es 0,159: la misma aritmética en un orden distinto, con precisión float32.

Tienes la arquitectura de la que parte cada modelo del resto de este curso, y es más pequeña que su reputación: una media ponderada cuyos pesos se aprenden, una MLP por posición que sostiene dos tercios de los parámetros, dos normalizaciones y dos sumas, apiladas.

Lo que no tienes es un modelo que sepa algo, y apilar no lo arreglará por sí solo. Dos bloques en este corpus alcanzan una perplejidad de entrenamiento de 14,49 y una perplejidad de validación de 40,57, frente a 18,77 y 38,07 de un bloque: más capacidad, mejor en lo que ha visto, peor en lo que no ha visto, que es la tabla del Capítulo 6 con un transformer dentro. La distancia entre este modelo y aquellos con los que hablan los capítulos 14 a 30 no es arquitectónica. Es el mismo bloque, repetido más veces, sobre muchísimo más texto.

Eso lo convierte en un problema de contabilidad, y la contabilidad es más extraña de lo que parece. ¿Cuánto texto, y de dónde lo saca alguien? ¿Cuánta aritmética, y cómo la estimas antes de gastar el dinero? Dado un presupuesto fijo, ¿es mejor hacer el modelo más grande o enseñarle más datos? ¿Y hay una respuesta correcta, o solo una moda? El Capítulo 10 responde a las tres por medición, y pone precio a la forma útil más barata de la pregunta: ¿cuánto cuesta, hoy, entrenar desde cero un modelo como GPT-2?


Tres explicaciones de este material son mejores que esta en aquello para lo que sirven, y este capítulo está escrito para leerse junto a ellas. The Illustrated Transformer, de Jay Alammar, es la mejor imagen del flujo de datos que se ha dibujado. The Annotated Transformer, de Harvard NLP, es el artículo de 2017 con código ejecutable intercalado línea a línea. Let's build GPT: from scratch, in code, spelled out, de Andrej Karpathy, construye el mismo modelo en directo en dos horas, y la escalera de ablaciones de arriba es la misma columna vertebral medida en otro corpus. Para la cuestión de interpretabilidad que este capítulo solo roza, las fuentes primarias son Elhage et al., A Mathematical Framework for Transformer Circuits (2021), y Olsson et al., In-context Learning and Induction Heads (2022), ambos del grupo de interpretabilidad de Anthropic.

  1. Hochreiter, S. and Schmidhuber, J. Long Short-Term Memory. Neural Computation 9(8), pp. 1735–1780 (1997).

  2. Sutskever, I., Vinyals, O. and Le, Q. V. Sequence to Sequence Learning with Neural Networks. arXiv:1409.3215 (2014). El encoder-decoder cuyo único vector de contexto es el cuello de botella.

  3. Bahdanau, D., Cho, K. and Bengio, Y. Neural Machine Translation by Jointly Learning to Align and Translate. arXiv:1409.0473 (2014). Attention, tres años antes del transformer.

  4. La perplejidad es el exponencial de la entropía cruzada media por token, del Capítulo 8. Cada número aquí usa el mismo tokenizador y la misma partición de validación, que es la única condición bajo la cual dos perplejidades pueden compararse.

  5. Vaswani, A., Shazeer, N., Parmar, N., Uszkoreit, J., Jones, L., Gomez, A. N., Kaiser, Ł. and Polosukhin, I. Attention Is All You Need. arXiv:1706.03762 (2017). La sección 3.2.1 es la única frase sobre dk\sqrt{d_k} que este capítulo dedica una sección a medir.

  6. Shazeer, N., Mirhoseini, A., Maziarz, K., Davis, A., Le, Q., Hinton, G. and Dean, J. Outrageously Large Neural Networks: The Sparsely-Gated Mixture-of-Experts Layer. arXiv:1701.06538 (2017).

  7. Ba, J. L., Kiros, J. R. and Hinton, G. E. Layer Normalization. arXiv:1607.06450 (2016). Introducida y medida en el Capítulo 6; usada aquí sin cambios.

  8. Xiong, R., Yang, Y., He, D., Zheng, K., Zheng, S., Xing, C., Zhang, H., Lan, Y., Wang, L. and Liu, T.-Y. On Layer Normalization in the Transformer Architecture. arXiv:2002.04745 (2020). El análisis de gradient detrás de pre-norm, y el argumento de que warmup es un síntoma.

  9. Su, J., Lu, Y., Pan, S., Murtadha, A., Wen, B. and Liu, Y. RoFormer: Enhanced Transformer with Rotary Position Embedding. arXiv:2104.09864 (2021).

  10. Press, O., Smith, N. A. and Lewis, M. Train Short, Test Long: Attention with Linear Biases Enables Input Length Extrapolation. arXiv:2108.12409 (2021). El resultado de extrapolación reproducido arriba.

  11. Dao, T., Fu, D. Y., Ermon, S., Rudra, A. and Ré, C. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. arXiv:2205.14135 (2022).

  12. Shazeer, N. Fast Transformer Decoding: One Write-Head is All You Need. arXiv:1911.02150 (2019).

  13. Ainslie, J., Lee-Thorp, J., de Jong, M., Zemlyanskiy, Y., Lebrón, F. and Sanghai, S. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. arXiv:2305.13245 (2023).

  14. Levesque, H. J., Davis, E. and Morgenstern, L. The Winograd Schema Challenge. KR (2012). La construcción detrás de la frase animal / street que usa cada tutorial de attention.

¿Listo para dejar que elija LIA?

Crea con todos los modelos de IA en un mismo sitio. Empieza gratis hoy.