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.
Tres lecturas de la misma lista
Apila la frase por filas, como el bloque anterior apilaba los estados del encoder:
donde la fila es , el vector del token que ocupa la posición —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:
con y , de modo que y . 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 , y lo que cambia no es la operación —es idéntica— sino quién le pasa los argumentos.
con el softmax por filas. El divisor 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. multiplica una matriz por una , así que es cuadrada: una fila y una columna por posición. En el bloque anterior era y sus dos ejes eran dos frases distintas; aquí los dos ejes son la misma, y por eso la casilla existe y significa algo —cuánto de sí misma se queda una posición—. Cada fila sigue sumando , 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
y no aparece en esa cuenta. Multiplican filas de , 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 , 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 , ¿por qué no leerla tal cual y ahorrarse las matrices? Pon y desarrolla. La puntuación de la pareja pasa a ser el producto escalar de los dos vectores de token,
y de ahí salen dos cosas, ninguna buena. La primera: , 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, —lo supongo yo aquí para que la cuenta salga exacta; la desigualdad de Cauchy-Schwarz hace el resto—:
con igualdad sólo cuando apunta en la dirección de . 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 : 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 , 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.
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))
La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador.
Cambia T a y ejecuta otra vez: el mapa pasa a ser , 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.
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}")
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 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 , y ?
Una frase de tokens entra en una capa de auto-atención con y . ¿Cuántos números tiene el mapa de pesos ?
A margin of ±0 is accepted.
Quitas las tres proyecciones —— con vectores de token de la misma longitud. ¿Qué calcula la capa?
Con , ¿qué es simétrico y qué no?
La misma capa recibe primero una frase de tokens y luego una de . 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
—una fila por posición—, y las tres proyecciones: Wq y Wk de
forma y Wv de forma . Devuelve la
pareja (salida, A), en ese orden: la salida de forma y el mapa de pesos
de forma , ya normalizado por filas.
La puntuación es el producto escalar de cada consulta con cada clave dividido por
, con 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 el mayor peso de una fila se queda en , y con los mismos vectores aleatorios a sube a . 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 el máximo también crece, pero despacio, de a , 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
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
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.