Atención como consulta, clave y valor

Atención como consulta, clave y valor

28 min read

El bloque ha montado un mecanismo completo y dos maneras distintas de rellenar su única casilla libre. Lo que no ha mirado en ningún momento es el vocabulario con el que está escrito todo eso: cada ecuación de las tres lecciones anteriores habla de un encoder que lee, de un decoder que escribe, de estados con barra y de estados sin ella. Ese vocabulario viene de una máquina concreta —dos redes recurrentes encadenadas para traducir—, y la pregunta de esta lección es cuánta de esa máquina necesita de verdad la atención.

Hazlo a lo bruto: tacha de las ecuaciones todo lo que hable de recurrencia y mira qué queda en pie. Queda un vector que pregunta, una lista de vectores contra los que se le compara y una lista de vectores que se mezclan; queda un softmax entre las dos listas y una suma ponderada al final. Ni una de esas piezas menciona una recurrencia, ni un idioma, ni una traducción. Son tres listas y una operación, cada una con su nombre fijado desde hace años —y con la inicial que el bloque siguiente pone en todas sus ecuaciones, empezando por la primera línea del artículo que las escribe—.

Hay una operación que ya has escrito y que hace casi exactamente esto. En los desafíos de la lección sobre la LSTM (long short-term memory) y la lección sobre la GRU (gated recurrent unit), los pesos llegaban en un diccionario P y se sacaban con P["z"]: das una consulta (query), Python la compara con las claves (keys) guardadas y te devuelve el valor (value) de la que coincide. Falla una letra y no hay respuesta.

La atención es esa búsqueda con la exigencia de coincidir quitada. La consulta no tiene que ser igual a ninguna clave: se compara con todas, cada comparación se convierte en un peso, y lo que vuelve no es un valor sino la mezcla de todos ellos, cada uno en la proporción que le tocó. Un diccionario devuelve una entrada; la atención devuelve una combinación de todas. En lo que llevas de bloque los tres papeles están repartidos así: pregunta el estado del decoder, y cada posición de la fuente pone las dos cosas restantes: con qué se la compara y qué vierte en la mezcla.

Arriba a la izquierda, una caja verde etiquetada q sub i, de la que sale una flecha hacia la derecha que entra en una franja ancha etiquetada a de q sub i y k sub j, y softmax sobre las posiciones. Debajo, en el centro, tres cajas en fila etiquetadas h con barra sub uno, sub dos y sub tres; a su izquierda, la nota k sub j igual a v sub j, un solo vector. De la parte de arriba de cada caja sale una flecha que sube hasta la franja, en un carril etiquetado claves. De la parte de abajo de cada caja sale una flecha que baja hasta una caja verde etiquetada c sub i, en un carril etiquetado valores; las tres flechas tienen grosores distintos y están etiquetadas alfa sub i uno, alfa sub i dos y alfa sub i tres, la del medio la más gruesa. Bajo la caja verde, la nota d sub v números.
Cada posición de la fuente pone dos cosas con un mismo vector: la clave con la que se la compara y el valor que se mezcla. Ninguna ecuación del bloque pide que sean el mismo.

Tres papeles, tres argumentos

Fijemos los tres papeles como tres objetos, sin decir de dónde salen. Una consulta qiRdk\mathbf{q}_i \in \mathbb{R}^{d_k}, y TxT_x parejas de clave y valor, kjRdk\mathbf{k}_j \in \mathbb{R}^{d_k} y vjRdv\mathbf{v}_j \in \mathbb{R}^{d_v}, una pareja por posición. La consulta y las claves miden lo mismo porque van a compararse entre sí; los valores no tienen por qué. Con eso, la llamada de atención son tres pasos:

eij=a(qi,kj),αij=exp(eij)m=1Txexp(eim),ci=j=1TxαijvjRdv.e_{ij} = a\left(\mathbf{q}_i, \mathbf{k}_j\right), \qquad \alpha_{ij} = \frac{\exp(e_{ij})}{\sum_{m=1}^{T_x} \exp(e_{im})}, \qquad \mathbf{c}_i = \sum_{j=1}^{T_x} \alpha_{ij}\,\mathbf{v}_j \in \mathbb{R}^{d_v}.

Es el mecanismo de la lección sobre la idea de atención, con otros nombres: se puntúa cada pareja, se normaliza la fila con un softmax sobre las posiciones y se mezclan con esos pesos. El índice del denominador cambia de letra —mm donde antes iba kk— porque la kk está ocupada de aquí en adelante por las claves.

Mira las formas, que es donde la reescritura empieza a pagar. La salida tiene dvd_v coordenadas: el tamaño de un valor, no el de la consulta y no el número de parejas. dkd_k vive entero dentro de la comparación y muere ahí, en el número eije_{ij}; TxT_x vive en la fila de pesos y muere en la suma. De modo que la lista de parejas puede crecer todo lo que quiera y el hueco donde cae el resultado no se entera —que es la propiedad de forma de esa lección, dicha ahora sin nombrar ninguna recurrencia—.

Y las tres lecciones anteriores son un caso particular de esto, sin cambiar una ecuación: qi=si1\mathbf{q}_i = \mathbf{s}_{i-1}, y kj=vj=hˉj\mathbf{k}_j = \mathbf{v}_j = \bar{\mathbf{h}}_j, con dk=dv=dhd_k = d_v = d_h. La función aa es la que se eligiera —la aditiva de Bahdanau o la multiplicativa de Luong—, el softmax es el mismo y la suma ponderada es la misma. Lo único que ha pasado es que cada pieza tiene ahora un nombre que no menciona quién la produjo.

La clave y el valor son el mismo vector, y ninguna ecuación lo pide

La igualdad del medio, kj=vj\mathbf{k}_j = \mathbf{v}_j, es lo interesante, porque no la exige nada de lo que hay escrito arriba. Repasa dónde aparece cada uno: kj\mathbf{k}_j está dentro de aa y en ningún otro sitio, y vj\mathbf{v}_j está dentro de la suma y en ningún otro sitio. Son dos trabajos separados —uno decide cuánto pesa la posición jj, el otro es lo que llega al decoder cuando pesa— y en el modelo del bloque los hace un mismo vector. Eso obliga a las dhd_h coordenadas de hˉj\bar{\mathbf{h}}_j a servir para las dos cosas a la vez: a ser buenas para que las comparen y buenas para que las lean.

La cuenta del gradiente lo enseña mejor que la fórmula. Desatados, cada papel aparece en un solo sitio, así que la regla de la cadena de su lección recorre un solo camino para cada uno; lo que suma es sobre los TyT_y pasos, porque las dos listas se usan en todos:

vj=i=1Tyαijci,kj=i=1Tyeijkjeij,\nabla_{\mathbf{v}_j}\ell = \sum_{i=1}^{T_y} \alpha_{ij}\,\nabla_{\mathbf{c}_i}\ell, \qquad \nabla_{\mathbf{k}_j}\ell = \sum_{i=1}^{T_y} \frac{\partial \ell}{\partial e_{ij}}\,\nabla_{\mathbf{k}_j} e_{ij},

y el paso por el softmax es el de la lección sobre la atención aditiva, sin tocar una coma, sólo que ahora dice con precisión qué se compara con qué:

eij=αij(ci)(vjci).\frac{\partial \ell}{\partial e_{ij}} = \alpha_{ij}\left(\nabla_{\mathbf{c}_i}\ell\right)^{\top}\left(\mathbf{v}_j - \mathbf{c}_i\right).

Quien se compara con la mezcla es el valor de esa posición; quien recibe la corrección es su clave. Con un solo vector haciendo los dos papeles esa frase no se podía ni escribir.

Y al atarlos otra vez sale la fórmula de la lección anterior

Poner kj=vj=hˉj\mathbf{k}_j = \mathbf{v}_j = \bar{\mathbf{h}}_j no es un caso nuevo del gradiente: es un peso compartido, como el Whh\mathbf{W}_{hh} que recorre todas las posiciones en la recurrencia del bloque anterior o el Wa\mathbf{W}_a que puntúa todas las parejas. Un objeto que aparece en dos sitios recibe la suma de lo que le llega por cada uno, de modo que las dos rutas de arriba se suman:

hˉj=i=1Tyαijci+i=1Tyeijhˉjeij,\nabla_{\bar{\mathbf{h}}_j}\ell = \sum_{i=1}^{T_y} \alpha_{ij}\,\nabla_{\mathbf{c}_i}\ell + \sum_{i=1}^{T_y} \frac{\partial \ell}{\partial e_{ij}}\,\nabla_{\bar{\mathbf{h}}_j} e_{ij},

que es, término a término, la línea de las dos rutas de la lección sobre la atención aditiva —allí el segundo sumando estaba desarrollado para la puntuación aditiva, Whapij\mathbf{W}_{ha}^{\top}\nabla_{\mathbf{p}_{ij}}\ell, que es este mismo con hˉjeij\nabla_{\bar{\mathbf{h}}_j} e_{ij} escrito por dentro—. Lo que en aquella lección era una peculiaridad de la arquitectura —«a este estado le llega el gradiente por dos caminos»— resulta ser una consecuencia de haber atado dos papeles.

La rejilla entera con tres matrices

Apila los tres papeles por filas, como la lección anterior apilaba los estados: QRTy×dk\mathbf{Q} \in \mathbb{R}^{T_y \times d_k} con la consulta del paso ii en la fila ii, KRTx×dk\mathbf{K} \in \mathbb{R}^{T_x \times d_k} y VRTx×dv\mathbf{V} \in \mathbb{R}^{T_x \times d_v} con la clave y el valor de la posición jj en su fila jj —transpuestos, porque en este curso los vectores son columnas—. Con la puntuación en su forma más barata, el producto escalar, las tres operaciones de arriba caben en una línea:

Attention(Q,K,V)=softmax(QK)VRTy×dv,\text{Attention}\left(\mathbf{Q}, \mathbf{K}, \mathbf{V}\right) = \text{softmax}\left(\mathbf{Q}\mathbf{K}^{\top}\right)\mathbf{V} \in \mathbb{R}^{T_y \times d_v},

con el softmax por filas, de modo que cada fila de QK\mathbf{Q}\mathbf{K}^{\top} —una consulta contra las TxT_x claves— sale sumando 11. El nombre de la función se queda en inglés, como el título del artículo que la fija: es así como vas a encontrarla escrita en todas partes.

Y ahora la pregunta que decide el bloque siguiente: ¿dónde ha quedado Wa\mathbf{W}_a? La rejilla de la lección anterior era SWaHˉ\mathbf{S}\mathbf{W}_a\bar{\mathbf{H}}^{\top}, así que agrupa los dos primeros factores y compara con la línea de arriba. Sale Q=SWa\mathbf{Q} = \mathbf{S}\mathbf{W}_a —esto es, qi=Wasi1\mathbf{q}_i = \mathbf{W}_a^{\top}\mathbf{s}_{i-1}— y K=V=Hˉ\mathbf{K} = \mathbf{V} = \bar{\mathbf{H}}. La matriz no ha desaparecido: ha dejado de ser parte de la comparación para ser parte de lo que se compara. Y con ella fuera, la función aa se queda en el producto escalar pelado para siempre; todo lo que un modelo quiera aprender sobre cómo comparar tiene que estar ya metido en cómo se fabrican Q\mathbf{Q} y K\mathbf{K}, que es un asunto del que esta línea no dice nada y del que el bloque siguiente no habla de otra cosa.

La aditiva no se deja mudar de sitio, y es el mismo argumento de su lección: entre las dos proyecciones y la lectura hay un tanh\tanh, y una función no lineal de una suma no se reparte entre sus sumandos. Por eso esta forma de tres matrices —la que el bloque siguiente hereda— existe gracias a la elección que hizo la lección anterior, y no a pesar de ella.

Queda una deuda, y conviene decirla en voz alta porque la vas a ver escrita. La forma del bloque siguiente lleva un factor más:

Attention(Q,K,V)=softmax ⁣(QKdk)V.\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}.

Ese divisor responde a lo que la lección anterior midió sin arreglarlo: con pesos recién inicializados el mayor peso de una fila se iba a 0.8510.851 frente al 0.1670.167 del reparto uniforme, porque un producto escalar sin acotar crece con el tamaño de los vectores que multiplica. En esta escritura ese tamaño tiene nombre, dkd_k, y por eso el arreglo cabe en una división. De dónde sale la raíz —y por qué dk\sqrt{d_k} y no otra cosa— es la lección del bloque siguiente, sobre el producto interno escalado; lo que hace falta aquí es saber que el factor está ahí y a qué problema contesta.

Los tres argumentos, en NumPy

La función entera son dos líneas, y lo que la celda comprueba es que no ha cambiado nada al cambiarle los nombres. Primero recupera la rejilla de la lección anterior repartiendo los papeles —la consulta es el estado del decoder reescrito por Wa\mathbf{W}_a, la clave y el valor son el mismo estado del encoder— y después los desata: una lectura para comparar de dk=3d_k = 3 coordenadas y otra para mezclar de dv=7d_v = 7. Mira sobre todo las tres últimas líneas.

import numpy as np

T_y, T_x, d_h = 8, 6, 5
rng = np.random.default_rng(4)
S = rng.normal(size=(T_y, d_h)) * 0.6 # fila i: el estado s_{i-1} del decoder
H = rng.normal(size=(T_x, d_h)) * 0.6 # fila j: el estado del encoder en la posicion j
Wa = rng.normal(size=(d_h, d_h)) * 0.4 # la matriz de la puntuacion multiplicativa


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) # (T_q, T_k): un peso por pareja, filas que suman 1
return A @ V, A # (T_q, d_v): una mezcla de valores por consulta


# 1. La leccion anterior, con los papeles repartidos: la consulta es el estado del
# decoder reescrito por Wa, y la clave y el valor son el mismo estado del encoder.
Y, A = atencion(S @ Wa, H, H)
print("misma rejilla que la leccion anterior:", bool(np.allclose(A, softmax_filas(S @ Wa @ H.T))))
print("cada fila de A suma 1:", bool(np.allclose(A.sum(axis=1), 1.0)))
print("formas: Q", (S @ Wa).shape, " K", H.shape, " V", H.shape, " -> salida", Y.shape)
print()

# 2. Desatados: dos lecturas distintas del mismo estado, y ya ni miden lo mismo.
d_k, d_v = 3, 7
Wq, Wk, Wv = (rng.normal(size=(d_h, d_k)), rng.normal(size=(d_h, d_k)), rng.normal(size=(d_h, d_v)))
Q, K, V = S @ Wq, H @ Wk, H @ Wv
Y2, A2 = atencion(Q, K, V)
print("con d_k =", d_k, "y d_v =", d_v, " -> A", A2.shape, " salida", Y2.shape,
" filas que suman 1:", bool(np.allclose(A2.sum(axis=1), 1.0)))

# 3. Los dos papeles, uno contra otro: quien decide el reparto y quien llega al decoder.
otra_K, otra_V = rng.normal(size=(T_x, d_k)), rng.normal(size=(T_x, d_v))
print("cambiar V mueve el mapa: ", bool(not np.allclose(atencion(Q, K, otra_V)[1], A2)))
print("cambiar K mueve el mapa: ", bool(not np.allclose(atencion(Q, otra_K, V)[1], A2)))
print("cambiar V mueve la salida:", bool(not np.allclose(atencion(Q, K, otra_V)[0], Y2)))
numpy

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

Cambiar los valores enteros no mueve ni un peso del mapa, y cambiar las claves lo mueve todo. Ésa es la separación que las tres lecciones anteriores no podían enseñar, porque el vector era uno y tocarlo cambiaba las dos cosas a la vez. Fíjate también en las formas del segundo bloque: la consulta mide 33, la salida mide 77 y el mapa sigue siendo 8×68 \times 6, una fila por paso y una columna por posición. Cambia d_v a 22 y ejecuta otra vez: la salida encoge y el mapa no se inmuta.

Comprueba tu intuición

Cinco preguntas: quién hace cada papel en lo que ya sabes, qué tamaño tiene lo que sale, qué pasa al tocar sólo los valores, por dónde le llega el gradiente a cada papel y dónde ha quedado Wa\mathbf{W}_a.

En la atención de las tres lecciones anteriores, ¿qué vector hace cada papel?

Una llamada de atención con Tx=12T_x = 12 parejas, consultas y claves de dk=64d_k = 64 coordenadas y valores de dv=32d_v = 32. ¿Cuántos números devuelve una consulta?

números

A margin of ±0 is accepted.

Cambias los valores vj\mathbf{v}_j por otros y dejas las consultas y las claves como estaban. ¿Qué le pasa al mapa de pesos A\mathbf{A}?

Con la clave y el valor desatados —dos vectores distintos en cada posición—, marca lo que es cierto sobre cómo les llega el gradiente de la pérdida.

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

¿Dónde queda la matriz Wa\mathbf{W}_a de la lección anterior cuando la rejilla se escribe como QK\mathbf{Q}\mathbf{K}^{\top}?

Escribe atencion(Q, K, V), la función de esta lección. Recibe las consultas Q, de forma (Tq,dk)(T_q, d_k) —una fila por consulta—, las claves K, de forma (Tk,dk)(T_k, d_k), y los valores V, de forma (Tk,dv)(T_k, d_v): una clave y un valor por posición. Devuelve la pareja (salida, A), en ese orden —la salida, de forma (Tq,dv)(T_q, d_v), y el mapa de pesos A\mathbf{A}, de forma (Tq,Tk)(T_q, T_k), ya normalizado por filas—.

La puntuación es el producto escalar de cada consulta con cada clave, sin matriz en medio. Escríbela 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.


Con los nombres puestos, lo que falta se ve de otra manera. K\mathbf{K} y V\mathbf{V} existen enteras en cuanto el encoder termina de leer: sólo dependen de la fuente, y ninguna de sus filas espera a ninguna otra. Q\mathbf{Q} no. Sus filas se fabrican de una en una, porque cada consulta es un estado del decoder y cada estado del decoder pide el contexto del paso anterior. Dos de las tres matrices están completas de golpe y la tercera sigue yendo en fila india, que es la cuenta pendiente de la lección anterior dicha con el vocabulario de ésta.

Y ahí es donde el curso gira. La primera lección del bloque siguiente, sobre quitar la recurrencia se toma en serio la única pregunta que queda: si lo que estrangula el paralelismo es la recurrencia, y la atención ya sabe leer una secuencia entera sin recorrerla paso a paso, ¿qué queda si se quita? La respuesta ocupa el resto del curso, y llegas a ella con las tres palabras y las tres letras ya puestas —que es lo único que este bloque tenía que dejar hecho—.

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

    De aquí salen los tres nombres y softmax(QK^T/√d_k)V: su §3.2 abre con el marco consulta-clave-valor y su §3.2.1 lo escribe. Esta lección toma solo la notación; la arquitectura llega en el bloque siguiente, y el √d_k en el próximo.