Varias cabezas, varios repartos

Varias cabezas, varios repartos

29 min read

La lección anterior justificó el dk\sqrt{d_k} de la fórmula de la atención contando cuánto crece un producto escalar cuando se le añaden coordenadas, y al cerrar esa cuenta dejó dos cosas señaladas. La primera es de forma: una capa de auto-atención no le da a cada posición más que una fila de A\mathbf{A}, fabricada por un único par de matrices, WQ\mathbf{W}^Q y WK\mathbf{W}^K —una manera de mirar la frase, la misma para todas las posiciones y para todas las relaciones que haya que encontrar entre ellas—. La segunda es más pequeña y sale del mismo sitio: dkd_k entra en la fórmula sin que nada diga cuánto debería valer.

Mira lo que esa única fila le pide a una frase como el ratón pequeño que persiguen los gatos duerme. duerme tiene que encontrar a ratón por encima de la subordinada entera, y persiguen, tres posiciones antes, tiene que encontrar a gatos, que le queda detrás: dos verbos, dos sujetos, en direcciones opuestas. Ninguna relación sola los resuelve —quién es sustantivo no distingue ratón de gatos, y quién concuerda en número no distingue un sujeto de un determinante—, así que la única fila que hay tiene que llevar las dos a la vez. Y puede: las puntuaciones se suman antes del softmax, y la suma decide. Esta lección trata de lo que esa suma cuesta, y de la alternativa, que consiste en no hacerla.

Las rejillas de abajo son esa misma frase vista por cuatro reglas distintas, una por cabeza, y en el quinto botón por la capa de la lección anterior, que lleva las cuatro sumadas en una sola fila. Ponte en la fila de duerme y lee el panel de abajo entero antes de tocar nada: la primera cabeza empata ratón con gatos a 0.300.30; la segunda descarta gatos —le deja 0.010.01— pero reparte su peso mayor entre tres singulares y ninguno de los tres es el sujeto; dos no dicen nada de esa fila; y la capa de una sola cabeza contesta ratón. Pásate luego a la fila de persiguen: la primera cabeza repite su fila anterior número a número, porque no sabe distinguir un verbo de otro; la segunda se da la vuelta entera; y la respuesta de la capa de una cabeza cambia a gatos.

Cuatro reglas, cuatro repartos, y un quinto botón con las cuatro sumadas en una sola fila. La cabeza del verbo empata «ratón» con «gatos» en las dos filas de verbo; la de concordancia sí los separa, pero se queda con tres palabras que no son el sujeto. Sólo juntas contestan, y juntarlas de esa manera no sale gratis.

La misma matriz leída por varias atenciones a la vez

Cada cabeza es la capa de la lección sobre la auto-atención, entera: sus tres proyecciones sobre la misma matriz y su propio mapa. Lo único que hay que añadir es el subíndice que allí no hizo falta escribir porque sólo había una.

headi=Attention(XWiQ,  XWiK,  XWiV)RT×dv,\text{head}_i = \text{Attention}\left(\mathbf{X}\mathbf{W}^Q_i,\; \mathbf{X}\mathbf{W}^K_i,\; \mathbf{X}\mathbf{W}^V_i\right) \in \mathbb{R}^{T \times d_v},

con i=1,,hi = 1, \dots, h, y con WiQ,WiKRdmodel×dk\mathbf{W}^Q_i, \mathbf{W}^K_i \in \mathbb{R}^{d_{\text{model}} \times d_k} y WiVRdmodel×dv\mathbf{W}^V_i \in \mathbb{R}^{d_{\text{model}} \times d_v}. Fíjate en que las hh leen la misma X\mathbf{X}: no se reparten la frase, ni se reparten las coordenadas de entrada. Cada una saca de ella sus propias consultas (queries), claves (keys) y valores (values), y con ellas su propio mapa AiRT×T\mathbf{A}_i \in \mathbb{R}^{T \times T}, cuadrado y con las filas sumando 11 igual que el de la lección sobre la auto-atención.

Quedan entonces hh salidas de T×dvT \times d_v y la capa tiene que devolver una. Se ponen una al lado de otra y se proyectan:

MultiHead(X)=Concat(head1,,headh)WORT×dmodel,\text{MultiHead}(\mathbf{X}) = \text{Concat}\left(\text{head}_1, \dots, \text{head}_h\right)\mathbf{W}^O \in \mathbb{R}^{T \times d_{\text{model}}},

donde Concat\text{Concat} pega las hh matrices por columnas —una fila por posición, con los hh trozos seguidos dentro de ella, de modo que el resultado es T×(hdv)T \times \left(h \cdot d_v\right)— y WOR(hdv)×dmodel\mathbf{W}^O \in \mathbb{R}^{\left(h \cdot d_v\right) \times d_{\text{model}}} es la única matriz de la capa que no pertenece a ninguna cabeza. Su superíndice es una etiqueta de papel, como los de las otras tres.

Dónde se normaliza

La puntuación de una cabeza ya era una suma, y la lección anterior la escribió así precisamente para poder contarla:

eij=1dkm=1dkqimkjm,e_{ij} = \frac{1}{\sqrt{d_k}}\sum_{m=1}^{d_k} q_{im}k_{jm},

donde qimq_{im} es la coordenada mm de la consulta de la posición ii. Cada coordenada de la comparación aporta su término, y el softmax no ve los términos: ve el total. Dos coordenadas que midan dos cosas distintas llegan sumadas a la normalización, y lo que salga de esa suma es lo único que la fila tiene para repartir. Un término grande se lleva la fila; lo que otra coordenada tenía que decir de esa misma pareja desaparece, y no queda dónde ir a buscarlo.

El explorable de arriba es ese caso con dk=4d_k = 4 y una regla por coordenada, así que la cuenta se puede seguir número a número. Vuelve a las llaves del coche están ahí, la frase de la lección sobre la auto-atención, y a su fila de llaves. La regla que lleva un sustantivo hacia su determinante y sus modificadores, ella sola, pone en esa fila su peso mayor sobre ahí, con 0.400.40, y deja las en 0.230.23. Añádele la de concordancia, que carga sobre las y están por ser los plurales de la frase, y la misma fila cambia de ganador: las sube a 0.370.37 y ahí baja a 0.170.17. Ninguna de las dos reglas ha cambiado un solo número. Lo que ha pasado es que se sumaron antes de normalizar.

Ahí está la operación entera. Separar las coordenadas en hh grupos y darle a cada grupo su propio softmax es normalizar antes de combinar, en vez de combinar antes de normalizar. Salen hh distribuciones en lugar de una, cada una con su denominador, y llegan hasta la salida sin tocarse; lo que las combina es WO\mathbf{W}^O, que es una matriz que se aprende y que las recibe ya repartidas. La suma que la capa de una cabeza hacía dentro del softmax era irreversible; hecha fuera, todavía se puede decidir cuánto pesa cada parte.

El reparto de la anchura

Queda elegir dkd_k, que era la otra pregunta abierta, y se elige contando. Cada cabeza cuesta lo que costaba la capa entera de la lección sobre la auto-atención:

dmodel(2dk+dv)paraˊmetros,d_{\text{model}}\left(2d_k + d_v\right) \quad \text{parámetros},

hay hh de ellas, y encima está WO\mathbf{W}^O con hdvdmodelh \cdot d_v \cdot d_{\text{model}} más, de modo que la capa completa pide

hdmodel(2dk+dv)+hdvdmodel.h \cdot d_{\text{model}}\left(2d_k + d_v\right) + h \cdot d_v \cdot d_{\text{model}}.

Asigna ahora a dkd_k y a dvd_v el valor dmodel/hd_{\text{model}}/h —repartir la anchura del modelo entre las cabezas, en vez de darle a cada una la anchura entera— y mira lo que ocurre:

hdmodel3dmodelh+hdmodelhdmodel=3dmodel2+dmodel2=4dmodel2.h \cdot d_{\text{model}} \cdot \frac{3d_{\text{model}}}{h} + h \cdot \frac{d_{\text{model}}}{h} \cdot d_{\text{model}} = 3d_{\text{model}}^{2} + d_{\text{model}}^{2} = 4d_{\text{model}}^{2}.

La hh se ha ido. Ocho cabezas cuestan lo mismo que una, y dieciséis lo mismo que ocho. Las multiplicaciones se comportan igual: la lección sobre el adiós a la recurrencia midió una atención en T2(dk+dv)T^{2}\left(d_k + d_v\right), y hh de ellas dan hT2(dk+dv)=2T2dmodelh \cdot T^{2}\left(d_k + d_v\right) = 2T^{2}d_{\text{model}}, que es exactamente lo que costaba una sola de anchura dmodeld_{\text{model}}.

De modo que hh se elige por lo que hace y no por lo que cuesta, y dkd_k deja de ser un número suelto: es lo que queda de dmodeld_{\text{model}} al partirlo. Con las cifras del artículo, dmodel=512d_{\text{model}} = 512 y h=8h = 8, sale dk=64d_k = 64 —el número con el que la lección anterior hizo sus cuentas— y la capa entera pesa 45122=10485764 \cdot 512^{2} = 1\,048\,576 parámetros.

Lo que se paga conviene decirlo, porque no es cero. Cada cabeza compara en 6464 coordenadas donde una sola habría comparado en 512512: no son ocho comparaciones tan finas como la que había, son ocho comparaciones ocho veces más estrechas. El trato es ése, y es una apuesta sobre el lenguaje —que compensan varias relaciones gruesas frente a una sola muy fina— y no un teorema.

Concatenar y proyectar, no promediar

La alternativa evidente a concatenar es promediar: dale a cada cabeza la anchura entera, dv=dmodeld_v = d_{\text{model}}, y devuelve 1hiheadi\frac{1}{h}\sum_i \text{head}_i, que ya tiene la forma correcta sin necesidad de ninguna matriz más. Merece la pena verla como lo que es, un caso particular de lo que la capa ya hace. Toma la WO\mathbf{W}^O que apila hh copias de 1hI\frac{1}{h}\mathbf{I}:

WO=1h[II]Concat(head1,,headh)WO=1hi=1hheadi.\mathbf{W}^O = \frac{1}{h}\begin{bmatrix}\mathbf{I} \\ \vdots \\ \mathbf{I}\end{bmatrix} \quad\Longrightarrow\quad \text{Concat}\left(\text{head}_1, \dots, \text{head}_h\right)\mathbf{W}^O = \frac{1}{h}\sum_{i=1}^{h}\text{head}_i.

Concatenar con una WO\mathbf{W}^O que se aprende puede, entonces, dar el promedio y cualquier otra cosa; promediar es fijar de antemano un punto de todo eso. Y fijarlo cuesta algo muy concreto: obliga a las hh cabezas a escribir en las mismas dmodeld_{\text{model}} coordenadas, de manera que dos cabezas que han encontrado dos cosas distintas pueden cancelarse antes de que nadie las lea. Concatenadas, cada una tiene sus dvd_v coordenadas para ella sola, y WO\mathbf{W}^O decide después, coordenada de salida por coordenada de salida, a cuál hace caso.

Esa matriz hace además una segunda cosa que aquí sólo queda señalada: devuelve la salida a dmodeld_{\text{model}}, la anchura con la que entró. Eso es lo que permite apilar estas capas una sobre otra, y quien lo necesita de verdad es una lección más adelante, sobre el bloque completo con residuales y layer norm.

Falta decir lo que esta construcción no hace, porque los dibujos de los artículos invitan a creer lo contrario. Nada en la pérdida le asigna un papel a una cabeza: no existe ningún término que diga «la cabeza 2 se ocupa de la concordancia». Las cabezas hacen cosas distintas porque empiezan distintas, y por nada más. Los nombres del explorable de arriba se los he puesto yo, escribiendo las cuatro reglas a mano; en una capa entrenada no los pone nadie y hay que salir a buscarlos después, cuando están. Ni siquiera el orden significa nada: intercambia dos cabezas junto con sus dos bloques de filas de WO\mathbf{W}^O y no cambia un número de la salida, así que «la cabeza 3» no nombra lo mismo en dos entrenamientos del mismo modelo.

Las cabezas en NumPy

La primera celda monta la capa con las proyecciones enteras y un bloque de columnas por cabeza, que es como se implementa de verdad, y después pasa la cuenta de parámetros por cinco valores de hh.

import numpy as np

T, d_model, h = 6, 8, 4
d_k = d_v = d_model // h
rng = np.random.default_rng(5)
X = rng.normal(size=(T, d_model))
Wq = rng.normal(size=(d_model, d_model)) * 0.6 # las h cabezas, una al lado de otra
Wk = rng.normal(size=(d_model, d_model)) * 0.6
Wv = rng.normal(size=(d_model, d_model)) * 0.6
Wo = rng.normal(size=(h * d_v, d_model)) * 0.6


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


salidas, mapas = [], []
for i in range(h):
col = slice(i * d_k, (i + 1) * d_k) # el bloque de columnas de la cabeza i
Q, K, V = X @ Wq[:, col], X @ Wk[:, col], X @ Wv[:, col]
A = softmax_filas(Q @ K.T / np.sqrt(d_k))
salidas.append(A @ V)
mapas.append(A)

concat = np.concatenate(salidas, axis=1) # (T, h*d_v): los h trozos seguidos
Y = concat @ Wo

print("una cabeza", salidas[0].shape, " concatenadas", concat.shape, " salida", Y.shape)
print("mapas:", len(mapas), "de", mapas[0].shape,
" cada fila suma 1:", bool(np.allclose([A.sum(axis=1) for A in mapas], 1.0)))
print("fila 0 de la cabeza 1:", np.round(mapas[0][0], 3))
print("fila 0 de la cabeza 2:", np.round(mapas[1][0], 3))

d = 512 # la anchura del articulo
print("\n h d_k proyecciones W^O total")
for hh in (1, 2, 4, 8, 16):
dk = d // hh
proy, sal = hh * d * 3 * dk, hh * dk * d
print(f"{hh:3d} {dk:4d} {proy:12d} {sal:8d} {proy + sal:9d}")
numpy

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

Las dos filas impresas son la misma posición vista por dos cabezas, y no se parecen en nada: la primera pone 0.5910.591 sobre la posición 11 y la segunda 0.5310.531 sobre la 55, con la que la primera gastaba 0.0050.005. Aquí las cabezas se han separado sólo porque rng les dio pesos distintos, que es todo lo que las separa también en un modelo entrenado. Y la tabla es la cancelación de arriba comprobada: las tres columnas de la derecha no se mueven de 786432786\,432, 262144262\,144 y 10485761\,048\,576 mientras dkd_k baja de 512512 a 3232.

La segunda celda mide la concesión. Si las hh cabezas empiezan con los mismos pesos, no hay hh cabezas.

import numpy as np

T, d_model, h = 6, 8, 4
d_k = d_model // h
rng = np.random.default_rng(5)
X = rng.normal(size=(T, d_model))


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


def mapas(Wq, Wk):
salida = []
for i in range(h):
col = slice(i * d_k, (i + 1) * d_k)
Q, K = X @ Wq[:, col], X @ Wk[:, col]
salida.append(softmax_filas(Q @ K.T / np.sqrt(d_k)))
return np.stack(salida) # (h, T, T)


bloque = rng.normal(size=(d_model, d_k)) * 0.6
iguales = np.tile(bloque, (1, h)) # las h cabezas, con los mismos pesos
A = mapas(iguales, iguales)
print("h mapas identicos:", bool(np.allclose(A, A[0])))
print("mayor diferencia con la cabeza 1:", np.abs(A - A[0]).max(axis=(1, 2)))

distintas = iguales.copy()
distintas[:, :d_k] += 0.4 * rng.normal(size=(d_model, d_k)) # solo la cabeza 1
B = mapas(distintas, distintas)
print("tras mover solo la cabeza 1:")
print(" lo que dista cada una de la cabeza 1:", np.round(np.abs(B - B[0]).max(axis=(1, 2)), 3))
print(" y lo que dista cada una de la cabeza 2:", np.round(np.abs(B - B[1]).max(axis=(1, 2)), 3))
print("fila 0, cabeza 1:", np.round(B[0][0], 3))
print("fila 0, cabeza 2:", np.round(B[1][0], 3))
numpy

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

Con los cuatro bloques de columnas iguales las cuatro distancias salen exactamente 00: los cuatro mapas son el mismo, y WO\mathbf{W}^O recibe cuatro copias de un mismo vector por posición. Mover un solo bloque separa a esa cabeza de las otras tres —0.4580.458 en las tres casillas de la primera línea— y deja a esas tres idénticas entre ellas, que es lo que dicen los tres ceros de la segunda. La asimetría no aparece sola: entra por la inicialización, y en un entrenamiento se mantiene porque las cabezas y los bloques de filas de WO\mathbf{W}^O empiezan todos distintos. Con todo simétrico, el gradiente no tendría por dónde romperlo.

Comprueba tu intuición

Cinco preguntas: qué separa a hh cabezas de una sola, cuánto pesa la capa, de dónde sale dkd_k, qué calcula la capa con las cabezas iguales, y si promediar era una opción.

Una capa tiene una cabeza con dk=4d_k = 4. Otra tiene cuatro cabezas de dk=1d_k = 1, y sus columnas son exactamente las cuatro de la primera. ¿En qué se diferencian?

Una capa multi-head con dmodel=512d_{\text{model}} = 512 y h=8h = 8, repartiendo la anchura como manda el artículo (dk=dv=dmodel/hd_k = d_v = d_{\text{model}}/h). ¿Cuántos parámetros tiene en total, contando WO\mathbf{W}^O?

parámetros

A margin of ±0 is accepted.

¿Por qué dk=dmodel/hd_k = d_{\text{model}}/h y no cualquier otra anchura?

Inicializas las hh cabezas con exactamente las mismas WiQ\mathbf{W}^Q_i, WiK\mathbf{W}^K_i y WiV\mathbf{W}^V_i. Marca lo que es cierto de la capa en ese momento.

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

Promediar las hh cabezas es algo que la concatenación seguida de WO\mathbf{W}^O no puede hacer.

Escribe multi_head(X, Wq, Wk, Wv, Wo, h). Recibe la frase ya vectorizada, X de forma (T,dmodel)(T, d_{\text{model}}); las tres proyecciones enteras, Wq, Wk y Wv, las tres de forma (dmodel,dmodel)(d_{\text{model}}, d_{\text{model}}); la proyección de salida Wo, de la misma forma; y el número de cabezas h, que divide a dmodeld_{\text{model}}.

La cabeza ii se queda con el bloque de columnas ii-ésimo de cada proyección —de i * d_k a (i + 1) * d_k, con d_k = d_v = d_model // h— y hace con él la atención de siempre, dividiendo entre dk\sqrt{d_k}.

Devuelve la pareja (salida, A), en ese orden: la salida de forma (T,dmodel)(T, d_{\text{model}}), con las hh cabezas concatenadas por columnas y proyectadas por Wo; y los hh mapas apilados, A de forma (h,T,T)(h, T, T), cada fila de cada uno sumando 11.

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 hh cabezas, una posición tiene hh respuestas donde tenía una, y WO\mathbf{W}^O decide qué hacer con ellas. Pero mira de dónde salen las hh: de la misma X\mathbf{X}, con la misma fórmula, y en eije_{ij} siguen sin aparecer ii ni jj por ningún sitio. Multiplicar las cabezas multiplica qué se compara, y no añade ni una coordenada de dónde. Ninguna de las hh puede aprender «el sustantivo que tengo a la izquierda», por muchas que pongas, porque en lo que la capa lee no hay izquierda.

La pieza que falta no está en la capa: está en la entrada, y consiste en que dos posiciones distintas le lleguen a X\mathbf{X} como dos filas distintas aunque el token sea el mismo. Eso es la lección siguiente, sobre la codificación posicional, y lo que la hace interesante no es que haga falta —eso ya se sabía desde la primera lección de este bloque— sino la forma concreta que elige el artículo: unos senos y unos cosenos con los que la distancia entre dos posiciones se puede escribir como una operación lineal sobre lo que se ha añadido. Esa propiedad se enseña, no se afirma, y ocupa la lección entera.

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.2 trae la concatenación con W^O y el reparto d_k = d_model/h, y afirma en una línea que h cabezas cuestan como una. Esa cancelación la desarrollas tú, y el marco «dónde se normaliza» no está en el artículo.