Atención de Luong (multiplicativa)

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 dad_a unidades y una lectura lineal, entrenada junto al resto— y dejó una factura apuntada: puntuar una pareja obliga a fabricar un vector entero de dad_a 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 uv\mathbf{u}^{\top}\mathbf{v}, 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: si1hˉj\mathbf{s}_{i-1}^{\top}\bar{\mathbf{h}}_j 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 kk del decoder habla de lo mismo que la coordenada kk 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:

eij=a(si1,hˉj)=si1WahˉjR,e_{ij} = a\left(\mathbf{s}_{i-1}, \bar{\mathbf{h}}_j\right) = \mathbf{s}_{i-1}^{\top}\mathbf{W}_a\bar{\mathbf{h}}_j \in \mathbb{R},

con WaRdh×dh\mathbf{W}_a \in \mathbb{R}^{d_h \times d_h}. Se lee de derecha a izquierda: Wahˉj\mathbf{W}_a\bar{\mathbf{h}}_j es el estado del encoder reescrito en las coordenadas del decoder, y el producto escalar con si1\mathbf{s}_{i-1} 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 dad_a de la lección anterior desaparece del problema.

Cuenta los pesos: dh2d_h^2, y ninguno depende de la pareja ni de las longitudes, igual que antes. Con dh=64d_h = 64 son 40964\,096, contra los 41284\,128 que pedía la aditiva con da=32d_a = 32. 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, si\mathbf{s}_i, 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 aa, 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:

Waeij=si1hˉj,hˉjeij=Wasi1,si1eij=Wahˉj.\nabla_{\mathbf{W}_a} e_{ij} = \mathbf{s}_{i-1}\bar{\mathbf{h}}_j^{\top}, \qquad \nabla_{\bar{\mathbf{h}}_j} e_{ij} = \mathbf{W}_a^{\top}\mathbf{s}_{i-1}, \qquad \nabla_{\mathbf{s}_{i-1}} e_{ij} = \mathbf{W}_a\bar{\mathbf{h}}_j.

Ninguna lleva un factor 1tanh21 - \tanh^{2}, porque no hay tanh\tanh que derivar. Y todo lo que hay por encima de eije_{ij} en la cadena es de la lección anterior, sin tocar una coma: el paso por el softmax y la derivada /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) nunca mencionaron qué forma tenía aa. 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

Wa\mathbf{W}_a es uno solo para las TxTyT_x \cdot T_y parejas, así que su gradiente las suma todas:

Wa=i=1Tyj=1Txeijsi1hˉj,\nabla_{\mathbf{W}_a}\ell = \sum_{i=1}^{T_y}\sum_{j=1}^{T_x} \frac{\partial \ell}{\partial e_{ij}}\,\mathbf{s}_{i-1}\bar{\mathbf{h}}_j^{\top},

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 Wa\mathbf{W}_a igual a la identidad, con lo que la puntuación se queda en

eij=si1hˉj,e_{ij} = \mathbf{s}_{i-1}^{\top}\bar{\mathbf{h}}_j,

y la función aa 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 —dhd_h 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 HˉRTx×dh\bar{\mathbf{H}} \in \mathbb{R}^{T_x \times d_h} —la fila jj es hˉj\bar{\mathbf{h}}_j^{\top}, transpuesta porque en este curso los vectores son columnas— y los del decoder en SRTy×dh\mathbf{S} \in \mathbb{R}^{T_y \times d_h}, cuya fila ii es si1\mathbf{s}_{i-1}^{\top}, el estado con el que se puntúa el paso ii. Entonces las TxTyT_x \cdot T_y puntuaciones son un solo objeto:

A=softmax(SWaHˉ)RTy×Tx,\mathbf{A} = \text{softmax}\left(\mathbf{S}\mathbf{W}_a\bar{\mathbf{H}}^{\top}\right) \in \mathbb{R}^{T_y \times T_x},

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 (i,j)(i, j): la fila ii de SWa\mathbf{S}\mathbf{W}_a es si1Wa\mathbf{s}_{i-1}^{\top}\mathbf{W}_a y la columna jj de Hˉ\bar{\mathbf{H}}^{\top} es hˉj\bar{\mathbf{h}}_j, así que su producto es eije_{ij} 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 tanh\tanh, de modo que hay que materializar un bloque de tres dimensiones —una pareja, y dentro de cada pareja dad_a coordenadas— antes de poder reducirlo a números.

A la izquierda, un bloque tridimensional de casillas etiquetado T sub y por T sub x por d sub a, con una etiqueta que dice suma, tanh y lectura en cada casilla; una flecha lo reduce a una rejilla plana de T sub y por T sub x. A la derecha, dos rectángulos planos, uno alto de T sub y por d sub h etiquetado S por W sub a y otro ancho de d sub h por T sub x etiquetado H con barra traspuesta, unidos por un signo de producto; una flecha los lleva a una rejilla plana idéntica a la de la izquierda.
Las dos puntuaciones llegan a la misma rejilla. La aditiva construye antes un número por pareja y por coordenada del modelo de alineación; la multiplicativa va de los estados a la rejilla con dos productos y nada en medio.

Los recuentos ponen precio a esa diferencia. Toma Tx=Ty=60T_x = T_y = 60 y dh=da=256d_h = d_a = 256, que es una frase larga y un modelo pequeño. En productos y sumas las dos andan cerca: unos 8.88.8 millones la aditiva contra unos 4.94.9 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 921600921\,600 veces el tanh\tanh y la multiplicativa ninguna. Y sobre todo, la aditiva tiene que guardar esos 921600921\,600 números —7.47.4 MB por cada pareja de frases, antes de multiplicar por el tamaño del batch— mientras que a la multiplicativa le bastan las 36003\,600 puntuaciones y los 6060 estados ya proyectados. Entre las dos rejillas el factor es exactamente dad_a: un número por pareja contra dad_a 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 tanh\tanh 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 S\mathbf{S} supone tener los TyT_y estados del decoder, y el decoder es recurrente: si1\mathbf{s}_{i-1} necesita ci1\mathbf{c}_{i-1}, que necesita la fila i1i-1 de A\mathbf{A}, que necesita si2\mathbf{s}_{i-2}. 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 WaHˉ\mathbf{W}_a\bar{\mathbf{H}}^{\top}, 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:

(HˉWa)si1RTx,\left(\bar{\mathbf{H}}\mathbf{W}_a^{\top}\right)\mathbf{s}_{i-1} \in \mathbb{R}^{T_x},

la fila ii 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 Wa\mathbf{W}_a igual a la identidad, la puntuación se queda en el producto escalar de los dos estados.

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
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))))
numpy

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 0.8510.851 cuando el reparto uniforme sería 0.1670.167—, 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 tanh\tanh: acota cada coordenada antes de que va\mathbf{v}_a 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 Wa\mathbf{W}_a 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 numpy as np
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))
numpy

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 tanh\tanh y los 7.47.4 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 dad_a a 6464 y ejecuta otra vez: la aditiva pasa a hacer menos aritmética que la multiplicativa —2.22.2 millones de operaciones contra 4.94.9— 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 Wa\mathbf{W}_a frente a puntuar con el producto escalar pelado, eij=si1hˉje_{ij} = \mathbf{s}_{i-1}^{\top}\bar{\mathbf{h}}_j?

Un encoder y un decoder de dh=64d_h = 64 coordenadas. ¿Cuántos pesos tiene la puntuación multiplicativa?

pesos

A margin of ±0 is accepted.

Puntúas una pareja de frases de Tx=Ty=60T_x = T_y = 60 tokens con dh=da=256d_h = d_a = 256. ¿Qué guarda la puntuación aditiva que la multiplicativa no guarda?

Con la puntuación multiplicativa, ¿puede un decoder recurrente calcular la rejilla entera SWaHˉ\mathbf{S}\mathbf{W}_a\bar{\mathbf{H}}^{\top} de una sola vez mientras traduce?

Escribe alineacion_multiplicativa(S, H, Wa). Recibe los estados del decoder S, de forma (Ty,ds)(T_y, d_s) —la fila ii es si1\mathbf{s}_{i-1}—, los del encoder H, de forma (Tx,dh)(T_x, d_h), y la matriz de la puntuación, de forma (ds,dh)(d_s, d_h). Devuelve la matriz A\mathbf{A} de forma (Ty,Tx)(T_y, T_x), ya normalizada por filas.

Es la puntuación de la lección, eij=si1Wahˉje_{ij} = \mathbf{s}_{i-1}^{\top}\mathbf{W}_a\bar{\mathbf{h}}_j, 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 aa, que es donde se decide cuánto peso se lleva, y otra como sumando de ci\mathbf{c}_i, 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
    paperLuong, Pham y Manning, 2015arXiv:1508.04025EN

    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.