Auto-atención: la secuencia se mira a sí misma

Auto-atención: la secuencia se mira a sí misma

29 min read

La arquitectura terminó la lección anterior sin una sola pieza recurrente y con una única operación en el centro, la atención, a la que hay que pasarle tres cosas: consultas (queries), claves (keys) y valores (values). Las primeras se han quedado sin quien las haga. Era el decoder, un paso detrás de otro, y el decoder era recurrente. Lo único que sigue teniendo una fila por posición es la frase misma, ya convertida en vectores, así que para los tres papeles hay un candidato y no hay dos.

Usarlo tal cual, eso sí, no funciona, y falla por algo que se dice en una línea: a lo que más se parece un vector es a sí mismo. Cada posición se daría a sí misma el peso más alto, en todas las filas a la vez, y la capa se limitaría a devolverle a cada una lo que ya traía —sin un solo parámetro con el que hacer otra cosa—. Y hace falta que haga otra: en las llaves del coche están ahí, para conjugar están hay que haber mirado llaves, tres posiciones más atrás, y no coche, que está pegado y en singular. Esta lección monta la capa que sabe hacer eso, y lo que cuesta son tres matrices.

La rejilla de abajo es una frase contra sí misma: la fila es quien pregunta, la columna es quien contesta, y cada casilla es cuánto de esa columna entra en lo que lee esa fila. Escribe la frase que quieras y, sobre todo, apaga y enciende las proyecciones. Apagadas, la casilla más alta de cada fila es la que esa posición se da a sí misma, en las seis filas a la vez. Encendidas, la fila de están se va a buscar llaves por encima de todo lo demás.

Las mismas posiciones en los dos ejes, que es lo que no pasaba en el bloque anterior. Sin proyectar, cada fila se gana a sí misma; con las tres proyecciones, «están» se salta «coche» y se va a «llaves».

Tres lecturas de la misma lista

Apila la frase por filas, como el bloque anterior apilaba los estados del encoder:

XRT×dmodel,\mathbf{X} \in \mathbb{R}^{T \times d_{\text{model}}},

donde la fila tt es xt\mathbf{x}_t^{\top}, el vector del token que ocupa la posición tt —transpuesto, porque en este curso los vectores son columnas—. Eso es todo lo que entra en la capa: una sola matriz, ninguna segunda red.

Las tres listas salen de ahí, cada una por su propia matriz:

Q=XWQ,K=XWK,V=XWV,\mathbf{Q} = \mathbf{X}\mathbf{W}^Q, \qquad \mathbf{K} = \mathbf{X}\mathbf{W}^K, \qquad \mathbf{V} = \mathbf{X}\mathbf{W}^V,

con WQ,WKRdmodel×dk\mathbf{W}^Q, \mathbf{W}^K \in \mathbb{R}^{d_{\text{model}} \times d_k} y WVRdmodel×dv\mathbf{W}^V \in \mathbb{R}^{d_{\text{model}} \times d_v}, de modo que Q,KRT×dk\mathbf{Q}, \mathbf{K} \in \mathbb{R}^{T \times d_k} y VRT×dv\mathbf{V} \in \mathbb{R}^{T \times d_v}. Los superíndices de estas tres matrices son etiquetas del papel que hace cada una, no índices de capa: es el único sitio del curso donde un superíndice no cuenta capas, y viene de la notación del artículo. Aquí hay una sola de estas tres parejas de proyecciones; una capa de verdad lleva varias en paralelo, y eso es una lección más adelante, sobre la atención multi-head.

Ahí está el «auto» del nombre, y es lo único que separa esta lección de la lección sobre la consulta, la clave y el valor. Allí la consulta era el estado de una red y las claves y los valores eran los estados de otra: tres listas de dos sitios distintos. Aquí las tres salen de la misma X\mathbf{X}, y lo que cambia no es la operación —es idéntica— sino quién le pasa los argumentos.

Attention(Q,K,V)=softmax ⁣(QKdk)VRT×dv,\text{Attention}\left(\mathbf{Q}, \mathbf{K}, \mathbf{V}\right) = \text{softmax}\!\left(\frac{\mathbf{Q}\mathbf{K}^{\top}}{\sqrt{d_k}}\right)\mathbf{V} \in \mathbb{R}^{T \times d_v},

con el softmax por filas. El divisor dk\sqrt{d_k} es lo único de esa línea que todavía no ha justificado nadie; queda anotado y vuelvo a él al final.

Mira la forma del mapa de pesos, que es donde se nota que la entrada y la salida son la misma frase. QK\mathbf{Q}\mathbf{K}^{\top} multiplica una matriz T×dkT \times d_k por una dk×Td_k \times T, así que ART×T\mathbf{A} \in \mathbb{R}^{T \times T} es cuadrada: una fila y una columna por posición. En el bloque anterior era Ty×TxT_y \times T_x y sus dos ejes eran dos frases distintas; aquí los dos ejes son la misma, y por eso la casilla αii\alpha_{ii} existe y significa algo —cuánto de sí misma se queda una posición—. Cada fila sigue sumando 11, porque el softmax sigue normalizando una fila cada vez, y ninguna posición está excluida de ninguna fila: todas se miran a todas.

Y cuenta los pesos, que es la otra propiedad de forma que decide el resto del bloque. Las tres matrices juntan

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

y TT no aparece en esa cuenta. Multiplican filas de X\mathbf{X}, y les da igual cuántas filas haya. Ponlo al lado de la lección sobre por qué falla el MLP (multilayer perceptron) en secuencias: allí cada posición tenía su propio bloque de columnas Wt(1)\mathbf{W}^{(1)}_t, así que alargar la entrada obligaba a añadir pesos y cada posición aprendía por su cuenta. Estas tres matrices son las mismas para todas las posiciones y para todas las longitudes.

Sin proyecciones, la capa se copia a sí misma

Queda por defender lo que parece un rodeo: si las tres listas son lecturas de X\mathbf{X}, ¿por qué no leerla tal cual y ahorrarse las matrices? Pon Q=K=V=X\mathbf{Q} = \mathbf{K} = \mathbf{V} = \mathbf{X} y desarrolla. La puntuación de la pareja (i,j)(i,j) pasa a ser el producto escalar de los dos vectores de token,

eij=xixjdmodel,e_{ij} = \frac{\mathbf{x}_i^{\top}\mathbf{x}_j}{\sqrt{d_{\text{model}}}},

y de ahí salen dos cosas, ninguna buena. La primera: xixj=xjxi\mathbf{x}_i^{\top}\mathbf{x}_j = \mathbf{x}_j^{\top}\mathbf{x}_i, de modo que la rejilla de puntuaciones es simétrica. Los pesos no lo son —cada fila se divide entre la suma de la suya, y esas sumas no tienen por qué coincidir—, pero el orden dentro de cada fila lo fija un número que no distingue quién pregunta de quién contesta. Un idioma sí los distingue: están necesita a llaves mucho más de lo que llaves necesita a están.

La segunda es peor. Supongamos que todos los vectores de token miden lo mismo, xt=r\lVert \mathbf{x}_t \rVert = r —lo supongo yo aquí para que la cuenta salga exacta; la desigualdad de Cauchy-Schwarz hace el resto—:

xixjxixj=r2=xixi,\mathbf{x}_i^{\top}\mathbf{x}_j \le \lVert \mathbf{x}_i \rVert \lVert \mathbf{x}_j \rVert = r^{2} = \mathbf{x}_i^{\top}\mathbf{x}_i,

con igualdad sólo cuando xj\mathbf{x}_j apunta en la dirección de xi\mathbf{x}_i. La casilla de la diagonal es la mayor de su fila, siempre, en todas las filas a la vez, así que el softmax le da el peso más alto a la propia posición y la mezcla se inclina hacia xi\mathbf{x}_i: la capa devuelve una versión desteñida de lo que le entró. Cuánto se inclina depende de lo que midan los vectores —cuanto mayor es rr, más se cierra la fila sobre su diagonal—, y con longitudes distintas entre sí la conclusión no cambia de forma: la fila se la lleva entonces la posición cuyo vector sea más largo, que sigue siendo una propiedad de los vectores y no de lo que la frase necesita.

Lo que remata el argumento no es ninguna de las dos, sino lo que falta en las dos. Esa capa no tiene un solo parámetro. No hay nada que ajustar, así que no hay forma de que aprenda que un verbo en plural tiene que ir a buscar un sustantivo en plural. Las tres matrices son lo que separa los tres papeles —comparar no es ser comparado, y ninguno de los dos es aportar contenido— y, de paso, lo único que la capa puede aprender.

La secuencia en los tres papeles, en NumPy

La capa entera son cinco líneas. Míralas con las formas delante, que es donde se ve que la frase entra una vez y sale por tres sitios.

import numpy as np

T, d_model, d_k, d_v = 6, 8, 4, 5
rng = np.random.default_rng(5)
X = rng.normal(size=(T, d_model)) # fila t: el vector del token de la posicion t
Wq = rng.normal(size=(d_model, d_k)) * 0.6
Wk = rng.normal(size=(d_model, d_k)) * 0.6
Wv = rng.normal(size=(d_model, d_v)) * 0.6

Q, K, V = X @ Wq, X @ Wk, X @ Wv # tres lecturas de la MISMA X

E = Q @ K.T / np.sqrt(d_k) # (T, T): una puntuacion por pareja de posiciones
Z = np.exp(E - E.max(axis=1, keepdims=True))
A = Z / Z.sum(axis=1, keepdims=True) # softmax por filas
Y = A @ V

print("X", X.shape, " Q", Q.shape, " K", K.shape, " V", V.shape)
print("A", A.shape, " cuadrada:", A.shape[0] == A.shape[1])
print("cada fila suma 1:", bool(np.allclose(A.sum(axis=1), 1.0)))
print("A simetrica:", bool(np.allclose(A, A.T)))
print("salida", Y.shape, "->", d_v, "coordenadas por posicion")
print("parametros:", d_model * (2 * d_k + d_v), " y T no aparece en la cuenta")
print("fila 0 de A:", np.round(A[0], 3))
numpy

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

Cambia T a 99 y ejecuta otra vez: el mapa pasa a ser 9×99 \times 9, la salida gana tres filas y la línea de los parámetros no se mueve. Ésa es la propiedad de forma de esta capa en una frase.

Ahora la comprobación que sostiene la sección anterior. La celda de abajo quita las proyecciones, mide lo que pasa y las vuelve a poner.

import numpy as np

rng = np.random.default_rng(5)
T, d_model = 6, 8
X = rng.normal(size=(T, d_model))
X = X / np.linalg.norm(X, axis=1, keepdims=True) # todos de longitud 1


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


def atencion(Q, K, V):
A = softmax_filas(Q @ K.T / np.sqrt(Q.shape[1]))
return A @ V, A


# 1. Sin proyectar: las tres listas son X.
_, A0 = atencion(X, X, X)
print("puntuaciones simetricas:", bool(np.allclose(X @ X.T, (X @ X.T).T)))
print("la diagonal gana su fila:", bool(np.all(A0.argmax(axis=1) == np.arange(T))))
print("peso sobre si misma: ", np.round(np.diag(A0), 3))
print("mayor de los demas: ", np.round((A0 - np.eye(T)).max(axis=1), 3)) # sin la diagonal
_, A0_largo = atencion(3 * X, 3 * X, 3 * X) # los mismos vectores, mas largos
print("y con vectores 3 veces mas largos:", np.round(np.diag(A0_largo), 3))

# 2. Con proyecciones: dos lecturas distintas de la misma X.
d_k = 4
Wq, Wk = rng.normal(size=(d_model, d_k)), rng.normal(size=(d_model, d_k))
_, A1 = atencion(X @ Wq, X @ Wk, X)
print("ahora la diagonal gana su fila:", bool(np.all(A1.argmax(axis=1) == np.arange(T))))
print("peso sobre si misma: ", np.round(np.diag(A1), 3))
print("fila 0 entera: ", np.round(A1[0], 3))

# 3. El divisor de la formula, medido y todavia sin explicar.
for d in (8, 64, 512):
Qd, Kd = rng.normal(size=(T, d)), rng.normal(size=(T, d))
sin_dividir = softmax_filas(Qd @ Kd.T).max(axis=1).mean()
dividiendo = softmax_filas(Qd @ Kd.T / np.sqrt(d)).max(axis=1).mean()
print(f"d_k = {d:3d} peso mayor sin dividir: {sin_dividir:.3f} dividiendo: {dividiendo:.3f}")
numpy

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

Los dos primeros bloques son la sección anterior con números. Sin proyectar, las seis filas ponen su máximo en su propia diagonal, y la línea de los vectores tres veces más largos dice de qué depende que ese máximo sea una preferencia o un monopolio: con longitud 11 la fila reparte casi por igual, y estirando los mismos vectores el peso propio se dispara. Con dos lecturas distintas de la misma matriz eso se acaba, y las posiciones empiezan a mirarse entre ellas. El tercer bloque es otra cosa y conviene leerlo despacio, porque es de lo que trata la lección siguiente.

Comprueba tu intuición

Cinco preguntas: de dónde salen las tres listas, qué tamaño tiene el mapa, qué calcula la capa sin sus matrices, qué es simétrico y qué no, y qué se mueve cuando la frase se alarga.

En una capa de auto-atención, ¿de dónde salen Q\mathbf{Q}, K\mathbf{K} y V\mathbf{V}?

Una frase de T=12T = 12 tokens entra en una capa de auto-atención con dmodel=512d_{\text{model}} = 512 y dk=dv=64d_k = d_v = 64. ¿Cuántos números tiene el mapa de pesos A\mathbf{A}?

números

A margin of ±0 is accepted.

Quitas las tres proyecciones —Q=K=V=X\mathbf{Q} = \mathbf{K} = \mathbf{V} = \mathbf{X}— con vectores de token de la misma longitud. ¿Qué calcula la capa?

Con Q=K=X\mathbf{Q} = \mathbf{K} = \mathbf{X}, ¿qué es simétrico y qué no?

La misma capa recibe primero una frase de T=20T = 20 tokens y luego una de T=200T = 200. Marca lo que es cierto.

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

Escribe auto_atencion(X, Wq, Wk, Wv). Recibe la frase ya vectorizada, X de forma (T,dmodel)(T, d_{\text{model}}) —una fila por posición—, y las tres proyecciones: Wq y Wk de forma (dmodel,dk)(d_{\text{model}}, d_k) y Wv de forma (dmodel,dv)(d_{\text{model}}, d_v). Devuelve la pareja (salida, A), en ese orden: la salida de forma (T,dv)(T, d_v) y el mapa de pesos A\mathbf{A} de forma (T,T)(T, T), ya normalizado por filas.

La puntuación es el producto escalar de cada consulta con cada clave dividido por dk\sqrt{d_k}, con dkd_k leído de la forma de Wq. Escríbelo con productos de matrices, sin recorrer las parejas una a una, 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.


Queda el divisor, y ya no se puede dejar pasar. El tercer bloque de la última celda lo mide sin tocar nada más: con consultas y claves de dk=8d_k = 8 el mayor peso de una fila se queda en 0.5580.558, y con los mismos vectores aleatorios a dk=512d_k = 512 sube a 0.9830.983. Una fila así ha dejado de ser una mezcla —es quedarse con una posición y tirar las otras cinco—, y eso tiene nombre desde la lección sobre las funciones de activación: está saturada. Dividiendo entre dk\sqrt{d_k} el máximo también crece, pero despacio, de 0.3300.330 a 0.4860.486, y la fila sigue repartiendo.

Esa raíz no es una constante que alguien fuera probando hasta que las cosas salieron bien: se calcula, y lo que hay que calcular es cuánto crece un producto escalar a medida que se le añaden coordenadas. La cuenta ocupa entera la lección siguiente, sobre el producto interno escalado, y merece hacerse despacio, porque es donde vuelven a encontrarse dos cosas que este curso dejó en bloques distintos: la saturación que apaga un reparto y el gradiente que se desvanece cuando lo atraviesa. Ninguna de las dos necesitaba una sigmoide ni una recurrencia para aparecer.

Further reading2 sources · 1 paper, 1 article

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

    La operación está en su ecuación (1); que las tres listas salgan de la misma secuencia, en su §3.2.3, en una frase. Lo que el artículo no hace es defender las proyecciones: por qué sin ellas la capa se copia a sí misma es tuyo.

  • The Illustrated Transformer
    articleJay Alammar, 2018jalammar.github.ioEN

    Recorre con dibujos el cálculo de Q, K y V y la rejilla de pesos, primero vector a vector y luego en forma matricial. Se queda en la mecánica: la simetría y la diagonal con las que aquí justificas las tres matrices no las toca.