Atención de Bahdanau (aditiva)
29 min de lectura
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 . Lo que no dejó escrito es lo que los produce. Antes del softmax hay una función, , 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 y . La de a, la siguiente, deja esa misma posición en . El estado 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.
Una red pequeña que puntúa parejas
La puntuación aditiva proyecta los dos estados a un espacio propio, de coordenadas, y los suma dentro de él:
con , 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.
con . Es un perceptrón multicapa (multilayer perceptron, MLP) del bloque 2 en su tamaño mínimo: una capa oculta de unidades con , 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: . No aparece , no aparece y no aparece la pareja. Hay un solo juego de pesos y con él se puntúan las parejas de esta frase y las de todas las demás, que es la misma economía de la recurrencia del bloque anterior —un 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 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 y mira qué queda. El vector de lectura entra en la suma y la reparte:
dos números que se suman y que no se han mirado el uno al otro. El primero no lleva : dentro de la fila es la misma cantidad en las 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:
El índice ha desaparecido del lado derecho. Los 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 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 multiplica a un vector conocido:
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 fijo, depende de las 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 , y cuando no, . Separando el término del resto:
El sumatorio es , y esa suma ponderada es .
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 de la mezcla que esa fila ya construyó, y en qué dirección. Un estado que coincida con no recibe corrección ninguna. Tiene sentido que sea una comparación y no una nota: los pesos suman , 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 , el aporta su derivada , y cada matriz recoge el producto exterior con el vector que multiplicaba:
Las 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 : en la mezcla, como uno de los sumandos, y en las puntuaciones, como argumento de . Son dos sitios, así que el gradiente le llega por dos rutas y hay que sumarlas:
A esas dos se les suma la que ya traía el bloque anterior, la que baja desde por la recurrencia del encoder, y ahí está lo interesante: las dos de arriba no pasan por ella. Van del paso del decoder a la posición de la entrada sin recorrer el trecho de recurrencia que separa esa posición del final de la frase, así que no multiplican por ni una sola vez y su longitud no crece cuando 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 .
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])))
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 cuando el reparto uniforme sería , 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 , la fila y la fila 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 es un vector dado— para que lo único que se esté midiendo sea el mecanismo.
# 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())
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 , 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 — la única que puede sacar al mapa de esa
uniformidad. Cambia la escala de va a 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 , cuánto pesa la función , 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 y dejas la puntuación en . ¿Qué le pasa a la alineación?
Un decoder y un encoder de coordenadas, y un modelo de alineación de . ¿Cuántos pesos tiene la función ?
Se acepta un margen de ±0.
Marca las rutas por las que el gradiente de la pérdida llega a desde el mecanismo de atención, dejando fuera la que le llega desde por la recurrencia del encoder.
Marca todas las opciones correctas. Se corrige todo o nada: no hay puntuación parcial.
La derivada de la pérdida respecto de una puntuación sale . ¿Qué dice sobre cuándo se corrige una puntuación?
Traduces una frase de tokens en 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
—la fila es , el estado con el que se puntúa el paso
—, los del encoder H, de forma , y los tres pesos del modelo de
alineación. Devuelve la matriz de forma , ya normalizada.
Es la puntuación aditiva de la lección, , seguida de un softmax por filas. No escribas en los arrays que recibes.
La primera comprobación descarga el intérprete de Python (~15 MB); después queda en la caché del navegador. Este desafío se resuelve mejor con un teclado físico: en el móvil puedes leerlo y volver luego.
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 , y eso obliga a construir la rejilla de vectores de coordenadas de una en una. Una frase de sesenta tokens traducida a sesenta pide de esos vectores, y ninguna de las operaciones que los producen es un producto de matrices grande: son sumas y 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.
Para profundizar1 fuente · 1 paper
De dónde sale lo de esta lección, y dónde seguir si quieres más. Nada de aquí hace falta para continuar el curso.
- Neural Machine Translation by Jointly Learning to Align and Translate
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é.