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 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 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.
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:
Con ella, la llamada a la atención cambia en un sumando y en nada más:
con 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 —.
Por qué y no otra cosa se ve escribiendo el peso entero:
donde es la puntuación ya dividida por . Como , 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 en la puntuación no tacha nada: es un peso perfectamente normal, y en la celda de más abajo la fila acaba dándole al futuro el 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 —, , — 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 el sumando vale , 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 . Divide arriba y abajo por la suma completa y aparecen los pesos sin máscara:
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 se queda con un solo candidato y por tanto con , 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 va el token , el que el modelo debería haber escrito en el paso anterior, y en la fila va <GO>, el símbolo de arranque, porque antes de la primera no hay ninguna. Así la fila tiene delante y tiene que producir , que es precisamente la distribución que el bloque anterior escribía
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 filas salían de 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 estaría condicionando también con , 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 necesita el token que produjo la pasada anterior, y son pasadas, una por token. Lo otro es que el triángulo no ahorra ni una multiplicación: las 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:
donde es lo que se lleva escrito y es lo que el encoder dejó leído. Los superíndices de dicen de qué columna viene la matriz; los de 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:
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 . 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 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 de la fuente está disponible desde el primer token de la salida igual que la posición . (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:
| Llamada | Consultas | Claves y valores | Forma de | ¿Máscara? |
|---|---|---|---|---|
| auto-atención del encoder | no | |||
| auto-atención enmascarada | sí, causal | |||
| atención encoder-decoder | no |
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 , la columna de las sumas, y qué pasa con las dos maneras de equivocarse.
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))
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 la fila sale
[1. 0. 0. 0.] y las cuatro sumas valen . Con un en la puntuación las sumas también valen
, y ahí está la trampa: la fila reparte bien, pero reparte entre gente que no debería estar —el
que se van tres columnas del futuro—. Con los pesos tachados después, las sumas caen a
, , y : 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.
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))
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 — el propio, el cruzado— y las
diez filas suman . De las casillas del primero quedan vivas , que es la cuenta
. Y las dos últimas líneas son la propiedad por la que existe todo esto: reescribir las
posiciones y del decoder mueve lo que sale en las tres primeras exactamente , 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 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é 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 a la puntuación antes del softmax. ¿Qué tiene de malo dejar el softmax como estaba y poner a cero los pesos que sobran, después?
Un decoder trabaja sobre una salida de tokens. ¿Cuántas casillas de su mapa acaban siendo distintas de cero?
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 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 preferiría con mucho la columna . ¿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 , K forma y V
forma , y y no tienen por qué coincidir: en la atención
encoder-decoder las consultas son y las claves son .
Devuelve la pareja (salida, A), en ese orden: la salida de forma y el
mapa de pesos de forma , normalizado por filas. La puntuación es el
producto escalar de cada consulta con cada clave dividido por , con
leído de la forma de Q.
Con causal=True, la casilla con vale 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
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.