La raíz que faltaba en la atención

La raíz que faltaba en la atención

28 min read

La auto-atención quedó montada del todo salvo por un símbolo. Las tres proyecciones tienen su argumento, la rejilla cuadrada tiene su forma y el softmax por filas reparte una unidad entre las posiciones; el dk\sqrt{d_k} del denominador entró en la fórmula sin nada detrás y así sigue. Es el único préstamo que sigue abierto, y esta lección lo salda.

La pregunta tiene más filo del que parece. ¿Por qué una raíz, y no el propio dkd_k? ¿Y por qué de dkd_k, y no de dmodeld_{\text{model}} o de la longitud de la frase? Las dos se contestan de una vez, porque las dos dependen de averiguar cuánto se aleja del cero una suma cuando se le van añadiendo sumandos. Nada de lo que viene necesita más probabilidad que un promedio, y el promedio se hace sobre una lista finita.

Toma la versión más cruda posible de una consulta (query) y una clave (key): todas sus coordenadas valen +1+1 o 1-1. La puntuación es entonces una cuenta de coincidencias, porque cada término aporta +1+1 si las dos coordenadas comparten signo y 1-1 si no. Con cuatro coordenadas, tres coincidencias y una discrepancia dan 31=23 - 1 = 2. Con cuatrocientas, la misma proporción —trescientas contra cien— da 200200.

Y ahí está lo que hay que ver, que no es que los términos se hayan hecho mayores: siguen valiendo ±1\pm 1. Es que hay cien veces más. Una desviación pequeña respecto de la mitad exacta, repetida sobre más sumandos, se convierte en un número grande, y ese número se lo va a comer una exponencial.

Arriba, seis puntuaciones con signo alrededor de un eje. Abajo, cuatro grupos de seis barras: el reparto tras dividir, y los repartos sin dividir con d_k igual a 4, 16 y 64. La barra mayor pasa de destacar poco a ocuparlo todo.
El mismo reparto deja de ser un reparto. Estirar las seis puntuaciones no cambia cuál gana, cambia cuánto gana, y con d_k = 64 el softmax ya no mezcla: se queda con una posición y tira las otras cinco.

Una puntuación crece con la raíz de la anchura

Las consultas y las claves de una capa viven en Rdk\mathbb{R}^{d_k}, y la puntuación de una pareja es su producto escalar:

e=qk=m=1dkqmkm,e = \mathbf{q}^{\top}\mathbf{k} = \sum_{m=1}^{d_k} q_m k_m,

una suma de dkd_k términos. De qué tamaño sale depende de qué tamaño entra, así que fijemos primero el caso exacto: las coordenadas valen +1+1 o 1-1, como en la cuenta de coincidencias de arriba. Llamemos sm=qmkms_m = q_m k_m al término mm-ésimo, que vale entonces +1+1 o 1-1, y e=msme = \sum_m s_m.

Elevar al cuadrado separa la diagonal del resto:

e2=(msm)2=msm2+mnsmsn=dk+mnsmsn,e^{2} = \left(\sum_{m} s_m\right)^{2} = \sum_{m} s_m^{2} + \sum_{m \neq n} s_m s_n = d_k + \sum_{m \neq n} s_m s_n,

porque sm2=1s_m^{2} = 1 sea cual sea el signo. El primer sumando ya no depende de nada. El segundo es el que hay que promediar, y promediar aquí quiere decir algo perfectamente finito: recorrer las 2dk2^{d_k} combinaciones de signos que puede tomar q\mathbf{q}, con k\mathbf{k} fija, y hacer la media. La escribiremos E[]\mathbb{E}\left[\cdot\right] —la media de una cantidad sobre el azar que se le supone—. Los términos cruzados se van todos.

Ver por qué cada término cruzado promedia cero

Fija dos índices distintos mnm \neq n y mira el término smsn=qmkmqnkns_m s_n = q_m k_m q_n k_n. Empareja cada combinación de signos con la que resulta de cambiarle el signo a qmq_m y dejar las otras dk1d_k - 1 coordenadas donde están. Es un emparejamiento perfecto de las 2dk2^{d_k} combinaciones en 2dk12^{d_k - 1} parejas: cada combinación tiene exactamente una compañera, y no es ella misma.

Dentro de una pareja, qnq_n, kmk_m y knk_n no se han movido y qmq_m ha cambiado de signo, de modo que los dos valores de smsns_m s_n son opuestos y suman cero. Sumar cero 2dk12^{d_k - 1} veces y dividir entre 2dk2^{d_k} deja

E[smsn]=0,mn.\mathbb{E}\left[s_m s_n\right] = 0, \qquad m \neq n.

El mismo emparejamiento da E[e]=0\mathbb{E}\left[e\right] = 0: en e=msme = \sum_m s_m cada sumando cambia de signo junto con su propia coordenada, así que cada uno promedia cero por separado.

De modo que del cuadrado sobrevive el primer sumando y nada más:

E[e2]=dk.\mathbb{E}\left[e^{2}\right] = d_k.

Con E[e]=0\mathbb{E}\left[e\right] = 0, eso es cuanto hace falta. La puntuación se reparte alrededor del cero, y lo que mide su alejamiento típico —la desviación típica— es la raíz de ese promedio:

E[e2]=dk,\sqrt{\mathbb{E}\left[e^{2}\right]} = \sqrt{d_k},

que es la respuesta a la mitad de la pregunta. Dividir entre dk\sqrt{d_k} deja E[(e/dk)2]=1\mathbb{E}\left[\left(e/\sqrt{d_k}\right)^{2}\right] = 1, y en ese 11 no aparece dkd_k: la misma escala con consultas de ocho coordenadas y con consultas de mil.

Con el divisor puesto, la operación entera —producto, división, softmax y mezcla— pasa a tener nombre propio: es el producto interno escalado, que es como el artículo escribe scaled dot-product. El producto interno es el producto escalar de siempre, el de la línea de arriba; lo escalado es esta división y nada más.

Los signos estaban ahí para que la cuenta fuese exacta, pero fíjate en qué propiedades le hicieron falta de verdad. Dos: que E[sm]=0\mathbb{E}\left[s_m\right] = 0 y que E[sm2]=1\mathbb{E}\left[s_m^{2}\right] = 1. Cualquier inicialización centrada cuyas coordenadas sean independientes y de tamaño típico 11 las cumple —las gaussianas de rng.normal que el curso lleva usando desde la lección sobre implementar el MLP (multilayer perceptron), entre ellas—, y con ellas la misma cancelación deja el mismo E[e2]=dk\mathbb{E}\left[e^{2}\right] = d_k. La primera celda lo mide.

Conviene decir en voz alta de qué habla esa hipótesis, porque no habla del modelo que te encontrarás. Habla de pesos recién inicializados: nada garantiza que después de miles de pasos las coordenadas de q\mathbf{q} sigan siendo independientes ni midiendo 11. Por eso el divisor se fija al diseñar la capa y no se mide sobre la marcha —arregla el arranque, que es donde el problema detiene el entrenamiento antes de que llegue a empezar—.

Queda la parte que suele darse por sabida: por qué la raíz y no dkd_k. El divisor tiene que quitarle a la puntuación exactamente lo que la anchura le puso, y lo que le puso es un factor dk\sqrt{d_k}. Dividir entre dkd_k dejaría las puntuaciones en un tamaño típico de 1/dk1/\sqrt{d_k}, cada vez más cerca del cero cuanto más ancho el modelo; y unas puntuaciones que tienden a cero dan un softmax que tiende al reparto uniforme, 1/T1/T en cada casilla, con lo que la fila deja de elegir nada. Son dos fallos, uno a cada lado, y la raíz es el único exponente que no cae en ninguno.

Una fila cerrada no devuelve gradiente

Si las puntuaciones miden dk\sqrt{d_k}, las diferencias entre las de una misma fila miden eso también, y lo que el softmax mira son las diferencias:

αij=exp(eij)m=1Texp(eim),\alpha_{ij} = \frac{\exp\left(e_{ij}\right)}{\sum_{m=1}^{T}\exp\left(e_{im}\right)},

con mm recorriendo las posiciones de la fila ii. Dos puntuaciones separadas por 88 —una distancia corriente cuando el tamaño típico es 88, que es lo que da dk=64d_k = 64— quedan a un factor exp(8)2981\exp(8) \approx 2\,981 la una de la otra, y con eso el reparto se ha terminado: una posición se lleva casi la unidad entera.

Que una fila elija fuerte no suena a defecto. El defecto está en la derivada.

Deriva αij\alpha_{ij} respecto de las puntuaciones de su propia fila. Es la misma cuenta que la lección sobre las funciones de pérdida hizo con el softmax de la capa de salida, y salen los mismos dos casos según el índice coincida o no:

αijeij=αij(1αij),αijeim=αijαim(mj).\frac{\partial \alpha_{ij}}{\partial e_{ij}} = \alpha_{ij}\left(1 - \alpha_{ij}\right), \qquad \frac{\partial \alpha_{ij}}{\partial e_{im}} = -\alpha_{ij}\,\alpha_{im} \quad (m \neq j).

Pon ahora la fila cerrada, con un peso en casi 11 y el resto en casi 00. El caso de la diagonal da 11 por 00; el cruzado, en cualquier casilla, arrastra al menos un factor casi nulo. Las dos derivadas se apagan a la vez, y con ellas la fila entera.

Lo que eso cuesta se ve mirando qué hay antes de las puntuaciones. Sólo WQ\mathbf{W}^Q y WK\mathbf{W}^K: la rejilla sale de QK\mathbf{Q}\mathbf{K}^{\top} y no entra nada más. Si el gradiente no atraviesa el softmax, esas dos matrices no reciben nada por esa fila, y son precisamente las que tendrían que aprender a quién debe mirar cada posición. La capa sigue calculando —una salida saturada es una salida perfectamente definida—; lo que ha dejado de poder es cambiar de opinión.

Merece la pena pararse en por qué el bloque 2 no se topó con esto, porque la diferencia es de una línea. Allí el softmax iba seguido de la entropía cruzada, y el log-\log de la pérdida aporta el factor inverso exacto: las dos juntas dieron z=y^y\nabla_{\mathbf{z}}\,\ell = \hat{\mathbf{y}} - \mathbf{y}, que vale y^c1\hat{y}_c - 1 en la clase correcta y no se anula por muy segura que esté la red. Aquí detrás del softmax no hay ningún logaritmo, sólo la multiplicación por V\mathbf{V}, la lista de valores (values). La jacobiana no se cancela contra nada: se queda entera, ceros incluidos.

Y así reaparecen dos cosas que el curso dio por cerradas en bloques distintos. La curva α(1α)\alpha\left(1 - \alpha\right) es la misma que la derivada de la sigmoide, con su máximo de 0.250.25 en la mitad y sus dos colas planas: la saturación de la lección sobre las funciones de activación, en una capa donde no hay ninguna sigmoide. Y un gradiente que llega apagado a las matrices que tenían que aprender es el bloque 3 entero, en una arquitectura sin una sola recurrencia. Lo que las une no era la sigmoide ni el estado recurrente: era que en las dos había una función aplanándose.

Una fila entera se resume en un número, y sale de sumar el caso de la diagonal sobre las TT posiciones:

j=1Tαij(1αij)=jαijjαij2=1j=1Tαij2,\sum_{j=1}^{T} \alpha_{ij}\left(1 - \alpha_{ij}\right) = \sum_{j} \alpha_{ij} - \sum_{j} \alpha_{ij}^{2} = 1 - \sum_{j=1}^{T} \alpha_{ij}^{2},

donde la última igualdad usa que la fila suma 11. Con la fila repartida por igual vale 11/T1 - 1/T; con la fila entera en una posición vale 00. Es la suma de la diagonal de la jacobiana, y mide lo mismo por los dos lados: cuánto reparte esa fila, y cuánto le queda por aprender.

El divisor, medido en NumPy

La primera celda saca la cuenta del caso de los signos y la comprueba donde de verdad se aplica: coordenadas gaussianas, cinco anchuras, la columna del promedio medido contra la columna de dk\sqrt{d_k}.

import numpy as np

rng = np.random.default_rng(7)
N = 8000 # parejas de vectores por cada anchura

print(" d_k media tamano tipico raiz(d_k) tras dividir")
for d_k in (4, 16, 64, 256, 1024):
Q = rng.normal(size=(N, d_k)) # coordenadas centradas, de tamano tipico 1
K = rng.normal(size=(N, d_k))
e = np.sum(Q * K, axis=1) # una puntuacion por pareja
tipico = np.sqrt(np.mean(e ** 2))
print(f"{d_k:5d} {e.mean():7.3f} {tipico:13.2f} {np.sqrt(d_k):9.2f} {tipico / np.sqrt(d_k):12.2f}")
numpy

La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador.

Las dos columnas centrales coinciden hasta donde llega la precisión de medir con ocho mil parejas, y la de la derecha se queda clavada en 11: dividir deja el tamaño donde estaba, sea cual sea la anchura. La columna de la media es pequeña frente al tamaño típico de su fila —a dk=1024d_k = 1\,024, 0.2360.236 contra 32.0332.03—, que es lo que dice E[e]=0\mathbb{E}\left[e\right] = 0 cuando se mide en vez de demostrarse.

La segunda celda pasa de las puntuaciones a lo que sale de ellas. Mira sobre todo la última columna, que es la que decide si la capa puede aprender.

import numpy as np

rng = np.random.default_rng(11)
T, N = 8, 4000 # 8 posiciones por fila, 4000 filas medidas


def softmax_filas(E):
Z = np.exp(E - E.max(axis=1, keepdims=True))
return Z / Z.sum(axis=1, keepdims=True)


print(" d_k peso mayor reparto mayor a(1-a)")
for d_k in (8, 64, 512):
Q = rng.normal(size=(N, d_k)) # N consultas contra las mismas T claves
K = rng.normal(size=(T, d_k))
E = Q @ K.T
for nombre, Esc in (("sin dividir", E), ("dividiendo", E / np.sqrt(d_k))):
A = softmax_filas(Esc)
print(f"{nombre:>17} {d_k:5d} {A.max(axis=1).mean():10.3f} "
f"{(1 - (A ** 2).sum(axis=1)).mean():7.3f} {(A * (1 - A)).max(axis=1).mean():12.4f}")
numpy

La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador.

Léela por parejas de líneas. Sin dividir, el peso mayor sube de 0.6700.670 a 0.9530.953 y el reparto 1jαij21 - \sum_j \alpha_{ij}^{2} se hunde de 0.4450.445 a 0.0660.066 —el máximo posible con ocho posiciones es 11/8=0.8751 - 1/8 = 0.875—, mientras la mayor casilla de la diagonal de la jacobiana cae de 0.17550.1755 a 0.03210.0321: más de cinco veces menos gradiente por el mero hecho de haber ensanchado el modelo. Dividiendo, las tres columnas dejan de moverse: 0.360.36 de peso mayor, 0.770.77 de reparto y 0.2160.216 de derivada en las tres anchuras, incluida la de 512512. Ésa es la propiedad que se compra con el divisor, y no es que los números sean mejores, es que dejan de depender de dkd_k.

Comprueba tu intuición

Cinco preguntas: cuánto mide una puntuación, por qué la raíz y no otra cosa, qué pierde una fila cerrada, por qué el bloque 2 se libró de esto y qué es lo que dividir no toca.

Las consultas y las claves de una capa tienen dk=64d_k = 64 coordenadas, centradas e independientes, de tamaño típico 11. Sin dividir, ¿cuánto mide típicamente una puntuación qk\mathbf{q}^{\top}\mathbf{k}?

A margin of ±0 is accepted.

¿Por qué el divisor es dk\sqrt{d_k} y no dkd_k?

Una fila del mapa se ha ido casi entera a una posición: αi11\alpha_{i1} \approx 1 y las demás 0\approx 0. Marca lo que es cierto.

Select every correct option. This is graded all-or-nothing: there is no partial credit.

El bloque 2 también puso un softmax en la red y nunca habló de que saturase. ¿Por qué no le pasaba esto?

Dividir las puntuaciones de una fila entre dk\sqrt{d_k} puede cambiar cuál de las posiciones se lleva el peso mayor.

Escribe gradiente_puntuaciones(A, G), la vuelta del softmax por filas. Recibe el mapa de pesos ya normalizado, A de forma (T,T)(T, T) con cada fila sumando 11, y el gradiente que baja desde arriba, G de la misma forma, con gij=/αijg_{ij} = \partial\ell/\partial\alpha_{ij}. Devuelve /e\partial\ell/\partial e, de forma (T,T)(T, T).

Los dos casos de la derivada están desarrollados arriba; encadenados dan, casilla a casilla, eij=αij(gijmgimαim)\dfrac{\partial\ell}{\partial e_{ij}} = \alpha_{ij}\left(g_{ij} - \sum_{m} g_{im}\alpha_{im}\right), donde mm recorre la fila ii.

Escríbelo con operaciones sobre arrays —sin recorrer las casillas una a una y sin construir ninguna jacobiana— y no escribas en los arrays que recibes.

The first run downloads the Python interpreter (~15 MB); after that it stays in the browser cache. This challenge is much easier to solve on a physical keyboard: on a phone, read it and come back later.


Con el divisor puesto, la escala de las puntuaciones deja de depender de la anchura, y el reparto y el gradiente dejan de depender de ella con la misma cuenta. Pero mira lo que el divisor da por hecho: recibe dkd_k ya decidido. Nada de lo que hemos contado dice cuánto debería medir una consulta —sólo cómo compensar lo que mida—, y hasta aquí dkd_k es un número que apareció porque las proyecciones de la lección anterior tenían que salir a alguna anchura.

Hay una segunda cosa que esta capa no sabe hacer, y apunta al mismo sitio. Una fila de A\mathbf{A} es un reparto y sólo uno: en las llaves del coche están ahí, la posición de están tiene que decidir con un único juego de pesos si mira a llaves, que es con quien concuerda, o a del coche, que es lo que la separa de ella; y con una fila no puede hacer las dos cosas. La respuesta a las dos preguntas es la misma y ocupa la lección siguiente, sobre la atención multi-head: varias parejas de proyecciones trabajando en paralelo, cada una con su propio reparto y su propia anchura, repartida entre todas a partir de dmodeld_{\text{model}}. Ahí dkd_k deja de ser un número dado y pasa a ser una decisión con un motivo.

Further reading1 source · 1 paper

Where this lesson comes from, and where to go next. None of it is needed to carry on with the course.

  • Attention Is All You Need
    paperVaswani, Shazeer, Parmar y otros, 2017arXiv:1706.03762EN

    Su §3.2.1 y su nota al pie 4 son esta lección en dos frases: con q y k de media 0 y varianza 1, q·k tiene varianza d_k, y de ahí √d_k. No llega a la saturación del softmax ni al gradiente que se apaga; eso es tuyo.