Atención de Luong (multiplicativa)
27 min read
Un modelo con atención puntúa parejas: un número por cada paso de la salida y cada posición de la entrada, y de esos números salen los pesos con los que se mezclan los estados del encoder. La lección anterior, sobre la atención aditiva de Bahdanau, llenó esa casilla con una red pequeña —una capa oculta de unidades y una lectura lineal, entrenada junto al resto— y dejó una factura apuntada: puntuar una pareja obliga a fabricar un vector entero de coordenadas antes de reducirlo al número que hacía falta.
La pregunta de esta lección es si hace falta una capa oculta para decir que dos vectores se corresponden. El álgebra lineal lleva desde el bloque 1 devolviendo ese número sin ninguna red de por medio, y en 2015 Luong, Pham y Manning propusieron usarlo tal cual, con un solo retoque. Lo que sale no se parece a una red: se parece a un producto. Y como la calidad de las dos formas es comparable —eso lo miden los experimentos de sus artículos, no una derivación—, lo que decide entre ellas es todo lo demás: cuánto cuestan, qué memoria mueven y qué forma tiene su cálculo.
Vuelve a lo que la puntuación tiene que hacer. Un número que sube cuando dos vectores apuntan al mismo sitio ya está en el curso desde la lección sobre la bolsa de palabras: el producto escalar , que la similitud coseno se limita a normalizar. Sería una puntuación gratis —cero parámetros— y sería función de la pareja, que es la condición que la lección anterior declaró irrenunciable: no se puede partir en una opinión sobre un estado más otra sobre el otro.
Lo que falla es una suposición escondida. El producto escalar suma coordenada contra coordenada, de modo que da por hecho que la coordenada del decoder habla de lo mismo que la coordenada del encoder, y con la misma escala. Son dos recurrencias distintas, con dos juegos de pesos y dos trabajos —una lee español, la otra escribe inglés—, y nada las obliga a haber elegido las mismas direcciones para las mismas cosas; ni siquiera a tener el mismo número de coordenadas. De modo que primero se cambia de base y después se mide.
Multiplicar en vez de sumar
La puntuación multiplicativa mete una matriz entre los dos estados y no hace nada más:
con . Se lee de derecha a izquierda: es el estado del encoder reescrito en las coordenadas del decoder, y el producto escalar con mide lo que mide siempre un producto escalar. No hay capa oculta, no hay saturación y no hay ningún espacio propio de la función: el de la lección anterior desaparece del problema.
Cuenta los pesos: , y ninguno depende de la pareja ni de las longitudes, igual que antes. Con son , contra los que pedía la aditiva con . Prácticamente los mismos, lo cual es útil saberlo pronto: lo que separa a las dos puntuaciones no es cuántos pesos tienen, sino qué operaciones hacen con ellos.
Aquí conviene decir qué se está dejando fuera. Luong y sus coautores cambian además dónde entra el contexto —puntúan con el estado del paso en curso, , y mezclan el contexto después de la recurrencia en vez de dentro de ella— y este curso no lo sigue: mantiene la colocación de la lección sobre la idea de atención, para que entre aquella lección, la anterior y ésta cambie una sola cosa. Lo que se compara es la función , con todo lo demás quieto.
El gradiente es más corto que el de la aditiva, y por la misma razón por la que el cálculo lo es. Las tres derivadas de una puntuación salen de derivar un producto de tres factores:
Ninguna lleva un factor , porque no hay que derivar. Y todo lo que hay por encima de en la cadena es de la lección anterior, sin tocar una coma: el paso por el softmax y la derivada nunca mencionaron qué forma tenía . La atención se queda igual; lo que cambia es el modelo de alineación.
Y la acumulación sobre las parejas también es un producto
es uno solo para las parejas, así que su gradiente las suma todas:
que es la rejilla de derivadas puesta entre los estados del decoder y los del encoder —dos productos de matrices, con la rejilla en medio—. El camino de vuelta tiene la misma forma que el de ida, que es lo que hace que la sección siguiente valga para los dos sentidos.
Sin matriz: el producto escalar
El caso extremo es tomar igual a la identidad, con lo que la puntuación se queda en
y la función se queda sin un solo parámetro que entrenar. Lo que pide a cambio es justo lo que la matriz servía para no pedir: que los dos lados midan lo mismo — coordenadas en los dos— y que además hayan aprendido a usarlas de la misma forma. Es una exigencia fuerte cuando el encoder y el decoder son dos redes distintas, entrenadas cada una para su idioma, y por eso la matriz es la opción que Luong deja por defecto. La variante sin parámetros no se va del curso —es la forma más barata de puntuar que existe—, pero sólo se sostiene donde los dos lados comparten base.
La rejilla entera en un producto de matrices
Hasta aquí las puntuaciones se han escrito de una en una. Apílalas. Guarda los estados del encoder por filas en —la fila es , transpuesta porque en este curso los vectores son columnas— y los del decoder en , cuya fila es , el estado con el que se puntúa el paso . Entonces las puntuaciones son un solo objeto:
con el softmax aplicado por filas, que es sobre las posiciones de la entrada, como en la lección sobre la idea de atención. Comprueba la casilla : la fila de es y la columna de es , así que su producto es y no otra cosa.
La aditiva no admite esta escritura, y no por falta de maña. Sus dos proyecciones sí se sacan fuera, pero entre la suma y la lectura está el , de modo que hay que materializar un bloque de tres dimensiones —una pareja, y dentro de cada pareja coordenadas— antes de poder reducirlo a números.
Los recuentos ponen precio a esa diferencia. Toma y , que es una frase larga y un modelo pequeño. En productos y sumas las dos andan cerca: unos millones la aditiva contra unos millones la multiplicativa, un factor menor que dos, y en ambas la mayor parte se la llevan las proyecciones. Donde se separan es en lo otro. La aditiva evalúa veces el y la multiplicativa ninguna. Y sobre todo, la aditiva tiene que guardar esos números — MB por cada pareja de frases, antes de multiplicar por el tamaño del batch— mientras que a la multiplicativa le bastan las puntuaciones y los estados ya proyectados. Entre las dos rejillas el factor es exactamente : un número por pareja contra números por pareja.
Y hay una diferencia de forma que no aparece en ninguno de esos números. El trabajo de la multiplicativa son dos productos de matrices; el de la aditiva son sumas y sobre un bloque tridimensional, coordenada a coordenada. La misma cantidad de aritmética repartida de las dos maneras no cuesta lo mismo en el hardware que entrena estas redes, que está construido justamente para lo primero —es la lección sobre el forward pass, medida allí con un bucle contra un producto—.
Falta lo honesto, y es lo que abre el bloque siguiente. Escribir supone tener los estados del decoder, y el decoder es recurrente: necesita , que necesita la fila de , que necesita . Mientras traduce, la rejilla se llena fila a fila y el producto grande no llega a ocurrir; lo que sí se saca del bucle es , que sólo depende de la fuente y se calcula una vez por frase, dejando cada paso del decoder en un solo producto de matriz por vector:
la fila de puntuaciones sin normalizar. La puntuación ya no es el cuello de botella del cálculo. Lo que estrangula el paralelismo, ahora, es la recurrencia que produce los estados.
Las dos puntuaciones, en NumPy
La primera celda calcula la alineación multiplicativa de la forma corta y comprueba que el producto no dice nada distinto de puntuar las parejas una a una. Mira también las dos últimas líneas: con igual a la identidad, la puntuación se queda en el producto escalar de los dos estados.
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
Wa = rng.normal(size=(d_h, d_h)) * 0.4 # sin entrenar
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)
E = (S @ Wa) @ H.T # (T_y, T_x): la rejilla entera, dos productos
A = softmax_filas(E)
E_bucle = np.array([[S[i] @ Wa @ H[j] for j in range(T_x)] for i in range(T_y)])
print("formas: S", S.shape, " H", H.shape, " Wa", Wa.shape, " -> A", A.shape)
print("cada fila suma 1:", bool(np.allclose(A.sum(axis=1), 1.0)))
print("el producto y el bucle coinciden:", bool(np.allclose(E, E_bucle)))
print("pesos de a: multiplicativa", d_h * d_h, " aditiva", 2 * d_a * d_h + d_a)
print("mayor peso de cada fila:", np.round(A.max(axis=1), 3), " uniforme:", round(1 / T_x, 3))
print()
# Wa = I: no hay cambio de base y la puntuacion es el producto escalar de los dos estados.
identidad = softmax_filas((S @ np.eye(d_h)) @ H.T)
print("con Wa = I, fila 1:", np.round(identidad[0], 3))
print("igual que puntuar con S @ H.T:", bool(np.allclose(identidad, softmax_filas(S @ H.T))))
La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador.
La línea que importa es la del bucle: puntuar pareja a pareja da los mismos números que el producto, que es lo único que hay que comprobar de la forma corta. La del mayor peso de cada fila trae una sorpresa. Con los pesos recién inicializados el mapa ya reparte de forma desigual —una fila se lleva cuando el reparto uniforme sería —, y eso no pasaba en la lección anterior con un tamaño parecido, donde el mayor peso de cada fila rondaba el uniforme. La diferencia está en el : acota cada coordenada antes de que la lea, así que la puntuación aditiva nace pequeña. La multiplicativa no tiene nada que la acote, y su tamaño crece con el de los estados. Ese pico no es una alineación, es la dirección que tomó al azar; lo que la fórmula garantiza no es empezar diciendo algo, sino que haya por dónde entrenarlo. Que las puntuaciones sin acotar se disparen al crecer el modelo es un problema de verdad, con un arreglo de una línea que llega en el bloque siguiente.
import time
T_y = T_x = 60
d_h, d_a = 256, 256
rng = np.random.default_rng(0)
S, H = rng.normal(size=(T_y, d_h)) * 0.5, rng.normal(size=(T_x, d_h)) * 0.5
Wa = rng.normal(size=(d_h, d_h)) * 0.05
Wsa, Wha = rng.normal(size=(d_a, d_h)) * 0.05, rng.normal(size=(d_a, d_h)) * 0.05
va = rng.normal(size=d_a)
def aditiva(S, H):
P = (S @ Wsa.T)[:, None, :] + (H @ Wha.T)[None, :, :] # (T_y, T_x, d_a): hay que construirlo
return np.tanh(P) @ va # (T_y, T_x)
def multiplicativa(S, H):
return S @ (Wa @ H.T) # (T_y, T_x): dos productos y nada mas
t0 = time.perf_counter(); Ea = aditiva(S, H)
t1 = time.perf_counter(); Em = multiplicativa(S, H)
t2 = time.perf_counter()
print("misma forma de salida:", Ea.shape, Em.shape)
print("productos y sumas: aditiva %9d multiplicativa %9d"
% ((T_x + T_y) * d_a * d_h + T_x * T_y * d_a, T_x * d_h * d_h + T_x * T_y * d_h))
print("evaluaciones de tanh: aditiva %9d multiplicativa %9d" % (T_x * T_y * d_a, 0))
print("numeros intermedios: aditiva %9d multiplicativa %9d"
% (T_x * T_y * d_a, T_x * d_h + T_x * T_y))
print(" %.1f MB frente a %.2f MB (8 bytes por numero)"
% (T_x * T_y * d_a * 8 / 1e6, (T_x * d_h + T_x * T_y) * 8 / 1e6))
print("tiempo: aditiva %7.1f ms multiplicativa %7.1f ms"
% ((t1 - t0) * 1e3, (t2 - t1) * 1e3))
La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador.
Las tres primeras líneas son la sección anterior contada por el código: la aritmética se parece, el y los MB intermedios no. La última depende de la máquina y del navegador, y aun así dice algo, porque no está midiendo una diferencia pequeña: la multiplicativa tarda alrededor de un orden de magnitud menos, y el factor exacto cambia de una ejecución a otra sin acercarse nunca a uno. Y hay un experimento que lo deja claro. Cambia a y ejecuta otra vez: la aditiva pasa a hacer menos aritmética que la multiplicativa — millones de operaciones contra — y sigue tardando varias veces más que ella. Lo que se está midiendo no es cuántas operaciones hay.
Comprueba tu intuición
Cuatro preguntas: qué añade la matriz, cuánto pesa la puntuación, qué guarda la aditiva que la otra no guarda, y si la rejilla se puede calcular de una vez.
¿Qué añade frente a puntuar con el producto escalar pelado, ?
Un encoder y un decoder de coordenadas. ¿Cuántos pesos tiene la puntuación multiplicativa?
A margin of ±0 is accepted.
Puntúas una pareja de frases de tokens con . ¿Qué guarda la puntuación aditiva que la multiplicativa no guarda?
Con la puntuación multiplicativa, ¿puede un decoder recurrente calcular la rejilla entera de una sola vez mientras traduce?
Escribe alineacion_multiplicativa(S, H, Wa). Recibe los estados del decoder S, de forma
—la fila es —, los del encoder H, de forma ,
y la matriz de la puntuación, de forma . Devuelve la matriz de
forma , ya normalizada por filas.
Es la puntuación de la lección, , seguida de un softmax sobre las posiciones de la entrada. 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 las dos puntuaciones delante hay algo que se ve mejor que con una sola. En las dos, el estado del encoder entra en el mecanismo dos veces y por puertas distintas: una como argumento de , que es donde se decide cuánto peso se lleva, y otra como sumando de , que es donde su contenido acaba llegando al decoder. El estado del decoder, en cambio, entra una sola vez, y siempre del mismo lado: es el que pregunta. Tres papeles, y ninguna de las dos fórmulas exige que los dos primeros los haga el mismo vector —lo hace porque nadie ha propuesto todavía separarlos.
Separarlos y ponerles nombre —consulta, clave y valor— es la lección siguiente, la última del bloque, y es la más importante de este curso. No cambia ni una ecuación de las que ya has visto: reescribe las mismas con las palabras y los símbolos con los que el bloque 5 lo escribe todo, incluida la elección de puntuación que esta lección acaba de justificar. Lo que aquí quedó pendiente —que los estados sigan llegando de uno en uno— tampoco lo arregla ese cambio de nombre, y por eso conviene llegar al bloque siguiente con los nombres ya puestos.
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.
- Effective Approaches to Attention-based Neural Machine Translation
Su §3.1 trae la puntuación multiplicativa (su «general», s por W_a por h) y la variante sin pesos. Ellos además mueven el contexto detrás de la RNN y añaden input feeding (§3.3); el curso solo se queda con la puntuación.