Atención de Bahdanau (aditiva)

Atención de Bahdanau (aditiva)

29 min read

Un modelo con atención necesita, en cada paso de la salida, un peso por cada posición de la entrada. La lección anterior, sobre la idea de atención, dejó montado todo lo que rodea a esos pesos: salen de un softmax que recorre las posiciones de la fuente, sirven para mezclar los estados del encoder en un vector de contexto propio de cada paso, y ese vector entra en la recurrencia del decoder con su matriz Wch\mathbf{W}_{ch}. Lo que no dejó escrito es lo que los produce. Antes del softmax hay una función, aa, que recibe una pareja —dónde está el decoder y qué hay en una posición de la entrada— y devuelve una puntuación, y de ella sólo se pidió que fuese un número y que se pudiera derivar. Esta lección la escribe, con la forma que le dio Bahdanau en 2014.

La pregunta es pequeña y concreta. ¿Cuál es la función más simple que se traga dos vectores y devuelve un número que signifique «estos dos se corresponden»? Hay una respuesta corta —una red diminuta— y hay una restricción que decide su forma, y la restricción enseña más que la respuesta. La puntuación tiene que ser función de la pareja. Una función que opine por separado sobre cada vector y sume las dos opiniones no puntúa una correspondencia: puntúa dos cosas sueltas, y más abajo se ve que eso destruye el mecanismo entero.

Vuelve al mapa de la lección anterior y mira tres filas. La de i y la de read se llevan casi todo de la misma posición, la de leí, con pesos de 0.740.74 y 0.810.81. La de a, la siguiente, deja esa misma posición en 0.050.05. El estado hˉ2\bar{\mathbf{h}}_2 no ha cambiado entre esas tres filas —el encoder lo calculó una vez y ahí sigue—, y sin embargo recibe dos puntuaciones altas y una baja. Lo único distinto es con qué estado del decoder se le comparaba cada vez.

Las tres filas del ejemplo son la 2, la 3 y la 4. Los pesos siguen siendo los que puse yo a mano; lo que esta lección añade es de dónde saldrían si nadie los pusiera.

Una red pequeña que puntúa parejas

La puntuación aditiva proyecta los dos estados a un espacio propio, de dad_a coordenadas, y los suma dentro de él:

pij=Wsasi1+WhahˉjRda,\mathbf{p}_{ij} = \mathbf{W}_{sa}\mathbf{s}_{i-1} + \mathbf{W}_{ha}\bar{\mathbf{h}}_j \in \mathbb{R}^{d_a},

con Wsa,WhaRda×dh\mathbf{W}_{sa}, \mathbf{W}_{ha} \in \mathbb{R}^{d_a \times d_h}, una por lado. Esa suma es la preactivación de una capa oculta, y lo que sigue es lo que ya hace cualquier capa: aplastarla y leerla.

eij=a(si1,hˉj)=vatanh(pij)R,e_{ij} = a\left(\mathbf{s}_{i-1}, \bar{\mathbf{h}}_j\right) = \mathbf{v}_a^{\top}\tanh\left(\mathbf{p}_{ij}\right) \in \mathbb{R},

con vaRda\mathbf{v}_a \in \mathbb{R}^{d_a}. Es un perceptrón multicapa (multilayer perceptron, MLP) del bloque 2 en su tamaño mínimo: una capa oculta de dad_a unidades con tanh\tanh, y una salida lineal de una sola neurona sin sesgo. El sesgo sobra porque sumaría el mismo número a toda la fila, y una constante por fila no sobrevive al softmax.

Cuenta los pesos: 2dadh+da2 d_a d_h + d_a. No aparece TxT_x, no aparece TyT_y y no aparece la pareja. Hay un solo juego de pesos y con él se puntúan las TxTyT_x \cdot T_y parejas de esta frase y las de todas las demás, que es la misma economía de la recurrencia del bloque anterior —un Whh\mathbf{W}_{hh} para todas las posiciones— aplicada a una rejilla en vez de a una fila. Nada de esto obliga además a que el encoder y el decoder midan lo mismo: las dos matrices tienen dad_a filas cada una, y sus columnas son cosa de cada lado.

Sin la no linealidad, todos los pasos mirarían al mismo sitio

Quita el tanh\tanh y mira qué queda. El vector de lectura entra en la suma y la reparte:

eij=vaWsasi1+vaWhahˉj,e_{ij} = \mathbf{v}_a^{\top}\mathbf{W}_{sa}\mathbf{s}_{i-1} + \mathbf{v}_a^{\top}\mathbf{W}_{ha}\bar{\mathbf{h}}_j,

dos números que se suman y que no se han mirado el uno al otro. El primero no lleva jj: dentro de la fila ii es la misma cantidad en las TxT_x puntuaciones. Y una constante sumada a una fila entera no cambia su softmax —es la invariancia al desplazamiento de la lección sobre funciones de pérdida, la misma que permite restar el máximo antes de exponenciar—, así que se cancela arriba y abajo:

αij=exp(vaWhahˉj)k=1Txexp(vaWhahˉk).\alpha_{ij} = \frac{\exp\left(\mathbf{v}_a^{\top}\mathbf{W}_{ha}\bar{\mathbf{h}}_j\right)}{\sum_{k=1}^{T_x} \exp\left(\mathbf{v}_a^{\top}\mathbf{W}_{ha}\bar{\mathbf{h}}_k\right)}.

El índice ii ha desaparecido del lado derecho. Los TyT_y pasos reciben la misma fila de pesos, y con ella el mismo vector de contexto, de modo que el modelo vuelve a ser el resumen fijo de las lecciones anteriores: una alineación aprendida, sí, pero congelada, sólo que congelada donde diga el entrenamiento en vez de en el último token. El tanh\tanh es lo que impide esa separación. Aplasta la suma antes de leerla, y una función no lineal de una suma no se deja repartir entre sus sumandos, así que la puntuación no puede volver a partirse en una opinión sobre el decoder más otra sobre el encoder.

Por dónde le llega el gradiente

Ahora hay parámetros que entrenar, y la pregunta es si la pérdida los alcanza. La cadena empieza en la mezcla, donde αij\alpha_{ij} multiplica a un vector conocido:

αij=(ci)hˉj,\frac{\partial \ell}{\partial \alpha_{ij}} = \left(\nabla_{\mathbf{c}_i}\ell\right)^{\top}\bar{\mathbf{h}}_j,

y sigue por el softmax, que reparte cada puntuación entre todos los pesos de su fila.

El paso por el softmax, con sus dos casos

Con ii fijo, αij\alpha_{ij} depende de las TxT_x puntuaciones de la fila, así que la derivada respecto de una de ellas suma sobre todas. La derivada del softmax tiene dos casos —los mismos de la lección sobre funciones de pérdida—: cuando el peso y la puntuación son el mismo índice vale αij(1αij)\alpha_{ij}\left(1 - \alpha_{ij}\right), y cuando no, αikαij-\alpha_{ik}\alpha_{ij}. Separando el término k=jk = j del resto:

eij=αijαijαijk=1Txαikαik.\frac{\partial \ell}{\partial e_{ij}} = \alpha_{ij}\frac{\partial \ell}{\partial \alpha_{ij}} - \alpha_{ij}\sum_{k=1}^{T_x} \alpha_{ik}\frac{\partial \ell}{\partial \alpha_{ik}}.

El sumatorio es (ci)kαikhˉk\left(\nabla_{\mathbf{c}_i}\ell\right)^{\top}\sum_k \alpha_{ik}\bar{\mathbf{h}}_k, y esa suma ponderada es ci\mathbf{c}_i.

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

Léela despacio, porque dice cómo aprende a mirar. Una puntuación no se corrige según lo bueno que sea su estado, sino según cuánto se aparte hˉj\bar{\mathbf{h}}_j de la mezcla que esa fila ya construyó, y en qué dirección. Un estado que coincida con ci\mathbf{c}_i no recibe corrección ninguna. Tiene sentido que sea una comparación y no una nota: los pesos suman 11, así que subir uno es bajar los demás, y no hay forma de premiar a una posición sin decir a costa de cuál.

De ahí para atrás todo es cadena conocida. La lectura devuelve tanh(pij)\tanh(\mathbf{p}_{ij}), el tanh\tanh aporta su derivada 1tanh21 - \tanh^{2}, y cada matriz recoge el producto exterior con el vector que multiplicaba:

va=i,jeijtanh(pij),pij=eijva(1tanh2(pij)),\nabla_{\mathbf{v}_a}\ell = \sum_{i,j} \frac{\partial \ell}{\partial e_{ij}}\tanh\left(\mathbf{p}_{ij}\right), \qquad \nabla_{\mathbf{p}_{ij}}\ell = \frac{\partial \ell}{\partial e_{ij}}\,\mathbf{v}_a \odot \left(1 - \tanh^{2}\left(\mathbf{p}_{ij}\right)\right), Wha=i,j(pij)hˉj,Wsa=i,j(pij)si1.\nabla_{\mathbf{W}_{ha}}\ell = \sum_{i,j} \left(\nabla_{\mathbf{p}_{ij}}\ell\right)\bar{\mathbf{h}}_j^{\top}, \qquad \nabla_{\mathbf{W}_{sa}}\ell = \sum_{i,j} \left(\nabla_{\mathbf{p}_{ij}}\ell\right)\mathbf{s}_{i-1}^{\top}.

Las TxTyT_x \cdot T_y parejas suman sobre los mismos tres arrays, que es lo que significa compartir pesos.

Y hay una consecuencia que no está en los parámetros. Fíjate en dónde aparece hˉj\bar{\mathbf{h}}_j: en la mezcla, como uno de los sumandos, y en las TyT_y puntuaciones, como argumento de aa. Son dos sitios, así que el gradiente le llega por dos rutas y hay que sumarlas:

hˉj=i=1Tyαijci+i=1TyWhapij.\nabla_{\bar{\mathbf{h}}_j}\ell = \sum_{i=1}^{T_y} \alpha_{ij}\nabla_{\mathbf{c}_i}\ell + \sum_{i=1}^{T_y} \mathbf{W}_{ha}^{\top}\nabla_{\mathbf{p}_{ij}}\ell.

A esas dos se les suma la que ya traía el bloque anterior, la que baja desde hˉj+1\bar{\mathbf{h}}_{j+1} por la recurrencia del encoder, y ahí está lo interesante: las dos de arriba no pasan por ella. Van del paso ii del decoder a la posición jj de la entrada sin recorrer el trecho de recurrencia que separa esa posición del final de la frase, así que no multiplican por Whhenc\mathbf{W}^{\text{enc}}_{hh} ni una sola vez y su longitud no crece cuando jj se aleja. El producto de matrices que la lección sobre el gradiente desvanecido señalaba como causa del problema no aparece en estas dos rutas.

La puntuación, en NumPy

La primera celda calcula la rejilla entera. Guarda los estados por filas —una fila por paso, una fila por posición— y por eso las matrices aparecen transpuestas respecto de las ecuaciones. Mira dos cosas: cuánto se parece a un mapa el reparto que sale de unos pesos recién inicializados, y qué pasa al pedir la versión sin tanh\tanh.

import numpy as np

T_y, T_x, d_h, d_a = 8, 6, 5, 4
rng = np.random.default_rng(4)
S = rng.normal(size=(T_y, d_h)) * 0.6 # fila i: el estado s_{i-1} con el que se puntua el paso i
H = rng.normal(size=(T_x, d_h)) * 0.6 # fila j: el estado del encoder en la posicion j
Wsa = rng.normal(size=(d_a, d_h)) * 0.5 # sin entrenar: nadie ha corregido esto todavia
Wha = rng.normal(size=(d_a, d_h)) * 0.5
va = rng.normal(size=d_a) * 0.5


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


def puntua(S, H, con_tanh=True):
Ps = S @ Wsa.T # (T_y, d_a): una proyeccion por paso
Ph = H @ Wha.T # (T_x, d_a): una por posicion
P = Ps[:, None, :] + Ph[None, :, :] # (T_y, T_x, d_a): la suma, pareja a pareja
return (np.tanh(P) if con_tanh else P) @ va # (T_y, T_x)


alfa = softmax_filas(puntua(S, H))
assert np.allclose(alfa.sum(axis=1), 1.0), "hay una fila que no suma 1"

print("pesos:", 2 * d_a * d_h + d_a, " formas: alfa", alfa.shape)
print("fila 1:", np.round(alfa[0], 3))
print("mayor peso de cada fila:", np.round(alfa.max(axis=1), 3), " uniforme:", round(1 / T_x, 3))

lineal = softmax_filas(puntua(S, H, con_tanh=False))
print()
print("sin tanh, fila 1:", np.round(lineal[0], 3))
print("sin tanh, fila 8:", np.round(lineal[7], 3))
print("sin tanh, todas las filas iguales:", bool(np.allclose(lineal, lineal[0])))
numpy

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

El mapa recién inicializado no dice nada: el mayor peso de cada fila ronda 0.180.18 cuando el reparto uniforme sería 0.1670.167, así que las ocho filas son prácticamente la media de la frase. La alineación interpretable del explorable no viene de la arquitectura, viene del entrenamiento; lo que la arquitectura garantiza es que haya por dónde entrenarla. Y las dos últimas líneas son la sección anterior comprobada a mano: sin el tanh\tanh, la fila 11 y la fila 88 salen iguales hasta el último decimal.

La segunda celda persigue las dos rutas por separado. Necesita una pérdida y usa la más barata posible —una cuyo gradiente respecto de cada ci\mathbf{c}_i es un vector dado— para que lo único que se esté midiendo sea el mecanismo.

import numpy as np

# Necesita la celda anterior: S, H, Wsa, Wha, va, puntua, softmax_filas.
# Perdida de juguete: l = suma_i G[i] . c_i, de modo que el gradiente de l respecto
# de c_i es exactamente G[i], sin derivar el decoder ni la entropia cruzada.
G = np.random.default_rng(9).normal(size=(T_y, d_h))


def mezcla(H):
alfa = softmax_filas(puntua(S, H))
return alfa @ H, alfa


def perdida(H):
C, _ = mezcla(H)
return float((G * C).sum())


C, alfa = mezcla(H)
valor = alfa.T @ G # ruta 1: como sumando de c_i
dl_de = alfa * (G @ H.T - (G * C).sum(axis=1, keepdims=True)) # alfa_ij * G_i . (h_j - c_i)
P = (S @ Wsa.T)[:, None, :] + (H @ Wha.T)[None, :, :]
D = dl_de[:, :, None] * va * (1 - np.tanh(P) ** 2) # (T_y, T_x, d_a)
puntuacion = np.einsum("ijk,kl->jl", D, Wha) # ruta 2: como argumento de a

num = np.zeros_like(H) # sondeo numerico, derivada central
for j in range(T_x):
for l in range(d_h):
Hmas = H.copy(); Hmas[j, l] += 1e-6
Hmenos = H.copy(); Hmenos[j, l] -= 1e-6
num[j, l] = (perdida(Hmas) - perdida(Hmenos)) / 2e-6

print("norma de la ruta del valor: %.4f" % np.linalg.norm(valor))
print("norma de la ruta de la puntuacion: %.4f" % np.linalg.norm(puntuacion))
print("error de cada una por separado: %.1e %.1e"
% (np.abs(valor - num).max(), np.abs(puntuacion - num).max()))
print("error de las dos sumadas: %.1e" % np.abs(valor + puntuacion - num).max())
numpy

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

La última línea es la comprobación: sumadas, las dos rutas reproducen el gradiente numérico hasta el orden de 101010^{-10}, y por separado cada una se equivoca en la primera cifra. Las normas dicen algo más, y es honesto mirarlo: con los pesos sin entrenar la ruta del valor es once veces mayor que la de la puntuación. Al principio casi todo el gradiente llega por donde llegaría aunque la alineación fuese uniforme, y es la ruta pequeña —la que pasa por aa— la única que puede sacar al mapa de esa uniformidad. Cambia la escala de va a 2.02.0 y ejecuta otra vez: las puntuaciones se separan, los pesos se afilan y las dos normas se acercan.

Comprueba tu intuición

Cinco preguntas: qué se rompe sin el tanh\tanh, cuánto pesa la función aa, por dónde le llega el gradiente a un estado del encoder, qué dice la derivada de una puntuación, y qué parte del cálculo se puede sacar del bucle.

Quitas el tanh\tanh y dejas la puntuación en eij=va(Wsasi1+Whahˉj)e_{ij} = \mathbf{v}_a^{\top}\left(\mathbf{W}_{sa}\mathbf{s}_{i-1} + \mathbf{W}_{ha}\bar{\mathbf{h}}_j\right). ¿Qué le pasa a la alineación?

Un decoder y un encoder de dh=64d_h = 64 coordenadas, y un modelo de alineación de da=32d_a = 32. ¿Cuántos pesos tiene la función aa?

pesos

A margin of ±0 is accepted.

Marca las rutas por las que el gradiente de la pérdida llega a hˉj\bar{\mathbf{h}}_j desde el mecanismo de atención, dejando fuera la que le llega desde hˉj+1\bar{\mathbf{h}}_{j+1} por la recurrencia del encoder.

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

La derivada de la pérdida respecto de una puntuación sale /eij=αij(ci)(hˉjci)\partial\ell/\partial e_{ij} = \alpha_{ij}\left(\nabla_{\mathbf{c}_i}\ell\right)^{\top}\left(\bar{\mathbf{h}}_j - \mathbf{c}_i\right). ¿Qué dice sobre cuándo se corrige una puntuación?

Traduces una frase de TxT_x tokens en TyT_y pasos. ¿Cuántas veces hay que calcular cada pieza de la puntuación aditiva?

Escribe alineacion(S, H, Wsa, Wha, va). Recibe los estados del decoder S, de forma (Ty,dh)(T_y, d_h) —la fila ii es si1\mathbf{s}_{i-1}, el estado con el que se puntúa el paso ii—, los del encoder H, de forma (Tx,dh)(T_x, d_h), y los tres pesos del modelo de alineación. Devuelve la matriz α\alpha de forma (Ty,Tx)(T_y, T_x), ya normalizada.

Es la puntuación aditiva de la lección, eij=vatanh(Wsasi1+Whahˉj)e_{ij} = \mathbf{v}_a^{\top}\tanh\left(\mathbf{W}_{sa}\mathbf{s}_{i-1} + \mathbf{W}_{ha}\bar{\mathbf{h}}_j\right), seguida de un softmax por filas. 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 una cuenta pendiente de la última pregunta. Las dos proyecciones se sacan del bucle porque cada una depende de un solo índice, pero entre la suma y el vector de lectura hay un tanh\tanh, y eso obliga a construir la rejilla de TxTyT_x \cdot T_y vectores de dad_a coordenadas de una en una. Una frase de sesenta tokens traducida a sesenta pide 36003\,600 de esos vectores, y ninguna de las operaciones que los producen es un producto de matrices grande: son sumas y tanh\tanh sobre un bloque de tres dimensiones. El hardware que entrena estas redes hace una cosa mucho mejor que las demás, y no es ésa.

Hay otra forma de puntuar una pareja que evita el problema entero, y su fórmula cabe en cinco símbolos: la atención de Luong, multiplicativa, que es la lección siguiente. En vez de proyectar, sumar y aplastar, multiplica los dos estados a través de una matriz, con lo que la rejilla completa de puntuaciones pasa a ser un producto de matrices. La lección compara las dos puntuaciones —qué cuesta cada una, qué puede hacer cada una— y ahí empieza a decidirse cómo será el bloque 5.

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.

  • Neural Machine Translation by Jointly Learning to Align and Translate
    paperBahdanau, Cho y Bengio, 2015arXiv:1409.0473EN

    La puntuación de esta lección está en su apéndice A.2.2: v_a por el tanh de dos proyecciones sumadas. Allí nota que fijar el contexto al último estado da el encoder-decoder pelado —tu alineación congelada—. Del tanh no dice por qué.