Apilar sin perder lo de abajo

Apilar sin perder lo de abajo

28 min read

A la atención ya no le falta nada: tiene sus tres proyecciones, su divisor, sus cabezas y, desde la lección anterior sobre la codificación posicional, unas filas que saben dónde están. Lo que no existe todavía es una capa que se pueda poner encima de otra. Y el artículo pone seis, con lo que sale de cada una haciendo de entrada de la siguiente.

Repetir seis veces una operación que ya funciona suena a no hacer nada, y es justo donde se estrelló el bloque 3. Allí lo que se repetía era un paso de la recurrencia, y componer una función consigo misma multiplica derivadas: un producto de muchos factores parecidos se va a cero o se dispara, y lo que venía de abajo no llega arriba. La profundidad de un Transformer no la pone la frase —la ponen las capas, y son seis se alargue el texto lo que se alargue—, pero seis productos siguen siendo seis productos. Esta lección es lo que hay que ponerle alrededor a una subcapa para que esa cuenta salga distinta, y son dos cosas.

Vuelve al diagrama de la lección sobre el adiós a la recurrencia y mira las cajas que dejó sin explicar. Cinco de las quince se llaman «Suma y layer norm», y de cada una sale una línea que baja por el lado de fuera de su columna y vuelve a entrar por debajo de la subcapa que esa caja cierra. Pincha una: esa línea es la mitad de lo que la caja suma.

Cinco cajas de «Suma y layer norm» y cinco líneas que rodean una subcapa por fuera, una por cada una. Enciende el marcador y no se ilumina ninguna de las cinco: las que mezclan posiciones siguen siendo las tres de atención.

Ninguna subcapa reemplaza lo que recibe

La conexión residual es la línea del dibujo escrita en símbolos, y no tiene más que esto:

X+Sublayer(X),\mathbf{X} + \text{Sublayer}(\mathbf{X}),

donde Sublayer\text{Sublayer} es un hueco —como el φ\varphi del bloque 2 o el aa del bloque anterior— que ocupa la atención multi-head en la primera subcapa y el perceptrón por posiciones en la segunda. Lo interesante está en su derivada. Pon las TT filas de X\mathbf{X} una debajo de otra, como en la lección sobre por qué falla el MLP (multilayer perceptron) en secuencias, y llama xRTdmodel\mathbf{x} \in \mathbb{R}^{T \cdot d_{\text{model}}} a esa columna. Derivar una suma es derivar cada sumando:

(x+Sublayer(x))x=I+Sublayer(x)x.\frac{\partial\left(\mathbf{x} + \text{Sublayer}(\mathbf{x})\right)}{\partial\mathbf{x}} = \mathbf{I} + \frac{\partial\,\text{Sublayer}(\mathbf{x})}{\partial\mathbf{x}}.

Ahí está todo. El jacobiano de la subcapa depende de sus pesos y puede valer lo que sea, incluso cero; el otro sumando es la identidad y no depende de nada. Y como la regla de la cadena de su lección convierte una composición en un producto de jacobianos, apilar NN bloques multiplica NN paréntesis así, cada uno con su identidad dentro.

Dos bloques, multiplicados

Con dos bloques —y subíndices para distinguir la subcapa de cada uno— el producto se abre así:

(I+Sublayer2x)(I+Sublayer1x)=I+Sublayer1x+Sublayer2x+Sublayer2xSublayer1x.\left(\mathbf{I} + \frac{\partial\,\text{Sublayer}_2}{\partial\mathbf{x}}\right)\left(\mathbf{I} + \frac{\partial\,\text{Sublayer}_1}{\partial\mathbf{x}}\right) = \mathbf{I} + \frac{\partial\,\text{Sublayer}_1}{\partial\mathbf{x}} + \frac{\partial\,\text{Sublayer}_2}{\partial\mathbf{x}} + \frac{\partial\,\text{Sublayer}_2}{\partial\mathbf{x}}\frac{\partial\,\text{Sublayer}_1}{\partial\mathbf{x}}.

Cuatro sumandos: uno que no atraviesa ninguna subcapa, dos que atraviesan una y uno que atraviesa las dos. Con NN bloques salen 2N2^{N}, uno por cada manera de elegir en cada paréntesis entre la identidad y la subcapa, y entre ellos hay siempre exactamente uno que no atraviesa nada.

En la lección sobre el gradiente que se desvanece no había paréntesis: cada factor era un producto pelado y el único camino entre el paso 11 y el paso TT los atravesaba todos. La LSTM (long short-term memory) respondió abriendo una vía aditiva al lado. La conexión residual es esa misma idea con la profundidad en el sitio del tiempo: aquí lo que se apila no son posiciones, son capas.

Con eso queda pagada la deuda de la lección anterior. La codificación posicional se suma una vez, abajo del todo, y ninguna caja vuelve a recordar después dónde estaba cada token: el sumando que vale I\mathbf{I} es lo que la lleva hasta arriba.

Una sola anchura de abajo arriba

La suma trae una restricción escondida: X\mathbf{X} y Sublayer(X)\text{Sublayer}(\mathbf{X}) se suman casilla a casilla, así que la subcapa tiene prohibido devolver otra forma que no sea T×dmodelT \times d_{\text{model}}. Ninguna de las dos la respeta por dentro y las dos hacen el viaje de vuelta: la atención multi-head baja a dvd_v por cabeza y WO\mathbf{W}^O la devuelve —para esto lo dejó señalado la lección sobre la atención multi-head—, y el perceptrón por posiciones sube a dffd_{\text{ff}} y vuelve a bajar.

Eso convierte a dmodeld_{\text{model}} en el único número que hay que elegir para todo el modelo. En el perceptrón multicapa del bloque 2 cada capa tenía su anchura dld_l y podía estrecharse a voluntad; aquí, si una capa cambiara de anchura, la siguiente no tendría a qué sumarle nada. Los 512512 del artículo son los mismos en la primera capa y en la sexta, y esa rigidez es el precio de apilar.

Cada fila se normaliza con sus propias coordenadas

La suma crea un problema ella sola: si cada subcapa añade sin quitar, lo que circula va creciendo capa a capa, y la lección sobre el producto interno escalado ya midió lo que le pasa a un softmax cuando le llegan números grandes. Hay que devolver la escala a su sitio después de cada suma, y eso es el layer norm.

Fijemos la fila de la posición tt, con sus dmodeld_{\text{model}} coordenadas xt1,,xtdmodelx_{t1}, \dots, x_{t\,d_{\text{model}}}. Su media y su desviación típica son las de esos números y de ningún otro:

μt=1dmodeljxtj,st=1dmodelj(xtjμt)2.\mu_t = \frac{1}{d_{\text{model}}}\sum_{j} x_{tj}, \qquad s_t = \sqrt{\frac{1}{d_{\text{model}}}\sum_{j}\left(x_{tj} - \mu_t\right)^{2}}.

Con ellas se centra y se escala la fila, y después entran dos vectores que se aprenden, γ\boldsymbol{\gamma} y β\boldsymbol{\beta}, los dos de dmodeld_{\text{model}} coordenadas:

x^tj=xtjμtst2+ε,LayerNorm(xt)=γx^t+β.\hat{x}_{tj} = \frac{x_{tj} - \mu_t}{\sqrt{s_t^{2} + \varepsilon}}, \qquad \text{LayerNorm}(\mathbf{x}_t) = \boldsymbol{\gamma} \odot \hat{\mathbf{x}}_t + \boldsymbol{\beta}.

ε\varepsilon es una constante pequeña —10510^{-5} en casi todas las implementaciones— que evita dividir entre cero si una fila tiene sus coordenadas iguales. Los otros dos parecen deshacer el trabajo que se acaba de hacer, y en parte lo hacen: dejar cada fila con media 00 y desviación 11 es una condición que nadie ha pedido, y γ\boldsymbol{\gamma} y β\boldsymbol{\beta} devuelven esa libertad convertida en parámetros. Del todo no pueden deshacerla: μt\mu_t y sts_t cambian de fila en fila, y estos dos vectores son los mismos en todas.

Por qué la fila y no el batch

Sumar a lo ancho de la fila es una decisión, y la alternativa es la contraria: normalizar cada coordenada a lo largo de la columna, que es la normalización por batch. Sobre secuencias falla por tres sitios. Lo que saliera en la posición tt dependería de qué otras frases hubieran caído en el mismo batch; dependería de TT y de lo que hubiera en las demás posiciones, incluidas las de padding, que no dicen nada; y habría que guardar estadísticas del entrenamiento para responder después a una frase suelta.

Con la fila no aparece ninguna de las tres, porque μt\mu_t y sts_t salen de la fila tt y de nada más. Guarda la tercera: la caja «Suma y layer norm» no mezcla posiciones, y por eso la afirmación de la lección sobre el adiós a la recurrencia —tres de las quince cajas miran más allá de su propia fila— sigue en pie ahora que sabemos qué hacen las otras doce.

La subcapa que no mira a ninguna otra posición

Falta la segunda subcapa: la que el artículo llama FFN (feed-forward network) y esta lección llama el perceptrón por posiciones —los dos nombres son la misma caja—. Por dentro es el perceptrón multicapa del bloque 2 con dos capas y ReLU en medio, aplicado a una fila cada vez:

FFN(X)=ReLU(XW1+1Tb1)W2+1Tb2,\text{FFN}(\mathbf{X}) = \text{ReLU}\left(\mathbf{X}\mathbf{W}_1 + \mathbf{1}_T\mathbf{b}_1^{\top}\right)\mathbf{W}_2 + \mathbf{1}_T\mathbf{b}_2^{\top},

con W1Rdmodel×dff\mathbf{W}_1 \in \mathbb{R}^{d_{\text{model}} \times d_{\text{ff}}} y W2Rdff×dmodel\mathbf{W}_2 \in \mathbb{R}^{d_{\text{ff}} \times d_{\text{model}}}, y con el subíndice nombrando cuál de las dos capas es. 1T\mathbf{1}_T es la columna de unos del bloque 2, que aquí repite el sesgo una vez por posición en lugar de una vez por ejemplo, y la ReLU es el max(0,)\max(0, \cdot) del artículo.

Las dos matrices son las mismas para todas las posiciones: la fila 33 y la fila 300300 pasan por los mismos pesos, así que la subcapa no puede enterarse de cuál está transformando. Eso es lo que significa «por posiciones» y lo que la deja fuera de las cajas que mezclan. Su papel sale de ahí por descarte: la atención no transforma, mezcla, y éste es el único sitio donde el contenido de una posición pasa por una no linealidad propia.

El reparto no es el que sugiere el dibujo. Con dmodel=512d_{\text{model}} = 512 y dff=2048d_{\text{ff}} = 2\,048:

2dmodeldff=2097152frente a4dmodel2=1048576.2 d_{\text{model}} d_{\text{ff}} = 2\,097\,152 \quad\text{frente a}\quad 4 d_{\text{model}}^{2} = 1\,048\,576.

A la izquierda las dos matrices del perceptrón por posiciones, a la derecha las cuatro de la atención sumadas: la caja que parece fontanería lleva el doble de pesos que la que le da nombre a la arquitectura. Ninguno de los dos números menciona a TT.

El bloque completo, en dos líneas

Con las tres piezas puestas, un bloque del encoder es la misma línea dos veces:

LayerNorm(X+Sublayer(X)),\text{LayerNorm}\left(\mathbf{X} + \text{Sublayer}(\mathbf{X})\right),

con MultiHead\text{MultiHead} en el hueco la primera vez y FFN\text{FFN} la segunda, y con la salida de la primera haciendo de X\mathbf{X} de la segunda. Sesgos y vectores de normalización incluidos, son 31503363\,150\,336 parámetros por bloque, y el encoder del artículo apila N=6N = 6: unos 18.918.9 millones, ninguno de ellos dependiente de la longitud del texto.

Y ahora la concesión. El artículo normaliza después de sumar, y eso deja el layer norm encima del camino de la identidad: el sumando que vale I\mathbf{I} está dentro del paréntesis, pero lo que sale de la caja ha pasado además por una normalización, así que el camino limpio de la primera sección no es tan limpio. Casi todas las implementaciones de hoy escriben X+Sublayer(LayerNorm(X))\mathbf{X} + \text{Sublayer}\left(\text{LayerNorm}(\mathbf{X})\right) —normalizar lo que entra en la subcapa y dejar la suma sin tocar—, que es lo que se conoce como pre-norm. Cuál entrena mejor queda fuera de este curso; que son dos operaciones distintas, no.

Las dos piezas en NumPy

La primera celda es el layer norm sobre cinco posiciones. Mira tres cosas: la media y la desviación típica de cada fila, qué devuelven γ\boldsymbol{\gamma} y β\boldsymbol{\beta}, y qué le pasa a la fila 00 cuando se toca la fila 44.

import numpy as np

rng = np.random.default_rng(7)
X = rng.normal(loc=3.0, scale=2.0, size=(5, 8)) # 5 posiciones, d_model = 8


def layer_norm(X, gamma, beta, eps=1e-5):
mu = X.mean(axis=1, keepdims=True) # una media por fila
s = X.std(axis=1, keepdims=True) # una desviacion por fila
return gamma * (X - mu) / np.sqrt(s**2 + eps) + beta


gamma, beta = np.ones(8), np.zeros(8)
Y = layer_norm(X, gamma, beta)

print("entra media por fila:", np.round(X.mean(axis=1), 3))
print(" desviacion :", np.round(X.std(axis=1), 3))
print("sale media por fila:", np.round(Y.mean(axis=1), 3))
print(" desviacion :", np.round(Y.std(axis=1), 3))

Z = layer_norm(X, np.full(8, 4.0), np.full(8, -1.0))
print("\ncon gamma = 4 y beta = -1:", np.round(Z.mean(axis=1), 3), np.round(Z.std(axis=1), 3))

# Toco SOLO la fila 4 y miro que le pasa a la fila 0.
X2 = X.copy()
X2[4] += 10.0
por_columnas = lambda M: (M - M.mean(axis=0)) / M.std(axis=0)
print("\ntoco la fila 4; cuanto se mueve la fila 0")
print(" normalizando por filas :",
round(float(np.abs(layer_norm(X2, gamma, beta)[0] - Y[0]).max()), 12))
print(" normalizando por columnas:",
round(float(np.abs(por_columnas(X2)[0] - por_columnas(X)[0]).max()), 3))
numpy

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

Las cinco filas entran con medias entre 0.980.98 y 2.892.89 y salen con media 00 y desviación 11; con γ=4\boldsymbol{\gamma} = 4 y β=1\boldsymbol{\beta} = -1 salen con media 1-1 y desviación 44, que es lo que esos dos vectores significan. Las dos últimas líneas son la sección anterior comprobada: tocar la fila 44 mueve la fila 00 en 0.00.0 por filas y en 1.8891.889 por columnas.

La segunda mide lo que hace la suma, y para verla sola le quito al bloque el layer norm y la atención: quedan 1212 perceptrones por posiciones, uno encima de otro, con y sin la línea que los rodea. La escala de los pesos la he puesto baja a propósito, 0.150.15: al arrancar, una subcapa es una corrección pequeña.

import numpy as np

rng = np.random.default_rng(0)
T, d_model, d_ff, N = 8, 16, 32, 12

pesos = [(rng.normal(size=(d_model, d_ff)) * 0.15,
rng.normal(size=(d_ff, d_model)) * 0.15) for _ in range(N)]


def apilar(X, n, con_suma):
for W1, W2 in pesos[:n]:
sub = np.maximum(0.0, X @ W1) @ W2 # el perceptron por posiciones
X = X + sub if con_suma else sub
return X


X = rng.normal(size=(T, d_model))
h = 1e-5
Xh = X.copy()
Xh[0, 0] += h # muevo UNA coordenada de la fila 0

print("capas con la suma sin la suma")
for n in (1, 2, 4, 6, 8, 12):
con = np.linalg.norm((apilar(Xh, n, True) - apilar(X, n, True)) / h)
sin = np.linalg.norm((apilar(Xh, n, False) - apilar(X, n, False)) / h)
print(f"{n:5d} {con:11.4f} {sin:11.3e}")
numpy

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

Cada número es una derivada medida por sondeo numérico: cuánto se mueve la salida entera al empujar una sola coordenada de la entrada. Sin la suma cae de 0.1090.109 con una capa a 5.182×1075.182 \times 10^{-7} con doce —siete órdenes de magnitud, y con doce capas, no con doce mil—, porque es un producto de doce factores pequeños. Con la suma se queda entre 1.001.00 y 1.801.80 todo el camino: el sumando que no atraviesa ningún peso vale 11, y los demás sólo lo corrigen.

Comprueba tu intuición

Cinco preguntas: qué añade la suma a la derivada, qué le exige a la subcapa, qué números entran al normalizar una fila, dónde están los parámetros y qué cambia al mover el layer norm de sitio.

Cada subcapa se envuelve en x+Sublayer(x)\mathbf{x} + \text{Sublayer}(\mathbf{x}). Al componer NN bloques y derivar, ¿qué aparece que en la cuenta del bloque 3 no había?

¿Qué le exige a una subcapa el hecho de que su salida se le sume a su propia entrada?

El layer norm del bloque normaliza cada fila con sus propias dmodeld_{\text{model}} coordenadas. Marca lo que es cierto de lo que sale en la posición tt.

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

Con dmodel=512d_{\text{model}} = 512 y dff=2048d_{\text{ff}} = 2\,048, ¿cuántos pesos tienen las dos matrices del perceptrón por posiciones, sin contar los sesgos?

pesos

A margin of ±0 is accepted.

Normalizar la salida de la subcapa y sumarle la entrada después, x+LayerNorm(Sublayer(x))\mathbf{x} + \text{LayerNorm}\left(\text{Sublayer}(\mathbf{x})\right), es una operación distinta de la del artículo.

Monta el bloque entero con las dos piezas de la lección.

layer_norm(X, gamma, beta, eps=1e-5) recibe X de forma (T,dmodel)(T, d_{\text{model}}) y devuelve otra igual: cada fila centrada con su propia media μt\mu_t y dividida por st2+ε\sqrt{s_t^{2} + \varepsilon}, con sts_t la desviación típica de esa misma fila, y después multiplicada coordenada a coordenada por gamma y sumada a beta.

bloque(X, atencion, p) aplica las dos subcapas. atencion es una función que recibe X y devuelve algo de la misma forma; p es un diccionario con W1 de forma (dmodel,dff)(d_{\text{model}}, d_{\text{ff}}), b1 de (dff,)(d_{\text{ff}},), W2 de (dff,dmodel)(d_{\text{ff}}, d_{\text{model}}), b2 de (dmodel,)(d_{\text{model}},) y los cuatro vectores gamma1, beta1, gamma2, beta2, todos de (dmodel,)(d_{\text{model}},).

Cada subcapa va envuelta igual: se suma lo que entró a lo que salió y se normaliza el resultado. La segunda subcapa es el perceptrón por posiciones, ReLU(XW1+b1)W2+b2\text{ReLU}(\mathbf{X}\mathbf{W}_1 + \mathbf{b}_1)\mathbf{W}_2 + \mathbf{b}_2.

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.


El bloque está entero y es apilable, y con eso está terminada una de las dos columnas del dibujo. Vuelve a mirarlo: la de la derecha tiene tres subcapas donde la izquierda tiene dos, la de en medio recibe una flecha que cruza desde la otra columna, y la de abajo se llama «auto-atención enmascarada», que es una palabra que no ha salido en estas seis lecciones.

Las tres cosas son la misma cosa, y es la que separa leer un texto de escribirlo. Un encoder puede mirar la frase entera porque la tiene delante; un decoder produce la salida token a token, y al escribir el cuarto no puede mirar el quinto, que todavía no existe. Eso es la lección siguiente, sobre el encoder, el decoder y las máscaras, y lo bueno es que no toca ni una línea de lo que llevamos: tachar el futuro se hace sumándole a la rejilla de puntuaciones una matriz de ceros y de menos infinitos antes del softmax.

Further reading4 sources · 4 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.1 da la envoltura LayerNorm(x + Sublayer(x)) y el apilado N = 6; su §3.3, el perceptrón por posiciones con d_ff = 2048. Cada una en una frase: por qué la suma salva el apilado que hundió al bloque 3 es cosa tuya.

  • Layer Normalization
    paperBa, Kiros y Hinton, 2016arXiv:1607.06450EN

    Normaliza cada caso con sus propias unidades en vez de con la columna del batch, y por tus mismas razones: vale en recurrentes y no guarda estadísticas para responder a una frase suelta. Lo ensaya en RNN, no en un Transformer.

  • Identity Mappings in Deep Residual Networks
    paperHe, Zhang, Ren y Sun, 2016arXiv:1603.05027EN

    Analiza lo que aquí desarrollas: con la conexión identidad el gradiente lleva un sumando que baja directo entre bloques y no se cancela. Es sobre redes de imagen y pide la identidad también tras la suma, algo que el Transformer no hace.

  • On Layer Normalization in the Transformer Architecture
    paperXiong, Yang, He y otros, 2020arXiv:2002.04745EN

    Contesta lo que la lección deja abierto: normalizar antes de la subcapa (pre-norm) y no después deja el gradiente sano al arrancar y quita el calentamiento de la tasa. El post-norm del artículo lo necesita. Medido, no demostrado.