Quince cajas, cinco números

Quince cajas, cinco números

24 min read

Descrito, el Transformer ya está entero. La lección anterior, sobre el encoder, el decoder y las máscaras, cerró la última operación que quedaba suelta, y con ella el modelo sabe leer una frase, escribir otra y no mirar lo que todavía no ha escrito. Construido no está. Lo que hay escrito son letras —dkd_k, dffd_{\text{ff}}, hh, NN—, y un programa no reserva memoria para una letra: hasta que cada una no valga un número, no hay ni una matriz que crear.

Los números del artículo son cinco, y ponerlos encima del dibujo es lo que cierra el bloque. Cinco y no quince: fijados ésos, todas las formas del modelo quedan determinadas —incluidas un par que parecen libres y no lo son—, y de ellos sale el recuento entero de parámetros y, con él, el reparto, que no está donde el dibujo sugiere. Antes hay que cerrar la caja que ninguna lección construyó, la que corona la columna del decoder, y su respuesta viene de mucho más atrás que el resto del bloque: es la tabla de embeddings de la lección sobre las representaciones densas, leída al revés.

Pincha cualquiera de las quince cajas: el panel dice qué hace y en qué lección se construyó. Sólo una nombra a ésta, y es la de arriba del todo.

La figura del artículo está en inglés y este bloque la ha ido construyendo en español, así que lo primero es la correspondencia, rótulo por rótulo:

Rótulo en la figuraQué caja esDónde se construyó
Input Embedding, Output Embeddingla tabla E\mathbf{E}, la misma en las dosla lección sobre las representaciones densas
Outputs (shifted right)lo ya escrito, corrido una posiciónla lección anterior, sobre el decoder y su máscara
Positional EncodingPE\text{PE}, sumada una sola vez abajo del todola lección sobre la codificación posicional
Multi-Head Attention (izquierda)la auto-atención del encoderauto-atención, producto interno escalado y cabezas
Masked Multi-Head Attentionla misma, con M\mathbf{M} sumada a las puntuacionesla lección anterior, sobre las máscaras
Multi-Head Attention (centro de la derecha)la atención encoder-decoderla lección anterior, sobre la subcapa que junta las dos columnas
Add & Norm, cinco vecesLayerNorm(X+Sublayer(X))\text{LayerNorm}\left(\mathbf{X} + \text{Sublayer}(\mathbf{X})\right)la lección sobre los residuales y el layer norm
Feed Forward, dos vecesel perceptrón por posicionesla lección sobre los residuales y el layer norm
NxNN bloques apiladosesta lección
Linear, Softmaxla proyección al vocabularioesta lección

La caja de arriba lee la tabla del bloque 1 al revés

De lo que sale del último bloque del decoder, la fila tt es un vector de dmodeld_{\text{model}} coordenadas, xtdec\mathbf{x}^{\text{dec}}_t, y lo que hace falta en esa posición es una probabilidad por cada entrada del vocabulario. El artículo dibuja dos cajas para ese salto: una matriz que lleva de dmodeld_{\text{model}} a V\lvert V \rvert coordenadas y el reparto de siempre. Esa matriz podría ser suya, y no lo es:

y^t=softmax(Extdec)RV,\hat{\mathbf{y}}_t = \text{softmax}\left(\mathbf{E}\,\mathbf{x}^{\text{dec}}_t\right) \in \mathbb{R}^{\lvert V \rvert},

donde ERV×dmodel\mathbf{E} \in \mathbb{R}^{\lvert V \rvert \times d_{\text{model}}} es la tabla de embeddings del bloque 1, la misma que abajo del todo cambia cada token por su vector. La coordenada que le toca a la entrada ww dentro del softmax es ewxtdec\mathbf{e}_w^{\top}\mathbf{x}^{\text{dec}}_t: el producto escalar entre la fila de esa entrada y lo que el decoder lleva calculado. La entrada más probable es la que más apunta en la dirección que la posición pide.

Es la tabla leída en la otra dirección. Abajo se entra con una entrada y se sale con su vector; aquí se entra con un vector y se sale con una puntuación por entrada. Las TyT_y posiciones salen a la vez del mismo producto, como todo en este bloque:

XdecERTy×V,\mathbf{X}^{\text{dec}}\mathbf{E}^{\top} \in \mathbb{R}^{T_y \times \lvert V \rvert},

con el softmax por filas. Una fila por posición escrita, una columna por entrada del vocabulario, y ni un parámetro nuevo en toda la caja.

La misma tabla en tres sitios

El artículo va más lejos que eso y usa esa matriz tres veces: en el embedding de entrada, en el de salida y en la proyección al vocabulario. Eso exige un vocabulario común a los dos idiomas, y lo tiene: sus 3700037\,000 entradas salen de un BPE (byte-pair encoding) entrenado sobre los dos textos a la vez, la técnica de la lección sobre la tokenización. Una tabla de 37000×51237\,000 \times 512 son 1894400018\,944\,000 parámetros, y lo que se compra compartiéndola es no pagarla tres veces.

Queda un factor que la lección sobre la codificación posicional dejó fuera. Allí la entrada de la primera capa se escribió xt=ewt+PEt\mathbf{x}_t = \mathbf{e}_{w_t} + \text{PE}_t; el artículo multiplica el primer sumando:

xt=dmodel  ewt+PEt.\mathbf{x}_t = \sqrt{d_{\text{model}}}\;\mathbf{e}_{w_t} + \text{PE}_t.

Con los números delante se ve para qué. Todas las filas de PE\text{PE} miden lo mismo, 256=16\sqrt{256} = 16 —es la cuenta PEposPEpos=dmodel/2\text{PE}_{\text{pos}}^{\top}\text{PE}_{\text{pos}} = d_{\text{model}}/2 de aquella lección—, mientras que una fila de E\mathbf{E} recién inicializada, con coordenadas pequeñas alrededor de cero, mide alrededor de 11. Sumados así, el contenido queda debajo de la posición; multiplicar por 51222.6\sqrt{512} \approx 22.6 los pone en el mismo orden de magnitud. El artículo escribe la multiplicación y no la justifica: ésta es la razón que se le suele dar, no una que él dé.

Cinco números fijan todas las formas

La tabla 3 del artículo tiene una fila por cada modelo que entrenaron. La primera, la que llaman base, es la que ha ido apareciendo suelta a lo largo del bloque:

NúmeroValorQué fija
NN66cuántos bloques se apilan en cada columna
dmodeld_{\text{model}}512512la anchura que entra y sale de toda subcapa
hh88cuántas cabezas tiene cada capa de atención
dffd_{\text{ff}}20482\,048la anchura interior del perceptrón por posiciones
V\lvert V \rvert3700037\,000cuántas filas tiene E\mathbf{E}

Todo lo demás sale de esos cinco. El reparto de la lección sobre la atención multi-head da dk=dv=dmodel/h=64d_k = d_v = d_{\text{model}}/h = 64, y con él WiQ\mathbf{W}^Q_i, WiK\mathbf{W}^K_i y WiV\mathbf{W}^V_i miden 512×64512 \times 64 y WO\mathbf{W}^O mide (hdv)×dmodel=512×512(h \cdot d_v) \times d_{\text{model}} = 512 \times 512; el perceptrón por posiciones sube a 20482\,048 y vuelve a bajar; cada normalización aporta dos vectores de 512512. Lo que no está en la lista tampoco está en ninguna forma: ni TT, ni el tamaño del batch. Esa tabla tiene además dos columnas que esta lección no toca, dropout y label smoothing: son del entrenamiento y no de la arquitectura, y quedan fuera del curso.

Un tercio del modelo es una tabla de búsqueda

Un bloque del encoder pesa 31503363\,150\,336 parámetros —la cuenta de la lección sobre los residuales y el layer norm, con sus sesgos y sus vectores de normalización dentro—. Uno del decoder lleva exactamente una atención y una normalización más:

3150336+1048576+1024=4199936.3\,150\,336 + 1\,048\,576 + 1\,024 = 4\,199\,936.

El modelo entero son seis de cada, más la tabla que las dos columnas comparten:

63150336+64199936+37000512=63045632.6 \cdot 3\,150\,336 + 6 \cdot 4\,199\,936 + 37\,000 \cdot 512 = 63\,045\,632.

Ahí está el reparto que el dibujo esconde. Son 18.918.9 millones de encoder, 25.225.2 de decoder y 18.918.9 de tabla: un 3030, un 4040 y un 3030 por ciento. Lo que más pesa de una sola pieza no es ninguna de las quince cajas, es la tabla de búsqueda, que ella sola supera a las seis capas del encoder juntas —1894400018\,944\,000 frente a 1890201618\,902\,016—. Y de ahí que el artículo la comparta: las dos copias que se ahorra suman 3788800037\,888\,000 parámetros, casi tanto como las doce capas.

Esa cuenta es una reconstrucción a partir de la figura, no la contabilidad del artículo, y conviene decir en qué se separan. El artículo declara 6565 millones y aquí salen 63.063.0. Los dos que faltan caben enteros en dos detalles que la figura no muestra: el tamaño exacto del vocabulario —con 4100041\,000 entradas la cifra cuadraría— y dónde pone cada implementación sus sesgos. La misma receta aplicada a la otra fila de la tabla 3, la que llaman big (dmodel=1024d_{\text{model}} = 1\,024, h=16h = 16, dff=4096d_{\text{ff}} = 4\,096), da 214.2214.2 millones frente a los 213213 declarados.

El recuento en una función

Toda la sección anterior cabe en ocho líneas de aritmética con números enteros. Mira tres cosas: los tres porcentajes, la diferencia entre la tabla y el encoder entero, y qué le pasa al total si cada caja tuviera su propia tabla.

def parametros(N, d_model, h, d_ff, V):
d_k = d_v = d_model // h # no se eligen: salen de h
atencion = 3 * h * d_model * d_k + h * d_v * d_model # las tres por cabeza, y W^O
ffn = 2 * d_model * d_ff + d_ff + d_model # dos capas y sus dos sesgos
norm = 2 * d_model # gamma y beta
enc = N * (atencion + ffn + 2 * norm) # dos subcapas, dos normalizaciones
dec = N * (2 * atencion + ffn + 3 * norm) # tres y tres
return enc, dec, V * d_model # la tabla, una para tres cajas


enc, dec, tabla = parametros(N=6, d_model=512, h=8, d_ff=2048, V=37000)
total = enc + dec + tabla

for nombre, valor in (("encoder", enc), ("decoder", dec), ("tabla E", tabla)):
print(f"{nombre:>8} {valor:>12,} {100 * valor / total:5.1f} %")
print(f"{'total':>8} {total:>12,}")

print("\nla tabla menos el encoder entero:", tabla - enc)
print("sin compartirla, dos copias mas :", total + 2 * tabla)

grande = parametros(N=6, d_model=1024, h=16, d_ff=4096, V=37000)
print("\nla fila big de la tabla 3 :", sum(grande))

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

Los tres porcentajes salen 30.030.0, 40.040.0 y 30.030.0, y el decoder pesa más que el encoder por lo único en lo que se diferencian: una subcapa de atención más, seis veces. La tabla le saca al encoder 4198441\,984 parámetros, que sobre dieciocho millones es empatar. Y sin compartirla el modelo se iría a 100933632100\,933\,632, de 6363 millones a 101101 por una decisión que no cambia ni una operación del dibujo.

La última línea cambia los cinco argumentos por los de la fila big y da 214171648214\,171\,648. Lo que no vas a poder cambiar es la longitud del texto, porque TT no aparece en ninguna línea de la función: no es que se haya quedado fuera de la cuenta, es que no hay ninguna matriz cuya forma dependa de ella.

Comprueba tu intuición

Cinco preguntas: qué forma no fijan los cinco números, qué costaría no compartir la tabla, qué vas a encontrarte al abrir la figura, si un texto más largo pide más parámetros y qué sale por arriba del decoder.

Con los cinco números del artículo puestos —NN, dmodeld_{\text{model}}, hh, dffd_{\text{ff}} y V\lvert V \rvert—, ¿cuál de estas formas sigue sin quedar determinada?

El artículo usa la misma matriz E\mathbf{E} en los tres sitios que la necesitan. Si cada uno tuviera la suya, ¿cuántos parámetros más tendría el modelo base?

parámetros

A margin of ±0 is accepted.

Abres la figura 1 del artículo. Marca lo que vas a encontrarte.

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

Pasas de traducir frases de 5050 tokens a traducir textos de 10001\,000. El modelo base necesita más parámetros.

Ésta es la caja de arriba del todo, con la tabla compartida y cuatro posiciones escritas. ¿Qué imprime?

import numpy as np
 
# E es la tabla compartida; X_dec sale del ultimo bloque
V, d_model, T_y = 37000, 512, 4
rng = np.random.default_rng(0)
E = rng.normal(size=(V, d_model))
X_dec = rng.normal(size=(T_y, d_model))
 
O = X_dec @ E.T
P = np.exp(O - O.max(axis=1, keepdims=True))
P /= P.sum(axis=1, keepdims=True)
print(P.shape, np.allclose(P.sum(axis=1), 1.0))
 

Ya no queda nada que describir. Las quince cajas tienen su operación, sus formas y sus números, y eso es justo lo que hace falta para escribir el modelo entero y ver qué contesta. Lo que no cabe es el tamaño: sesenta y tres millones de parámetros no se mueven dentro de un navegador, así que lo que se construye es la arquitectura con anchuras de juguete —las mismas fórmulas, dos capas en lugar de seis, dos cabezas en lugar de ocho— y una frase corta que se pueda seguir a ojo.

Eso es la lección siguiente, el proyecto: un Transformer desde cero. Casi todas las piezas están escritas ya, cada una en el desafío de su lección —el softmax por filas, la atención con su máscara, el layer norm, el perceptrón por posiciones, la codificación posicional—, y lo que queda es ensamblarlas y hacer una pasada hacia delante de principio a fin: de una frase de tokens a una probabilidad por entrada del vocabulario, con los mapas de atención a la vista por el camino. Entrenar de verdad pide una máquina que un navegador no tiene, y eso tiene su sitio más adelante en el bloque.

Further reading2 sources · 2 papers

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.4 comparte la tabla E entre las dos entradas y la salida y la multiplica por √d_model; su tabla 3, fila «base», da tus cinco números. Los 65 millones los reconstruyes tú, y no cuadran al detalle con lo que declara.

  • Using the Output Embedding to Improve Language Models
    paperPress y Wolf, 2017arXiv:1608.05859EN

    El Transformer lo cita para leer E al revés en la salida. Aquí está el porqué: atar la matriz de entrada con la de salida baja la perplejidad y reduce el modelo a menos de la mitad. Es sobre LSTM; la cuenta de parámetros es tuya.