La idea de atención: mirar hacia atrás

La idea de atención: mirar hacia atrás

24 min read

«Que el decoder mire la entrada entera» suena a solución y todavía no es una arquitectura. La lección anterior, sobre el cuello de botella del vector de contexto, midió lo que cuesta resumir y dejó a mano el material desaprovechado: los TxT_x estados que el encoder calcula y nadie vuelve a tocar. Falta todo lo demás. ¿Qué recibe exactamente el decoder en el paso ii, quién lo construye y con qué material?

Antes de la primera pregunta hay un obstáculo de forma, y es el que decide la respuesta. Los estados del encoder son tantos como tokens tenga la frase: seis para una, treinta para otra. El decoder, en cambio, tiene un tamaño decidido antes de ver ninguna —sus matrices de pesos tienen un número fijo de columnas—, de modo que ponerle los estados en fila y multiplicarlos por una matriz choca contra el mismo muro que la lección sobre por qué el perceptrón multicapa falla en secuencias: una entrada cuya longitud cambia no cabe en una matriz cuya forma no cambia. Leerlos todos tendrá que significar otra cosa. Tendrá que significar combinarlos en un solo vector del tamaño que el decoder ya esperaba, y entonces lo único que queda por decidir es cuánto de cada uno.

Cuánto de cada uno, y no cuál. La diferencia es lo que hace que la idea funcione: si el paso ii tuviera que señalar un estado, elegir sería una decisión discreta, y ya veremos que eso tiene un precio que no se puede pagar. Repartir, en cambio, se hace con números. El paso que escribe book vierte casi todo de libro y unas gotas del resto; el que escribe very vierte de muy. Cada paso se lleva su propia mezcla, y a ese reparto —una cantidad por posición de la entrada, distinta en cada paso de la salida— es a lo que llamamos atención.

El explorable trae la pareja de la lección sobre la arquitectura encoder-decoder, con una fila por paso de la salida. Recórrelas y mira dos sitios. i y read tiran del mismo token, leí, porque el español lleva el pronombre dentro del verbo y el inglés lo escribe aparte. Y book, en el paso 7, vuelve a libro —que está en la posición 4— después de que los pasos 5 y 6 hayan estado en muy y bueno: las líneas se cruzan, que es la reordenación que el resumen fijo tenía que aguantar de memoria. Después pulsa «resumen fijo».

Los pesos los he puesto yo a mano para el curso; lo que importa es su forma. Con «resumen fijo» sólo se enciende la última columna y las barras de abajo dejan de moverse: ocho pasos, un solo vector.

Un contexto distinto en cada paso

Fijemos lo que recibe el paso ii. Son TxT_x números, uno por posición de la entrada, y los llamamos αi1,,αiTx\alpha_{i1}, \dots, \alpha_{iT_x}. Con ellos el contexto de ese paso es

ci=j=1TxαijhˉjRdh,αij0,j=1Txαij=1.\mathbf{c}_i = \sum_{j=1}^{T_x} \alpha_{ij}\,\bar{\mathbf{h}}_j \in \mathbb{R}^{d_h}, \qquad \alpha_{ij} \geq 0, \qquad \sum_{j=1}^{T_x} \alpha_{ij} = 1.

Mira las dos formas, porque la gracia está en que no coinciden. La lista de pesos mide TxT_x: crece con la entrada, un peso más por cada token nuevo, y ahí es donde entra la información que antes no cabía. El resultado mide dhd_h pase lo que pase, que es exactamente lo que le permite ocupar el sitio del vector de contexto de siempre sin tocar el resto del decoder. La entrada creció; el hueco no.

Las dos condiciones sobre los pesos convierten esa suma en una mezcla, y merece la pena ver qué prohíben. Como ningún peso es negativo y todos suman uno, cada coordenada de ci\mathbf{c}_i queda entre la menor y la mayor de esa misma coordenada entre los hˉj\bar{\mathbf{h}}_j: la mezcla reparte lo que hay, no inventa valores nuevos ni amplifica ninguno. Los dos extremos de lo que puede hacer se escriben solos. Si un peso vale 11 y el resto 00, entonces ci=hˉj\mathbf{c}_i = \bar{\mathbf{h}}_{j} y la mezcla se ha convertido en una elección dura de un estado. Si todos valen 1/Tx1/T_x, entonces ci\mathbf{c}_i es la media de la frase entera, que mira todo y no distingue nada. Entre esos dos casos vive el mecanismo.

El resumen fijo era una alineación congelada

Hay un tercer juego de pesos que ya conoces, aunque nunca se escribió como tal. Toma αij=1\alpha_{ij} = 1 cuando j=Txj = T_x y αij=0\alpha_{ij} = 0 en el resto, para todos los pasos por igual:

ci=j=1Txαijhˉj=hˉTx=c.\mathbf{c}_i = \sum_{j=1}^{T_x} \alpha_{ij}\,\bar{\mathbf{h}}_j = \bar{\mathbf{h}}_{T_x} = \mathbf{c}.

Eso es el encoder-decoder de las dos lecciones anteriores, entero, sin quitar ni añadir nada. Y es una alineación perfectamente legal: los pesos no son negativos y suman uno. De modo que la atención no sustituye a aquella arquitectura ni compite con ella —la contiene—: lo que hace es descongelar unos pesos que estaban puestos a mano, en el último token, y que no dependían de ii. El cuello de botella no venía de que dhd_h fuera pequeño. Venía de que esa fila no se podía mover.

El contexto entra como una entrada más del decoder

Arriba, tres celdas del encoder etiquetadas h con barra sub uno, sub dos y sub T equis, enlazadas de izquierda a derecha, cada una recibiendo por arriba un token de la entrada, x sub uno, x sub dos, x sub T equis. De la parte de abajo de las tres celdas salen tres hilos verdes de grosor distinto, etiquetados alfa sub i uno, alfa sub i dos y alfa sub i T equis; el del medio es el más grueso. Los tres convergen en una caja verde etiquetada c sub i, con la nota «d sub h números» debajo. De esa caja sale una flecha verde horizontal hacia la derecha que entra en una celda del decoder etiquetada s sub i. Esa celda recibe además, por arriba, el estado anterior s sub i menos uno, y por abajo el token anterior, y con circunflejo sub i menos uno; por la derecha saca y con circunflejo sub i.
Los tres hilos salen de todas las celdas del encoder, no sólo de la última, y su grosor es su peso. La arquitectura anterior es este mismo dibujo con un solo hilo, siempre el de la derecha.

El contexto llega a la recurrencia del decoder como un sumando más, con su propia matriz de pesos WchRdh×dh\mathbf{W}_{ch} \in \mathbb{R}^{d_h \times d_h}:

si=tanh(Wxhdecoyi1+Whhdecsi1+Wchci+bhdec),\mathbf{s}_i = \tanh\left(\mathbf{W}^{\text{dec}}_{xh}\mathbf{o}_{y_{i-1}} + \mathbf{W}^{\text{dec}}_{hh}\mathbf{s}_{i-1} + \mathbf{W}_{ch}\mathbf{c}_i + \mathbf{b}^{\text{dec}}_h\right),

y el estado de partida cambia de sitio. En la lección sobre la arquitectura encoder-decoder el decoder arrancaba en s0=c\mathbf{s}_0 = \mathbf{c}, porque había un c\mathbf{c} y sólo uno; ahora no lo hay, así que tomamos s0=0\mathbf{s}_0 = \mathbf{0}, como la recurrencia del bloque anterior. Es una elección, no una necesidad —lo habitual es arrancar de una proyección aprendida de hˉTx\bar{\mathbf{h}}_{T_x}—, y no cambia nada de lo que sigue.

Con eso, la línea incómoda de la lección sobre la arquitectura encoder-decoder queda reescrita:

P(yiy<i,c)P(yiy<i,ci).P\left(y_i \mid y_{<i},\, \mathbf{c}\right) \quad\longrightarrow\quad P\left(y_i \mid y_{<i},\, \mathbf{c}_i\right).

El cambio parece un subíndice y es un cambio de condicionante. Aquella igualdad afirmaba que todo lo que la entrada tuviera que decir sobre cualquier posición de la salida cabía en dhd_h números fijados de una vez; esta afirma algo mucho más flojo, que para cada posición hay una mezcla de los estados que sirve. Y como ci\mathbf{c}_i se construye a partir de los TxT_x estados, el paso ii está condicionando otra vez sobre la fuente completa. La suposición no ha desaparecido, se ha vuelto barata.

De dónde salen los pesos

Queda la pregunta que lo sostiene todo: quién decide los αij\alpha_{ij}. No pueden estar fijados de antemano —eso es justo lo que se acaba de descongelar— ni pueden depender sólo de jj, porque entonces todos los pasos mirarían al mismo sitio. Tienen que depender de la pareja: de dónde está el decoder y qué hay en la posición jj. Así que primero se puntúa cada pareja,

eij=a(si1,hˉj),e_{ij} = a\left(\mathbf{s}_{i-1},\, \bar{\mathbf{h}}_j\right),

donde aa es una función que devuelve un número —lo compatibles que son ese estado del decoder y ese estado del encoder— y si1\mathbf{s}_{i-1} es el estado anterior por una razón de orden, no de gusto: calcular si\mathbf{s}_i exige tener ya ci\mathbf{c}_i, que exige tener los pesos, que exigen las puntuaciones. Lo último que existe cuando hay que puntuar es si1\mathbf{s}_{i-1}.

Las puntuaciones son números sueltos y pueden ser negativos, así que hace falta convertirlas en una mezcla. Eso lo hace un softmax, y aquí conviene mirar sobre qué:

αij=exp(eij)k=1Txexp(eik).\alpha_{ij} = \frac{\exp(e_{ij})}{\sum_{k=1}^{T_x} \exp(e_{ik})}.

Sale positivo por la exponencial y suma uno por el denominador, que son las dos condiciones de la mezcla, y de propina reparte de forma suave: una puntuación algo mayor se lleva algo más de peso, no todo.

Dos consecuencias, y la primera es una factura. Hay una puntuación por cada pareja de paso y posición, así que traducir una frase cuesta TxTyT_x \cdot T_y evaluaciones de aa: el precio crece con el producto de las dos longitudes. El resumen fijo no pagaba ninguna, y eso es precisamente lo que lo ahogaba; conviene saber desde el principio que la mirada hacia atrás no sale gratis.

La segunda explica el softmax en lugar de algo más directo. Lo natural sería quedarse con el estado mejor puntuado, argmaxjeij\arg\max_j e_{ij}, y ahorrarse la suma entera. Pero eso devuelve una posición, y una posición es un entero: mueve los parámetros de aa una millonésima y el índice no se inmuta, hasta que un día salta. La derivada de ci\mathbf{c}_i respecto de esos parámetros vale cero donde existe, de modo que la pérdida no tendría por dónde enseñarle a aa a puntuar mejor. Con el softmax, ci\mathbf{c}_i depende de las puntuaciones de forma continua, y la regla de la cadena de su lección atraviesa los pesos igual que atraviesa cualquier otra cosa. Elegir suave es lo que hace que elegir se pueda aprender.

Falta la función aa y no se define aquí. Lo que esta lección necesita de ella es sólo que devuelva un número y que sea derivable; qué forma tiene, cuántos parámetros lleva y cómo se entrenan es el asunto de la lección siguiente.

El contexto por pasos, en NumPy

La celda toma la alineación del explorable —mis números, otra vez— y un estado por token de la entrada, y hace tres cosas. Comprueba lo único que no se ve mirando el mapa, que cada fila sume uno. Construye los ocho contextos de golpe, porque la mezcla de todas las filas a la vez es un producto de matrices. Y repite la cuenta con la alineación congelada, para poder comparar.

import numpy as np

# La alineación del explorable: fila = paso de la salida, columna = posición de la entrada.
FUENTE = ["ayer", "leí", "un", "libro", "muy", "bueno"]
SALIDA = ["yesterday", "i", "read", "a", "very", "good", "book", "<EOS>"]
alfa = np.array([[0.86, 0.06, 0.03, 0.02, 0.02, 0.01], # yesterday
[0.09, 0.74, 0.05, 0.07, 0.03, 0.02], # i <- de "leí"
[0.04, 0.81, 0.04, 0.06, 0.03, 0.02], # read <- de "leí" otra vez
[0.02, 0.05, 0.72, 0.15, 0.03, 0.03], # a
[0.02, 0.03, 0.04, 0.06, 0.79, 0.06], # very
[0.02, 0.02, 0.03, 0.05, 0.10, 0.78], # good
[0.02, 0.03, 0.06, 0.76, 0.05, 0.08], # book <- vuelve a "libro"
[0.02, 0.03, 0.03, 0.12, 0.07, 0.73]]) # <EOS>
# Un estado por token de la entrada: d_h = 5 coordenadas en (-1, 1), como un tanh.
H = np.array([[ 0.82, -0.31, 0.10, 0.45, -0.62], [-0.55, 0.71, 0.28, -0.14, 0.36],
[ 0.12, 0.44, -0.77, 0.21, 0.08], [ 0.39, -0.68, 0.52, 0.77, -0.11],
[-0.24, 0.15, 0.63, -0.58, 0.49], [ 0.66, 0.29, -0.35, 0.18, 0.74]])
T_y, T_x = alfa.shape

# Lo que el mapa no deja ver, y sin lo cual esto no seria una mezcla:
assert np.allclose(alfa.sum(axis=1), 1.0), "hay una fila que no suma 1"

C = alfa @ H # c_i = sum_j alfa_ij * h_j, los T_y a la vez
fijo = np.zeros_like(alfa); fijo[:, -1] = 1.0 # el resumen fijo, escrito como alineacion
C_fijo = fijo @ H

print("paso escribe tira sobre todo de peso distancia al resumen fijo")
for i in range(T_y):
j = int(alfa[i].argmax())
print(" %d %-11s %-18s %.2f %.2f"
% (i + 1, SALIDA[i], FUENTE[j], alfa[i, j], np.linalg.norm(C[i] - C_fijo[i])))

print()
print("contextos distintos: con atencion %d de %d, con el resumen fijo %d de %d"
% (len(np.unique(C.round(9), axis=0)), T_y,
len(np.unique(C_fijo.round(9), axis=0)), T_y))
print("y ese unico es el estado de '%s':" % FUENTE[-1], np.allclose(C_fijo[0], H[-1]))
numpy

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

La columna del medio es el mapa leído en voz alta, y las dos últimas líneas son el hallazgo: con la alineación de arriba los ocho pasos reciben ocho vectores distintos, y con la congelada reciben el mismo ocho veces —el estado del último token de la entrada—. La última columna dice cuánto se estaba perdiendo cada paso, y ahí hay algo que mirar despacio, porque no se pierde lo mismo en todos. Los dos valores pequeños, 0.250.25 y 0.320.32, son los de good y <EOS>, que son los dos únicos pasos que tiran de bueno —justo el token que el resumen fijo llevaba dentro—; los de yesterday y book pasan de 1.351.35. El resumen fijo no era igual de malo en todas partes: acertaba donde la frase termina y fallaba en todo lo demás, que es por dónde se rompía en la lección anterior. Cambia un peso de la matriz alfa y ejecuta otra vez: la assert se queja antes de que puedas mirar nada, que es justo su trabajo.

Comprueba tu intuición

Cuatro preguntas: sobre qué normaliza el softmax, qué modelo sale al congelar la alineación, por qué mezclar en vez de elegir, y cuánto cuesta.

En αij=exp(eij)/kexp(eik)\alpha_{ij} = \exp(e_{ij}) / \sum_k \exp(e_{ik}), ¿sobre qué recorre kk el denominador?

Fijas αij=1\alpha_{ij} = 1 cuando j=Txj = T_x y αij=0\alpha_{ij} = 0 en todo lo demás, igual en todos los pasos ii. ¿Qué modelo te queda?

En vez de mezclar, el paso ii podría quedarse con el estado mejor puntuado: ci=hˉj\mathbf{c}_i = \bar{\mathbf{h}}_{j^{\star}} con j=argmaxjeijj^{\star} = \arg\max_j e_{ij}. ¿Qué se rompe?

Traduces ayer leí un libro muy bueno —seis tokens— y el decoder escribe siete palabras más el símbolo de fin de secuencia, <EOS>: ocho pasos. ¿Cuántas puntuaciones eije_{ij} calcula el mecanismo para esa sola frase?

puntuaciones

A margin of ±0 is accepted.


El mecanismo está completo salvo por una pieza, y es la que decide todo lo demás. Los pesos de esta lección no los ha calculado nadie: los escribí yo en una matriz, eligiéndolos para que el mapa enseñara una reordenación y un token compartido. Un modelo de verdad no tiene a nadie que se los escriba. Tiene que puntuar cada pareja (si1,hˉj)(\mathbf{s}_{i-1}, \bar{\mathbf{h}}_j) por su cuenta, con una función que empieza sabiendo tan poco como el resto de la red y que aprende al mismo tiempo que ella.

Esa función es aa, y la lección siguiente, sobre la atención de Bahdanau, es la que la escribe: le da sus propios parámetros —una capa pequeña, con su matriz y su vector—, la mete en el mismo descenso de gradiente que ya entrena al encoder y al decoder, y desarrolla por dónde le llega el gradiente, que resulta venir por dos caminos a la vez. Que haya más de una forma razonable de puntuar una pareja, y que la elección entre ellas termine decidiendo cómo es el bloque 5, ocupa el resto de este.

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.