Adiós a la recurrencia

Adiós a la recurrencia

22 min de lectura

El bloque anterior terminó dejando la atención en los huesos: tres listas de vectores, una función que compara y una suma ponderada, sin una sola recurrencia por dentro. Lo que no tocó fue la máquina que la rodea, y esa sigue siendo recurrente de arriba abajo: quien llena las listas es un decoder que produce un estado por paso y no sabe producir el siguiente antes que el anterior. La operación ya no espera a nada; espera a quien se la pide.

Esta lección quita esa máquina entera y se queda mirando lo que sobrevive. Y la pregunta no es estética, porque la recurrencia cobra un precio que se puede contar. Toma una noticia de 800800 tokens y una red recurrente cualquiera: leerla son 800800 pasos que ocurren uno después de otro, y comprar una máquina con 800800 procesadores no cambia nada, porque el paso 401401 necesita un número que el paso 400400 todavía no ha terminado de calcular. Esa es la primera factura. La segunda la pagó el bloque 3 entero: lo que la posición 11 tenga que decirle a la posición 800800 viaja por 799799 multiplicaciones por la misma matriz, y para cuando llega no queda casi nada. Las dos facturas tienen el mismo origen, y quitarlo se lleva las dos por delante.

El diagrama de abajo es dónde acaba todo esto: la arquitectura completa que el bloque construye, la figura 1 del artículo Attention is All You Need, caja por caja. A estas alturas la mitad de los nombres no te dicen nada, y no pasa nada —cada caja lleva escrita la lección que la construye—. Lo que sí se lee hoy es lo único que esta lección afirma: enciende el marcador y, de las quince cajas, se iluminan tres. Las otras doce trabajan sobre una fila cada vez y no tienen forma de saber qué hay en la de al lado.

Quince cajas y ninguna flecha que vaya de una posición a la siguiente. Las tres que mezclan posiciones son las tres de atención; el resto trata cada posición por separado.

La cadena de esperas que ningún ordenador acorta

La recurrencia del bloque anterior se escribe en una línea, y esa línea lo dice todo:

ht=tanh(Whhht1+Wxhxt+bh),\mathbf{h}_t = \tanh\left(\mathbf{W}_{hh}\mathbf{h}_{t-1} + \mathbf{W}_{xh}\mathbf{x}_t + \mathbf{b}_h\right),

con htRdh\mathbf{h}_t \in \mathbb{R}^{d_h} y h0=0\mathbf{h}_0 = \mathbf{0}. De los dos sumandos, el segundo no espera a nadie: Wxhxt\mathbf{W}_{xh}\mathbf{x}_t sólo mira el token de su posición, así que los TT productos se pueden calcular a la vez en cuanto el texto está tokenizado. El primero sí espera. ht1\mathbf{h}_{t-1} es el resultado del paso anterior, y no hay manera de adelantarlo sin calcularlo. Una de las dos mitades se reparte entre procesadores y la otra no, y la que no se reparte es la que fija el resultado: la profundidad de la cadena —cuántas operaciones tienen que ocurrir en fila, una detrás de otra— es TT.

Ponlo frente a la llamada de atención tal como quedó escrita. La puntuación de la pareja (i,j)(i,j) es

eij=a(qi,kj),e_{ij} = a\left(\mathbf{q}_i, \mathbf{k}_j\right),

y ahí no aparece ningún ee de ninguna otra pareja. Cada una de las T2T^{2} puntuaciones se calcula con dos vectores que ya existen, de modo que las T2T^{2} se pueden calcular simultáneamente —que es lo que significa que quepan en un solo producto de matrices, QK\mathbf{Q}\mathbf{K}^{\top}—. El softmax normaliza cada fila con la suya y ninguna fila mira a otra; la mezcla es un segundo producto. Tres operaciones en fila, y ese tres no se mueve cuando el texto se alarga.

La otra distancia, la que el bloque 3 no pudo acortar

La profundidad de la cadena decide cuánto se tarda. Hay una segunda distancia, y decide algo peor.

Pregúntate por dónde viaja lo que la posición 11 le aporta a la posición TT. En la recurrencia viaja por dentro del estado, y el estado se reescribe en cada paso, así que ese aporte pasa por T1T-1 productos por Whh\mathbf{W}_{hh} antes de llegar. Es exactamente la cuenta de la lección sobre el gradiente que se desvanece: multiplicar T1T-1 veces por la misma matriz apaga la señal si su radio espectral es menor que 11 y la dispara si es mayor, y la LSTM (long short-term memory) y la GRU (gated recurrent unit) no acortaron ese camino, le abrieron una vía aditiva al lado para que sobreviviera el viaje.

En la atención el camino entre esas dos posiciones es un peso:

cT=αT1v1+αT2v2++αTTvT.\mathbf{c}_T = \alpha_{T1}\,\mathbf{v}_1 + \alpha_{T2}\,\mathbf{v}_2 + \dots + \alpha_{TT}\,\mathbf{v}_T.

Lo que aporta la posición 11 es su sumando, y delante lleva un solo factor: αT1\alpha_{T1}. Uno, no T1T-1. Y como el camino de vuelta es el de ida recorrido al revés, el gradiente que llega a la posición 11 desde la pérdida del paso TT tampoco atraviesa nada intermedio. La distancia entre dos posiciones cualesquiera deja de depender de cuántas haya en medio, que es más de lo que el bloque 3 consiguió: la compuerta hizo el viaje sobrevivible, no corto.

Lo que cuesta cambiar los pasos por una rejilla

Nada de esto es gratis, y el precio se cuenta en multiplicaciones. Empieza por la recurrencia. Un paso multiplica WhhRdh×dh\mathbf{W}_{hh} \in \mathbb{R}^{d_h \times d_h} por un vector y WxhRdh×dmodel\mathbf{W}_{xh} \in \mathbb{R}^{d_h \times d_{\text{model}}} por otro, y multiplicar una matriz por un vector cuesta una multiplicación por casilla, así que el paso cuesta un producto por cada peso que tiene. Esa es la cuenta de la lección sobre la RNN (recurrent neural network) vanilla, leída como coste en lugar de como memoria. Sobre la frase entera:

T(dhdh+dhdmodel)multiplicaciones.T\left(d_h \cdot d_h + d_h \cdot d_{\text{model}}\right) \quad \text{multiplicaciones}.

Ahora la rejilla. QK\mathbf{Q}\mathbf{K}^{\top} multiplica una matriz T×dkT \times d_k por una dk×Td_k \times T: cada una de las T2T^{2} casillas del resultado es un producto escalar de dkd_k términos, y son T2dkT^{2} \cdot d_k multiplicaciones. La mezcla AV\mathbf{A}\mathbf{V} tiene TdvT \cdot d_v casillas y cada una suma TT términos, otras T2dvT^{2} \cdot d_v. En total:

T2(dk+dv)multiplicaciones.T^{2}\left(d_k + d_v\right) \quad \text{multiplicaciones}.

Las dos cuentas se comparan mejor con las cuatro anchuras iguales. Supongámoslas iguales y llamémoslas dd: quedan 2Td22Td^{2} frente a 2T2d2T^{2}d, y la razón entre las dos es T/dT/d. La rejilla cuesta menos que la recurrencia mientras la frase mida menos que el modelo, y más en cuanto lo pasa; se cruzan cuando T=dT = d, y no antes. En el artículo estas dos cuentas aparecen escritas como O(T2d)O(T^{2}d) frente a O(Td2)O(Td^{2}), que es la misma comparación con las constantes borradas; te la vas a encontrar así en su tabla 1.

Lo que se ha comprado con eso conviene decirlo entero, porque no es «más rápido». Es un cambio de eje: el trabajo total crece —con TT al cuadrado, en textos largos—, y a cambio ese trabajo se puede repartir, mientras que los TT pasos de la recurrencia no se pueden repartir de ninguna manera. Una cuenta grande que cabe en un solo producto de matrices se resuelve antes que una cuenta pequeña partida en TT esperas.

El precio que no se paga en multiplicaciones

Vuelve a mirar la puntuación, y esta vez fíjate en los índices. En eij=a(qi,kj)e_{ij} = a(\mathbf{q}_i, \mathbf{k}_j) los subíndices sólo señalan qué dos vectores entran; dentro de la cuenta no aparece ni ii ni jj. Lo mismo vale para la mezcla, que es una suma, y una suma no distingue el orden de sus sumandos. La consecuencia es incómoda: si barajas las filas de K\mathbf{K} y V\mathbf{V}, las salidas se barajan igual y ningún número cambia.

De modo que el perro muerde al cartero y al cartero muerde el perro llegan a la capa como el mismo conjunto de vectores. La recurrencia no tenía este problema, y no por virtud propia: lo tenía resuelto de nacimiento, porque leía en orden y el orden de lectura era el orden de la frase. El MLP (multilayer perceptron) concatenado de la lección sobre por qué falla el MLP en secuencias, tampoco lo tenía, por la razón contraria: cada posición multiplicaba su propio bloque de pesos, y por eso sabía distinguirlas y no sabía generalizar entre ellas. La atención no tiene ni lo uno ni lo otro.

La concesión es grande y esta lección no la arregla: quitar la recurrencia se lleva por delante lo único que sabía el orden. Se arregla sumándole a cada fila algo que dependa de dónde está, y eso es una lección más adelante, sobre la codificación posicional. Hasta entonces, todo lo que sigue trata la entrada como una bolsa de vectores.

Compruébalo a mano con una frase de 50 tokens

Haz las dos cuentas con números, que es donde se ve que el cambio no es teórico. Toma d=512d = 512 —la anchura del modelo del artículo— y una frase de T=50T = 50 tokens.

La recurrencia: 2505122=262144002 \cdot 50 \cdot 512^{2} = 26\,214\,400 multiplicaciones, repartidas en 5050 pasos que van uno detrás de otro. La rejilla: 2502512=25600002 \cdot 50^{2} \cdot 512 = 2\,560\,000, y no en 5050 esperas sino en dos productos de matrices. Diez veces menos trabajo, y la profundidad de la cadena ha pasado de 5050 a un número que ni siquiera menciona a 5050.

Sube ahora a T=1000T = 1\,000, un texto largo. La rejilla pide 210002512=10240000002 \cdot 1\,000^{2} \cdot 512 = 1\,024\,000\,000 y la recurrencia 210005122=5242880002 \cdot 1\,000 \cdot 512^{2} = 524\,288\,000: casi el doble, y esta vez el caro es el nuevo. La razón entre los dos es 1000/5121.951\,000/512 \approx 1.95, tal como decía T/dT/d. Hazlo tú con T=2000T = 2\,000 y comprueba que sale 3.93.9 —y fíjate en que la profundidad de la cadena sigue sin moverse, mientras la de la recurrencia se ha ido a 20002\,000—.

Ese es el trato, con sus dos columnas a la vista: se paga trabajo, que se puede repartir entre procesadores, y se compra espera, que no se puede repartir entre nada.

Comprueba tu intuición

Cuatro preguntas: qué es lo que no se puede solapar, dónde se cruzan las dos cuentas, qué pasa si barajas las filas y qué le ocurre al camino entre dos posiciones lejanas.

Tienes un ordenador con tantos procesadores como quieras y una frase de TT tokens ya tokenizada entera. ¿Qué es lo que sigue costando TT pasos que no se pueden solapar?

Con dh=dmodel=dk=dv=512d_h = d_{\text{model}} = d_k = d_v = 512, ¿a partir de qué longitud TT cuesta más multiplicaciones la rejilla de atención que la recurrencia?

tokens

Se acepta un margen de ±0.

Le pasas a la llamada de atención las mismas filas en otro orden —los mismos tokens, barajados—. ¿Qué sale?

Sobre el camino que une la posición 11 con la posición TT, marca lo que es cierto.

Marca todas las opciones correctas. Se corrige todo o nada: no hay puntuación parcial.


Queda un agujero en el sitio más visible. La llamada de atención necesita tres listas, y en todo el bloque anterior las consultas venían del decoder: eran sus estados, uno por paso. Ese decoder acaba de desaparecer del dibujo. Nadie fabrica ya Q\mathbf{Q}, y sin consultas la operación que esta lección ha puesto en el centro de la arquitectura no tiene con qué empezar.

La respuesta es la que le da nombre a la lección siguiente, sobre la auto-atención: que pregunte la propia secuencia. Si las consultas, las claves y los valores son tres lecturas de la misma lista de vectores, no hace falta ninguna segunda red que produzca nada —y cada posición pasa a mirar a todas las demás, incluida ella misma—. Es la operación de siempre con los tres papeles repartidos de otra manera, y es lo que convierte el diagrama de arriba en algo que se puede calcular.

Para profundizar1 fuente · 1 paper

De dónde sale lo de esta lección, y dónde seguir si quieres más. Nada de aquí hace falta para continuar el curso.

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

    Su §4 y su tabla 1 reúnen las tres cuentas de esta lección: operaciones que no se solapan, camino entre dos posiciones y coste por capa, O(T²d) frente a O(Td²). La tabla borra las constantes; el cruce en T = d lo despejas tú.