Proyecto: un Transformer desde cero
30 min read
Ocho lecciones han desmontado Attention is All You Need pieza a pieza y en ninguna ha llegado a ejecutarse nada. La anterior, sobre los cinco números que fijan todas las formas del modelo, era la última que quedaba por poner: con ella el Transformer está entero sobre el papel y sigue sin existir como programa. La distancia entre las dos cosas es corta —unas sesenta líneas de NumPy— y conviene recorrerla una vez.
Lo que salga por arriba no significará nada, y vale más decirlo antes de ejecutar que después: los pesos son los que devuelve el generador de números aleatorios y nadie los ha corregido jamás. Con me gusta la casa roja a la entrada, este modelo propondrá lo primero que le salga, y ajustarlo para que deje de hacerlo queda fuera de esta lección. Aun así hay tres cosas comprobables con pesos cualesquiera, y son exactamente las que dicen si la arquitectura está bien montada: que las formas encajan de un extremo al otro, que ninguna posición del decoder ve lo que todavía no se ha escrito, y que el recuento de parámetros es el de la lección anterior con otros números dentro.
Montar quince cajas suena peor de lo que es, y la razón cabe en una frase: salvo en los dos extremos, todo lo que hay dentro recibe una matriz de filas y columnas y devuelve otra de la misma forma. La atención multi-head baja a por cabeza y la sube de vuelta; el perceptrón por posiciones sube a y la baja; el layer norm no toca ninguna forma. No es una casualidad afortunada: la conexión residual de la lección sobre los residuales y el layer norm lo exige, porque suma casilla a casilla lo que entró con lo que salió, y dos matrices de formas distintas no se suman.
De modo que el programa entero cambia de forma dos veces, y las dos están en los extremos: abajo, donde una lista de enteros se convierte en vectores, y arriba, donde cada vector se convierte en una probabilidad por entrada del vocabulario. Entre medias no hay nada que ajustar.
| Etapa | Forma |
|---|---|
| los tokens de la fuente | enteros |
| , tras la tabla y | |
| tras cada uno de los bloques del encoder | |
| los tokens que entran al decoder | enteros |
| , tras la tabla y | |
| tras cada uno de los bloques del decoder | |
| , tras la tabla leída al revés |
Todo lo de en medio conserva la forma
El modelo de esta lección es el del artículo con las anchuras divididas por dieciséis: bloques por columna, , cabezas —y por tanto —, , y un vocabulario de entradas compartido por los dos idiomas, como en el artículo. La frase es me gusta la casa roja, cinco tokens, y su traducción i like the red house, cinco más <EOS>, el símbolo de fin de secuencia: y .
Por abajo entra la caja de la lección sobre la codificación posicional, con el factor que la lección anterior le añadió:
la misma operación en las dos columnas, con la misma tabla . Apiladas por filas dan y . A partir de ahí sólo hay una línea, la de la lección sobre los residuales y el layer norm:
aplicada cinco veces con cinco cosas distintas en el hueco. Ésa es toda la arquitectura:
| Columna | Subcapa | Qué va en el hueco | Su mapa |
|---|---|---|---|
| encoder | 1 | , con las tres listas leídas de | |
| encoder | 2 | el perceptrón por posiciones | — |
| decoder | 1 | sobre , con sumada a las puntuaciones | |
| decoder | 2 | consultas de , claves y valores de | |
| decoder | 3 | el perceptrón por posiciones | — |
es como el artículo llama al perceptrón por posiciones: el mismo objeto con otro nombre. Y por arriba sale la tabla leída al revés, que es la caja que cerró la lección anterior:
con el softmax por filas. Ninguna de esas líneas cambia el número de columnas de lo que recibe, y de ahí sale la propiedad que hace corto el programa: no aparece en ninguna forma, sólo en cuántas veces da la vuelta un bucle. Poner seis bloques en lugar de dos no toca una sola matriz.
Por qué la máscara va en todos los bloques y no sólo en el primero
La lección sobre el encoder, el decoder y las máscaras dejó demostrado que una llamada enmascarada respeta el pasado: en su fila no entra ninguna columna a la derecha de . Lo que hace falta al montar el modelo es otra cosa, y es lo único que este ensamblaje añade a lo ya escrito: que la propiedad sobreviva a apilar. Entre lo que entra al decoder y lo que sale hay seis subcapas, tres por bloque, y sólo dos llevan máscara.
Llamemos respetar el pasado a esto: la fila de lo que sale depende únicamente de las filas a de lo que entró. Tres observaciones bastan.
La primera es que las cajas que trabajan fila a fila lo respetan sin hacer nada. El perceptrón por posiciones aplica los mismos pesos a cada fila por separado, el layer norm normaliza cada fila con su propia media y su propia desviación típica, y la suma residual suma la fila con la fila . Es la cuenta de la lección sobre el adiós a la recurrencia: de las quince cajas del dibujo, sólo las tres de atención miran a otra posición.
La segunda es que la auto-atención enmascarada lo respeta por construcción, que es de lo que trataba : su fila de pesos vale cero de la columna en adelante, así que lo que sale en mezcla valores de filas que no pasan de . Y la tercera es que la atención encoder-decoder lo respeta por un motivo distinto, y conviene separarlo: sus claves y sus valores no son del decoder, son del encoder. Mezcla las cinco posiciones de la fuente enteras —y debe hacerlo, la fuente está leída desde el principio—; de la salida sólo usa la consulta de su propia fila.
Respetar el pasado se conserva al componer, y con eso está hecho: con , la fila de depende sólo de las filas de , y cada una de ésas sólo de las filas de . Encadena las seis subcapas y sale la propiedad entera:
que es la barra vertical de diciendo la verdad.
La cadena tiene la fuerza de su eslabón más débil, y ahí está el motivo de que el artículo ponga máscara en la primera subcapa de cada uno de sus seis bloques. Quítasela al segundo bloque de este modelo y su auto-atención mezclará, en la fila , la fila de lo que le entregó el primero; todo lo que venga por encima hereda la mezcla, y la barra vertical pasa a ser mentira. La última celda lo mide con números.
El modelo entero en cuatro celdas
Cuatro celdas, en este orden: la frase y su entrada, las piezas, el modelo, y lo que se puede comprobar sin haber entrenado nada. Cada una necesita las anteriores, así que ejecútalas seguidas.
La primera fija el vocabulario compartido, convierte las dos frases en índices y construye lo que entra por abajo en las dos columnas. Mira la tercera línea que imprime: lo que entra al decoder no es la traducción, es la traducción corrida una posición, con <GO> —el símbolo de arranque— en la primera y sin el <EOS> del final.
V = ["<GO>", "<EOS>", "me", "gusta", "la", "casa", "roja",
"i", "like", "the", "red", "house"] # un vocabulario para los dos idiomas
ix = {w: i for i, w in enumerate(V)}
n_v, d_model, h, d_ff, N = len(V), 32, 2, 128, 2
x = [ix[w] for w in "me gusta la casa roja".split()]
y = [ix[w] for w in "i like the red house".split()] + [ix["<EOS>"]]
ent_dec = [ix["<GO>"]] + y[:-1] # la salida, corrida una posicion
rng = np.random.default_rng(0)
E = rng.normal(size=(n_v, d_model)) / np.sqrt(d_model) # la tabla, la misma en los tres sitios
def codificacion_posicional(T, d):
i = np.arange(d // 2)
omega = 1.0 / (10000 ** (2 * i / d))
PE = np.zeros((T, d))
PE[:, 0::2] = np.sin(np.arange(T)[:, None] * omega)
PE[:, 1::2] = np.cos(np.arange(T)[:, None] * omega)
return PE
def entrada(tokens): # la caja de abajo del todo, las dos veces
return np.sqrt(d_model) * E[tokens] + codificacion_posicional(len(tokens), d_model)
X_enc, X_dec = entrada(x), entrada(ent_dec)
print("fuente :", [V[i] for i in x], "->", X_enc.shape)
print("etiqueta:", [V[i] for i in y])
print("al decoder:", [V[i] for i in ent_dec], "->", X_dec.shape)
La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador.
La segunda no trae nada nuevo, y ése es el hallazgo: son seis funciones y las seis son el cuerpo de
un desafío que ya resolviste. softmax_filas viene de la lección sobre la auto-atención;
atencion es el de la lección sobre las máscaras, con su máscara opcional; multi_head reparte las columnas entre
las cabezas como el de la lección sobre la atención multi-head; layer_norm y subcapa son los de la lección sobre los residuales y el layer norm. Lo único
escrito para hoy es ffn, que son dos líneas.
# Necesita la celda anterior: d_model, h, d_ff.
d_k = d_v = d_model // h # no se eligen: salen de h
def softmax_filas(S):
Z = np.exp(S - S.max(axis=1, keepdims=True))
return Z / Z.sum(axis=1, keepdims=True)
def atencion(Q, K, Val, causal=False):
S = Q @ K.T / np.sqrt(Q.shape[1])
if causal:
S = np.where(np.triu(np.ones(S.shape, dtype=bool), k=1), -np.inf, S)
A = softmax_filas(S)
return A @ Val, A
def multi_head(Xq, Xkv, p, causal=False): # Xq trae las consultas; Xkv, claves y valores
salidas, mapas = [], []
for i in range(h):
col = slice(i * d_k, (i + 1) * d_k) # el bloque de columnas de la cabeza i
s, A = atencion(Xq @ p["Wq"][:, col], Xkv @ p["Wk"][:, col],
Xkv @ p["Wv"][:, col], causal)
salidas.append(s)
mapas.append(A)
return np.concatenate(salidas, axis=1) @ p["Wo"], np.stack(mapas)
def layer_norm(X, gamma, beta, eps=1e-5):
mu, s = X.mean(axis=1, keepdims=True), X.std(axis=1, keepdims=True)
return gamma * (X - mu) / np.sqrt(s**2 + eps) + beta
def subcapa(X, salida, norma): # sumar lo que entro y normalizar la suma
return layer_norm(X + salida, *norma)
def ffn(X, p): # el perceptron por posiciones
return np.maximum(0.0, X @ p["W1"] + p["b1"]) @ p["W2"] + p["b2"]
print("definidas:", [f.__name__ for f in (softmax_filas, atencion, multi_head,
layer_norm, subcapa, ffn)])
La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador.
La tercera saca los pesos del generador y encadena las subcapas en el orden de la tabla de arriba.
mascaras es un interruptor por bloque y viene puesto en los dos; la última celda lo usará. De lo
que imprime, mira la forma de la salida, que las seis filas suman , y la pérdida por posición al
lado de .
# Necesita las dos celdas anteriores.
def pesos_atencion(r): # las cuatro matrices, sin sesgos
return {n: r.normal(size=(d_model, d_model)) / np.sqrt(d_model)
for n in ("Wq", "Wk", "Wv", "Wo")}
def pesos_bloque(r, cruzada):
p = {"propia": pesos_atencion(r),
"W1": r.normal(size=(d_model, d_ff)) / np.sqrt(d_model), "b1": np.zeros(d_ff),
"W2": r.normal(size=(d_ff, d_model)) / np.sqrt(d_ff), "b2": np.zeros(d_model),
"normas": [(np.ones(d_model), np.zeros(d_model)) for _ in range(3 if cruzada else 2)]}
if cruzada:
p["cruzada"] = pesos_atencion(r) # la subcapa que junta las dos columnas
return p
encoder = [pesos_bloque(rng, False) for _ in range(N)]
decoder = [pesos_bloque(rng, True) for _ in range(N)]
def modelo(X_enc, X_dec, mascaras=(True, True)):
mapas = {}
for n, p in enumerate(encoder):
s, A = multi_head(X_enc, X_enc, p["propia"])
X_enc = subcapa(X_enc, s, p["normas"][0])
X_enc = subcapa(X_enc, ffn(X_enc, p), p["normas"][1])
mapas["enc %d" % n] = A
for n, p in enumerate(decoder):
s, A = multi_head(X_dec, X_dec, p["propia"], causal=mascaras[n])
X_dec = subcapa(X_dec, s, p["normas"][0])
s, C = multi_head(X_dec, X_enc, p["cruzada"])
X_dec = subcapa(X_dec, s, p["normas"][1])
X_dec = subcapa(X_dec, ffn(X_dec, p), p["normas"][2])
mapas["dec %d" % n], mapas["cruz %d" % n] = A, C
return softmax_filas(X_dec @ E.T), mapas # la tabla, leida al reves
P, mapas = modelo(X_enc, X_dec)
print("salida:", P.shape, " una fila por posicion escrita, una columna por entrada")
print("cada fila suma 1:", bool(np.allclose(P.sum(axis=1), 1.0)))
print("perdida por posicion:", round(float(-np.log(P[np.arange(len(y)), y]).mean()), 3),
" ln |V| =", round(float(np.log(n_v)), 3))
La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador.
La pérdida sale y . Un modelo recién inicializado no está en el reparto uniforme: está peor. Repartir por igual es lo mejor que puede hacer quien no sabe nada, y éste ni siquiera reparte por igual —reparte desigual y sin motivo, y paga la diferencia—. Entrenar consistiría, para empezar, en bajar hasta esa línea; cruzarla es lo que se parecería a traducir.
La cuarta comprueba lo comprobable: los seis mapas con su forma, la primera fila de la salida, y la tabla que mide la sección anterior —el mismo modelo con la máscara puesta en los dos bloques, en uno, en el otro y en ninguno, mirando cuánto se mueven las cuatro primeras filas de la salida al reescribir el final de lo que entra—.
# Necesita las tres celdas anteriores.
for nombre, A in mapas.items():
print("%-7s %-11s filas que suman 1: %s" % (nombre, A.shape,
bool(np.allclose(A.sum(axis=2), 1.0))))
print("\nfila 1 de la salida:", np.round(P[0], 3))
print("lo mas probable en cada posicion:", [V[i] for i in P.argmax(axis=1)])
otro = list(ent_dec)
otro[4:] = [ix["casa"], ix["me"]] # reescribo las posiciones 5 y 6
X_otro = entrada(otro)
print("\nmascara en... se mueven las cuatro primeras filas")
for etiqueta, m in (("los dos bloques", (True, True)), ("solo el primero", (True, False)),
("solo el segundo", (False, True)), ("ninguno ", (False, False))):
antes, _ = modelo(X_enc, X_dec, m)
ahora, _ = modelo(X_enc, X_otro, m)
print(" %s %s" % (etiqueta, round(float(np.abs(ahora[:4] - antes[:4]).max()), 4)))
def arrays(p): # todos los pesos de un bloque, uno a uno
for v in p.values():
if isinstance(v, dict):
yield from v.values()
elif isinstance(v, list):
yield from (a for par in v for a in par)
else:
yield v
total = E.size + sum(a.size for p in encoder + decoder for a in arrays(p))
print("\nparametros:", total, " frente a los 63045632 del modelo base")
La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador.
Los mapas cuadran con la tabla: dos por capa de atención —una por cabeza—, los del encoder, los enmascarados y los cruzados, el único rectangular del modelo, y lo es porque las dos frases no miden lo mismo. En las seis posiciones lo más probable es la misma entrada, the, con un en la primera: no es que sepa inglés, es la fila de que más apunta en la dirección que le sale al decoder, decidida por el mismo generador que puso los pesos.
La tabla de las máscaras es la afirmación de la sección anterior con números. Sólo la primera fila da , y da exacto; enmascarar un bloque y no el otro deja o según cuál se quite, y no enmascarar ninguno, . Un número pequeño no vale: o la fila anterior no se mueve, o el modelo ha leído lo que todavía no había escrito.
Y los parámetros son la fórmula de la lección anterior con otros cinco números dentro — por bloque de encoder, por bloque de decoder y la tabla compartida—:
El modelo base del artículo tiene mil sesenta y ocho veces más pesos que éste y ni una caja más.
Comprueba tu intuición
Tres preguntas: qué demuestra un forward pass con pesos sin entrenar, si la máscara del primer bloque le alcanza al decoder entero, y cuántos parámetros guarda una capa de atención.
El modelo corre entero con pesos que nadie ha entrenado y devuelve su matriz de probabilidades. Marca lo que ese resultado demuestra.
Select every correct option. This is graded all-or-nothing: there is no partial credit.
El decoder de esta lección apila dos bloques. Con enmascarar la auto-atención del primero, el modelo entero ya es autorregresivo: lo que sale en la posición no puede depender de lo que entra en la .
El modelo de juguete tiene y cabezas. ¿Cuántos parámetros guarda una de sus capas de atención multi-head, con sus cuatro matrices y sin sesgos?
A margin of ±0 is accepted.
Queda una cosa que este modelo no ha hecho ni puede hacer de una pasada. Las seis filas salen a la
vez porque la salida correcta está sobre la mesa —el teacher forcing que su lección puso al lado de la máscara—, y al escribir de verdad no la hay: el decoder arranca con
<GO> solo, elige un token, lo realimenta y vuelve a llamarse. Ese bucle es el desafío, escrito
contra una función paso que hace de modelo.
Una llamada al modelo entrega las distribuciones a la vez porque tiene delante la salida correcta. Generar no funciona así: cada token se elige con lo que hay escrito hasta ese momento y vuelve a entrar.
Escribe genera(paso, go, eos, max_pasos). paso(entrada) es el modelo hecho función: recibe
la lista de tokens que entra al decoder —la primera vez, [go]— y devuelve las
probabilidades de la posición siguiente, un array de forma . No la
escribes tú, te la dan.
En cada vuelta elige la entrada más probable —decodificación voraz, la de la lección
sobre el encoder y el decoder—, apúntala y realiméntala. Para al elegir
eos, que sí va en el resultado, o al llevar max_pasos tokens, lo que ocurra antes.
Devuelve la lista de los tokens elegidos, sin go.
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.
Lo que has montado es el modelo del artículo, y eso incluye una decisión que en 2017 no se discutía: dos columnas, porque la tarea era traducir. La fuente entera por un lado, lo escrito por el otro, y una subcapa que las junta. Nada de lo que has escrito obliga a esa forma —las dos columnas comparten la tubería, las piezas y hasta la tabla—, y en cuanto la tarea deja de ser traducir, la pregunta de cuántas columnas hacen falta se abre.
Las dos respuestas que se dieron son la lección siguiente, sobre BERT y GPT: una se quedó con la columna de la izquierda y con ninguna máscara, la otra con la de la derecha y con todas, y las dos dejaron de traducir para entrenarse contra textos que no tienen traducción al lado. La aritmética que las mueve es la que acabas de escribir, sin una operación nueva; lo que cambia es qué se les pide predecir.
Further reading1 source · 1 article
Where this lesson comes from, and where to go next. None of it is needed to carry on with the course.
- The Annotated Transformer
El mismo montaje que haces aquí, en PyTorch y línea a línea junto al texto del artículo: las mismas piezas y, encima, el bucle de entrenamiento que tu celda no puede correr. Entrena una traducción de juguete en una GPU.