Escribir no es leer: el decoder y su máscara

Escribir no es leer: el decoder y su máscara

28 min read

La suma y el layer norm de la lección anterior, sobre los residuales, envuelven lo que se les ponga dentro y no preguntan qué es, así que poner una capa encima de otra dejó de ser el problema. El que queda no es de profundidad. Todo lo que este bloque ha construido da por hecho que la frase está entera delante y que las TT posiciones se calculan a la vez, y hay media arquitectura para la que eso es falso: un modelo que traduce no tiene la traducción delante, la escribe, un token detrás de otro. Y al entrenarlo calculándolo todo de golpe —que es lo que se ganó quitando la recurrencia— la respuesta acaba sentada al lado de la posición que tiene que adivinarla.

Míralo en el gato bebe leche, que es lo que el modelo tendría que aprender a escribir. Si las cuatro posiciones se calculan a la vez y cada una mira a todas —que es exactamente lo que hace la capa de la lección sobre la auto-atención—, la posición que debe producir bebe está mirando leche mientras la produce. Adivinar con la respuesta delante sale muy bien: la pérdida baja a cero durante el entrenamiento. Al generar de verdad no hay nada a la derecha, porque todavía no se ha escrito, y el modelo no ha aprendido a hacer otra cosa. Esta lección prohíbe esa mirada sin renunciar a calcular las TT posiciones a la vez, y de paso monta la subcapa que une las dos columnas del artículo.

Vuelve a la rejilla de la auto-atención y enciende la máscara. Es la misma frase, con las mismas proyecciones y la misma cuenta; lo único que cambia es que desaparece todo lo que hay a la derecha de la diagonal y que las casillas que quedan se reparten entre ellas la unidad que había que repartir. Empieza por la primera fila, que es el caso extremo: no tiene a nadie a quien mirar más que a sí misma.

La misma frase con la máscara puesta: se va media rejilla y ninguna fila deja de sumar 1. La posición 1 se queda con 1.00 sobre sí misma, porque detrás de ella no hay nada.

Tachar el futuro dentro del softmax

La máscara es una matriz del tamaño de la rejilla, y sus casillas sólo toman dos valores:

Mij={0si ji,si j>i.M_{ij} = \begin{cases} 0 & \text{si } j \leq i, \\ -\infty & \text{si } j > i \end{cases}.

Con ella, la llamada a la atención cambia en un sumando y en nada más:

Attention(Q,K,V)=softmax ⁣(QKdk+M)V,\text{Attention}(\mathbf{Q}, \mathbf{K}, \mathbf{V}) = \text{softmax}\!\left(\frac{\mathbf{Q}\mathbf{K}^{\top}}{\sqrt{d_k}} + \mathbf{M}\right)\mathbf{V},

con MRT×T\mathbf{M} \in \mathbb{R}^{T \times T} y el softmax por filas, como siempre. Nada de lo de dentro se toca: las tres proyecciones son las mismas, el divisor de la lección sobre el producto interno escalado es el mismo, y las cabezas de la lección sobre la atención multi-head son las mismas —cada una recibe la misma M\mathbf{M}—.

Por qué -\infty y no otra cosa se ve escribiendo el peso entero:

αij=exp(eij+Mij)m=1Texp(eim+Mim),\alpha_{ij} = \frac{\exp\left(e_{ij} + M_{ij}\right)}{\sum_{m=1}^{T}\exp\left(e_{im} + M_{im}\right)},

donde eije_{ij} es la puntuación ya dividida por dk\sqrt{d_k}. Como exp()=0\exp(-\infty) = 0, una casilla tachada desaparece del numerador y de su término del denominador, las dos cosas a la vez. Eso es lo que hace que no sobre nada: la unidad que el softmax reparte se la reparten enteramente las casillas que quedan.

Las dos alternativas que se le ocurren a cualquiera fallan justo ahí. Poner un 00 en la puntuación no tacha nada: exp(0)=1\exp(0) = 1 es un peso perfectamente normal, y en la celda de más abajo la fila 00 acaba dándole al futuro el 0.5250.525 de su peso. Y calcular el softmax como siempre y poner a cero los pesos después deja el denominador como estaba, así que la fila suma lo que sumara su principio —0.0850.085, 0.2910.291, 0.8710.871— y lo que sale ya no es una mezcla sino una mezcla encogida, con un factor distinto en cada fila.

Enmascarar es renormalizar el principio de la fila

Hay una manera de leer la máscara que la deja en su sitio, y sale de la fórmula de arriba en dos pasos. Para jij \leq i el sumando MijM_{ij} vale 00, de modo que el numerador es el de siempre y lo único que cambia es el denominador, que se queda con los términos hasta ii. Divide arriba y abajo por la suma completa m=1Texp(eim)\sum_{m=1}^{T}\exp(e_{im}) y aparecen los pesos sin máscara:

αijcausal=exp(eij)miexp(eim)=αijmiαim,ji.\alpha^{\text{causal}}_{ij} = \frac{\exp\left(e_{ij}\right)}{\sum_{m \leq i}\exp\left(e_{im}\right)} = \frac{\alpha_{ij}}{\sum_{m \leq i}\alpha_{im}}, \qquad j \leq i.

Es decir: coge la fila que la auto-atención calculaba, tira la parte de la derecha y divide lo que queda entre lo que sumaba. De ahí salen tres cosas, y ninguna necesita otra cuenta. El orden dentro de lo que sobrevive no cambia, porque todas esas casillas se dividen entre el mismo número: la máscara no altera preferencias, retira candidatos. La fila 11 se queda con un solo candidato y por tanto con α11=1\alpha_{11} = 1, la sepa el modelo o no. Y la máscara no es una segunda operación al lado del softmax, es el mismo softmax sobre menos casillas.

Una pasada en lugar de T

Falta decir qué entra en el decoder, porque hasta ahora hemos hablado de lo que no puede mirar y no de lo que tiene delante. Lo que entra es la salida correcta corrida una posición: en la fila ii va el token yi1y_{i-1}, el que el modelo debería haber escrito en el paso anterior, y en la fila 11 va <GO>, el símbolo de arranque, porque antes de la primera no hay ninguna. Así la fila ii tiene delante y<iy_{<i} y tiene que producir yiy_i, que es precisamente la distribución que el bloque anterior escribía

P(yiy<i,x1:Tx),P\left(y_i \mid y_{<i}, x_{1:T_x}\right),

una por fila. Eso es el teacher forcing de la lección sobre seq2seq, sin ningún cambio en la idea: se le da al modelo el token verdadero en lugar del que él escribió. Lo que cambia es el precio. Allí las TyT_y filas salían de TyT_y pasos de una recurrencia; aquí salen de un producto de matrices, y la máscara es lo que hace que esa barra vertical diga la verdad —sin ella, la fila ii estaría condicionando también con yiy_{i}, yi+1y_{i+1} y todo lo demás—.

Y ahora las dos concesiones, que van juntas. La máscara hace paralelo el entrenamiento, no la generación: al escribir de verdad no hay ninguna salida correcta que correr una posición, la fila ii necesita el token que produjo la pasada anterior, y son TyT_y pasadas, una por token. Lo otro es que el triángulo no ahorra ni una multiplicación: las T2T^{2} puntuaciones se calculan enteras y después se tachan más de la mitad. La máscara es una prohibición, no una optimización.

La subcapa que mira a la otra columna

El decoder lleva tres subcapas donde el encoder lleva dos, y la de en medio es la única caja de todo el artículo que junta las dos columnas. Es la atención del bloque anterior, entera, con las tres listas leídas de dos sitios distintos:

Q=XdecWQ,K=XencWK,V=XencWV,\mathbf{Q} = \mathbf{X}^{\text{dec}}\mathbf{W}^Q, \qquad \mathbf{K} = \mathbf{X}^{\text{enc}}\mathbf{W}^K, \qquad \mathbf{V} = \mathbf{X}^{\text{enc}}\mathbf{W}^V,

donde XdecRTy×dmodel\mathbf{X}^{\text{dec}} \in \mathbb{R}^{T_y \times d_{\text{model}}} es lo que se lleva escrito y XencRTx×dmodel\mathbf{X}^{\text{enc}} \in \mathbb{R}^{T_x \times d_{\text{model}}} es lo que el encoder dejó leído. Los superíndices de X\mathbf{X} dicen de qué columna viene la matriz; los de W\mathbf{W} siguen diciendo qué papel proyecta cada una, como desde la lección sobre la auto-atención. Ninguno de los dos cuenta capas.

La forma del mapa cae sola y es la que delata de dónde viene esta subcapa:

ARTy×Tx,\mathbf{A} \in \mathbb{R}^{T_y \times T_x},

rectangular, una fila por posición de la salida y una columna por posición de la entrada. Es la rejilla del bloque anterior, la de la lección sobre la atención multiplicativa, con los mismos dos ejes y la misma fila sumando 11. Lo que ha cambiado desde entonces son tres cosas y ninguna es la operación: la puntuación es el producto interno escalado en lugar del modelo de alineación, hay hh cabezas en paralelo en lugar de una, y la consulta ya no es el estado de una red recurrente sino una fila de una matriz que existe entera.

Aquí no hay máscara, y merece la pena decir por qué, porque suena raro después de la sección anterior. Lo que se prohíbe es leer lo que todavía no se ha escrito, y la entrada no se está escribiendo: se leyó entera antes de empezar. La posición TxT_x de la fuente está disponible desde el primer token de la salida igual que la posición 11. (Las implementaciones de verdad sí ponen otra máscara aquí, la del padding, para que las posiciones de relleno de un batch no cuenten; es el mismo mecanismo aplicado a otro problema y queda fuera de esta lección.)

Con esto, las tres llamadas a la atención del artículo se pueden poner en una tabla, y son las tres cajas que se encendían en el diagrama de la lección sobre el adiós a la recurrencia:

LlamadaConsultasClaves y valoresForma de A\mathbf{A}¿Máscara?
auto-atención del encoderXenc\mathbf{X}^{\text{enc}}Xenc\mathbf{X}^{\text{enc}}Tx×TxT_x \times T_xno
auto-atención enmascaradaXdec\mathbf{X}^{\text{dec}}Xdec\mathbf{X}^{\text{dec}}Ty×TyT_y \times T_ysí, causal
atención encoder-decoderXdec\mathbf{X}^{\text{dec}}Xenc\mathbf{X}^{\text{enc}}Ty×TxT_y \times T_xno

Una sola función, tres juegos de argumentos. Y las tres subcapas del decoder van envueltas en la línea de la lección anterior, sin ninguna excepción: se suma lo que entró y se normaliza el resultado, tres veces.

Las dos atenciones en NumPy

La primera celda es la máscara sola, sobre una rejilla de puntuaciones cuyos números he puesto yo a mano para que las cuentas se puedan seguir a ojo. Mira tres cosas: la fila 00, la columna de las sumas, y qué pasa con las dos maneras de equivocarse.

import numpy as np

E = np.array([[1.0, 2.0, 0.5, 3.0], # puntuaciones ya divididas por raiz de d_k
[0.5, 1.5, 2.5, 1.0],
[2.0, 0.0, 1.0, 0.5],
[1.0, 1.0, 2.0, 1.5]])
T = len(E)
futuro = np.triu(np.ones((T, T), dtype=bool), k=1) # True donde j > i


def softmax_filas(S):
Z = np.exp(S - S.max(axis=1, keepdims=True))
return Z / Z.sum(axis=1, keepdims=True)


M = np.where(futuro, -np.inf, 0.0)
con_mascara = softmax_filas(E + M) # lo del articulo
con_un_cero = softmax_filas(np.where(futuro, 0.0, E)) # tachar con un 0 en la PUNTUACION
pesos_a_cero = np.where(futuro, 0.0, softmax_filas(E)) # tachar los PESOS, ya repartidos

for nombre, A in (("con -inf ", con_mascara),
("con un 0 ", con_un_cero),
("pesos a 0 ", pesos_a_cero)):
print(nombre, "fila 0:", np.round(A[0], 3), " sumas:", np.round(A.sum(axis=1), 3))

# Enmascarar es quedarse con el principio de la fila y volver a normalizarlo.
libre = softmax_filas(E)
principios = np.where(futuro, 0.0, libre)
principios = principios / principios.sum(axis=1, keepdims=True)
print("\nla mascara es el principio renormalizado:", bool(np.allclose(con_mascara, principios)))
print("sin mascara, la fila 0 le da al futuro:", round(float(libre[0, 1:].sum()), 3))
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 el argumento entero. Con -\infty la fila 00 sale [1. 0. 0. 0.] y las cuatro sumas valen 11. Con un 00 en la puntuación las sumas también valen 11, y ahí está la trampa: la fila reparte bien, pero reparte entre gente que no debería estar —el 0.5250.525 que se van tres columnas del futuro—. Con los pesos tachados después, las sumas caen a 0.0850.085, 0.2910.291, 0.8710.871 y 11: cuatro filas escaladas de cuatro maneras distintas, que es lo peor de los tres resultados aunque sea el que parece más inofensivo.

La segunda celda monta las dos atenciones del decoder y comprueba lo único que la máscara promete.

import numpy as np

rng = np.random.default_rng(8)
T_x, T_y, d_model, d_k, d_v = 7, 5, 16, 8, 8

X_enc = rng.normal(size=(T_x, d_model)) # lo que el encoder dejo leido
X_dec = rng.normal(size=(T_y, d_model)) # lo que se lleva escrito
Wq = rng.normal(size=(d_model, d_k)) * 0.5
Wk = rng.normal(size=(d_model, d_k)) * 0.5
Wv = rng.normal(size=(d_model, d_v)) * 0.5


def atencion(consulta, contesta, causal=False):
Q, K, V = consulta @ Wq, contesta @ Wk, contesta @ Wv
E = Q @ K.T / np.sqrt(d_k)
if causal:
E = np.where(np.triu(np.ones(E.shape, dtype=bool), k=1), -np.inf, E)
Z = np.exp(E - E.max(axis=1, keepdims=True))
A = Z / Z.sum(axis=1, keepdims=True)
return A @ V, A


propia, A_propia = atencion(X_dec, X_dec, causal=True) # auto-atencion enmascarada
cruzada, A_cruzada = atencion(X_dec, X_enc) # atencion encoder-decoder

print("enmascarada: A", A_propia.shape, " sumas", np.round(A_propia.sum(axis=1), 3))
print("cruzada : A", A_cruzada.shape, " sumas", np.round(A_cruzada.sum(axis=1), 3))
print("casillas vivas de la enmascarada:", int((A_propia > 0).sum()), "de", A_propia.size)

# Reescribo el final de lo escrito y miro que se mueve de lo de antes.
otro = X_dec.copy()
otro[3:] = rng.normal(size=(T_y - 3, d_model))
salida_otro, _ = atencion(otro, otro, causal=True)
print("\ncambio las posiciones 4 y 5; se mueven las tres primeras?",
round(float(np.abs(salida_otro[:3] - propia[:3]).max()), 12))
print("y la cuarta?", round(float(np.abs(salida_otro[3] - propia[3]).max()), 3))
numpy

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

Los dos mapas tienen la forma de la tabla —5×55 \times 5 el propio, 5×75 \times 7 el cruzado— y las diez filas suman 11. De las 2525 casillas del primero quedan vivas 1515, que es la cuenta T(T+1)/2T(T+1)/2. Y las dos últimas líneas son la propiedad por la que existe todo esto: reescribir las posiciones 44 y 55 del decoder mueve lo que sale en las tres primeras exactamente 0.00.0, mientras que la cuarta —que es una de las que he reescrito— sí se mueve. Cambia causal=True por causal=False y ejecuta otra vez: ese 0.00.0 deja de serlo, y eso es un modelo que ha leído lo que todavía no ha escrito.

Comprueba tu intuición

Cinco preguntas: por qué -\infty y no un cero, cuántas casillas sobreviven, qué lleva máscara y qué no, qué se paraleliza de verdad, y qué sale en la primera fila.

Para tachar el futuro se le suma -\infty a la puntuación antes del softmax. ¿Qué tiene de malo dejar el softmax como estaba y poner a cero los pesos αij\alpha_{ij} que sobran, después?

Un decoder trabaja sobre una salida de T=6T = 6 tokens. ¿Cuántas casillas de su mapa A\mathbf{A} acaban siendo distintas de cero?

casillas

A margin of ±0 is accepted.

Un Transformer completo llama tres veces a la atención. Marca lo que es cierto de esas tres llamadas.

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

Con la máscara puesta, escribir una traducción de TyT_y tokens es una sola pasada por el decoder.

La rejilla de abajo tiene sus cuatro puntuaciones más altas repartidas por todas partes, y la fila 00 preferiría con mucho la columna 33. ¿Qué imprime?

import numpy as np
 
E = np.array([[2.0, 1.0, 0.5, 3.0],
              [0.0, 2.0, 1.0, 1.0],
              [1.0, 0.0, 2.0, 0.5],
              [0.5, 1.0, 0.0, 2.0]])
M = np.triu(np.full((4, 4), -np.inf), k=1)
Z = np.exp(E + M - (E + M).max(axis=1, keepdims=True))
A = Z / Z.sum(axis=1, keepdims=True)
print(np.round(A[0], 2))
 

Escribe atencion(Q, K, V, causal=False), la llamada que sirve para las tres atenciones del artículo. Q tiene forma (nq,dk)(n_q, d_k), K forma (nk,dk)(n_k, d_k) y V forma (nk,dv)(n_k, d_v), y nqn_q y nkn_k no tienen por qué coincidir: en la atención encoder-decoder las consultas son TyT_y y las claves son TxT_x.

Devuelve la pareja (salida, A), en ese orden: la salida de forma (nq,dv)(n_q, d_v) y el mapa de pesos de forma (nq,nk)(n_q, n_k), normalizado por filas. La puntuación es el producto escalar de cada consulta con cada clave dividido por dk\sqrt{d_k}, con dkd_k leído de la forma de Q.

Con causal=True, la casilla (i,j)(i, j) con j>ij > i vale -\infty antes del softmax. 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.


Ya no queda ninguna caja del dibujo sin una lección detrás, salvo la de arriba del todo, que convierte cada fila en una probabilidad por entrada del vocabulario. Lo que no ha aparecido nunca en estas siete lecciones son los números: cuántos bloques se apilan, cuánto miden, cuántas cabezas hay, qué tablas comparten pesos y cuántos parámetros suma todo junto. Cada pieza se ha construido con las anchuras que le hacían falta al argumento, y el artículo eligió unas concretas.

Ésa es la lección siguiente, sobre la arquitectura completa del artículo: la figura 1 caja por caja, con sus tamaños, puesta al lado de lo que llevamos construido. Es la lección que convierte Attention is All You Need en un texto que puedes abrir y leer de principio a fin, que era la promesa con la que empezó el bloque.

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.

  • Attention Is All You Need
    paperVaswani, Shazeer, Parmar y otros, 2017arXiv:1706.03762EN

    Su §3.2.3 es tu tabla de tres llamadas, con el enmascarado escrito como «−∞ en la entrada del softmax»; su §3.1, el decoder con la salida corrida una posición. Por qué −∞ y no un cero después, y el teacher forcing, son de la lección.