La RNN vanilla

La RNN vanilla

26 min read

La lección anterior, sobre por qué el perceptrón multicapa (multilayer perceptron, MLP) falla con secuencias, dejó dos defectos medidos: lo que hay que recortar de un texto largo, y lo que la red averigua en una posición y no le sirve en otra. Los dos salen de la misma causa, y esa causa no está en el tamaño de la red ni en cómo se la entrena. Está en la entrada: aquella red no leía una secuencia, recibía un vector. Lo que sigue es la corrección, y llega con una propiedad que de una corrección no se espera: la red que la aplica es más pequeña que la que falla.

Se llama red neuronal recurrente (recurrent neural network, RNN), y en su forma original —la que todo el mundo llama vanilla— cabe en dos matrices y un bucle. Lo difícil de ella no es cómo funciona, que son dos líneas, sino por qué algo tan corto consigue lo que la red grande no conseguía. La respuesta está entera en una renuncia: esta red no mira la frase, la recorre.

Recórrela antes de seguir. En el explorable de abajo, cada avance aplica el mismo paso a una posición más: las dos matrices son siempre las mismas y sólo cambia el estado que entra por la izquierda. Así es, además, como la vas a encontrar dibujada en cualquier libro que la trate.

Avanza posición a posición: comprueba que Whh y Wxh son las mismas en cada paso, y que lo único que cambia de una posición a otra es el estado que entra por la izquierda.

Hay una lectura que esa secuencia invita a hacer y que es falsa: las cajas no son capas distintas. Son la misma función dibujada una vez por posición, como se dibujarían las vueltas de un bucle for, y por eso los rótulos se repiten en lugar de numerarse. Una red con una capa por posición tendría un juego de pesos por cada una; ésta tiene uno solo, valga la secuencia lo que valga.

Un estado que se reescribe en cada posición

Cada posición aporta el vector de su entrada del vocabulario, xt=ewtRdmodel\mathbf{x}_t = \mathbf{e}_{w_t} \in \mathbb{R}^{d_{\text{model}}}, la fila de E\mathbf{E} que le corresponde, igual que en la lección anterior. Lo nuevo es el otro vector, el que la red se pasa a sí misma de un paso al siguiente: el estado oculto htRdh\mathbf{h}_t \in \mathbb{R}^{d_h}, donde dhd_h se elige al construir la red y es lo único que decide cuánto cabe en esa memoria. Fijemos además h0=0\mathbf{h}_0 = \mathbf{0}, porque antes de leer el primer token no hay nada que recordar.

Con eso, un paso es una capa corriente a la que entran dos cosas en lugar de una:

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 WhhRdh×dh\mathbf{W}_{hh} \in \mathbb{R}^{d_h \times d_h} leyendo el estado anterior, WxhRdh×dmodel\mathbf{W}_{xh} \in \mathbb{R}^{d_h \times d_{\text{model}}} leyendo el token de turno y bhRdh\mathbf{b}_h \in \mathbb{R}^{d_h}. La tangente hiperbólica ocupa el hueco φ\varphi de la lección sobre funciones de activación; ahí caben las tres que se vieron allí, y la escribo con su nombre porque es la que la recurrencia lleva en todo el bloque.

El paso t=1t = 1 parece pedir una regla aparte, y no la pide. Con h0=0\mathbf{h}_0 = \mathbf{0} el primer término desaparece solo:

h1=tanh(Whh0+Wxhx1+bh)=tanh(Wxhx1+bh).\mathbf{h}_1 = \tanh\left(\mathbf{W}_{hh}\mathbf{0} + \mathbf{W}_{xh}\mathbf{x}_1 + \mathbf{b}_h\right) = \tanh\left(\mathbf{W}_{xh}\mathbf{x}_1 + \mathbf{b}_h\right).
Ver la secuencia entera dentro del último estado

Sustituye ht1\mathbf{h}_{t-1} por su propia definición, empezando por la de arriba:

h2=tanh(Whhtanh(Wxhx1+bh)+Wxhx2+bh),\mathbf{h}_2 = \tanh\left(\mathbf{W}_{hh}\tanh\left(\mathbf{W}_{xh}\mathbf{x}_1 + \mathbf{b}_h\right) + \mathbf{W}_{xh}\mathbf{x}_2 + \mathbf{b}_h\right),

y una vez más, que ya deja el patrón a la vista:

h3=tanh(Whhtanh(Whhtanh(Wxhx1+bh)+Wxhx2+bh)+Wxhx3+bh).\mathbf{h}_3 = \tanh\left(\mathbf{W}_{hh}\tanh\left(\mathbf{W}_{hh}\tanh\left(\mathbf{W}_{xh}\mathbf{x}_1 + \mathbf{b}_h\right) + \mathbf{W}_{xh}\mathbf{x}_2 + \mathbf{b}_h\right) + \mathbf{W}_{xh}\mathbf{x}_3 + \mathbf{b}_h\right).

Cada nueva posición envuelve todo lo anterior en una capa más de tanh\tanh y de Whh\mathbf{W}_{hh}, sin que se caiga nada por el camino: x1\mathbf{x}_1 sigue ahí dentro cuando tt vale 33, y seguiría con t=300t = 300.

Ese anidamiento contesta la pregunta de dónde queda guardado el orden. No está en los pesos —son los mismos en las tres posiciones— sino en cuántas veces ha atravesado cada token la matriz Whh\mathbf{W}_{hh} antes de llegar al final: en h3\mathbf{h}_3, el vector x1\mathbf{x}_1 la ha cruzado dos veces, x2\mathbf{x}_2 una y x3\mathbf{x}_3 ninguna. Cambia dos tokens de sitio y cambias esa cuenta.

Queda sacar una respuesta del estado. Para un veredicto por documento, se lee del último:

y^=σ(WhyhT+by),\hat{\mathbf{y}} = \sigma\left(\mathbf{W}_{hy}\mathbf{h}_T + \mathbf{b}_y\right),

con Why\mathbf{W}_{hy} de forma 1×dh1 \times d_h —una fila por salida, y aquí la salida es una— y by\mathbf{b}_y de una coordenada. La activación es otra, y no por descuido: σ\sigma es la sigmoide logística y no la tanh\tanh del bucle, porque lo que la red entrega es una probabilidad y la tanh\tanh devuelve números entre 1-1 y 11. Es la capa de salida del clasificador de reseñas del bloque anterior sin un solo cambio; lo distinto es de dónde viene lo que entra en ella. Por eso el dibujo de arriba no la lleva y la cuenta de más abajo tampoco: la recurrencia son las dos matrices del bucle, y esta tercera es una capa prestada que ya sabes montar.

Repasa ahora las formas de esas dos matrices y comprueba qué no aparece en ninguna: TT. Whh\mathbf{W}_{hh} mide dh×dhd_h \times d_h, Wxh\mathbf{W}_{xh} mide dh×dmodeld_h \times d_{\text{model}}, y ninguna de las dos sabe cuántas veces se la va a usar. Eso deja sin objeto las dos operaciones que la longitud fija imponía: no hay un TmaxT_{\max} por el que truncar, y no hay posiciones sobrantes que rellenar con el vector nulo —el padding de la lección anterior—, porque no hay posiciones de más. Hay exactamente las que trae el documento.

Debajo de eso hay un cambio de unidad que conviene ver. En el MLP concatenado, una entrada de la red era un documento: la primera capa tenía tantas columnas como ocupaba el documento entero. Aquí una entrada de la red es un token, porque Wxh\mathbf{W}_{xh} tiene dmodeld_{\text{model}} columnas y ninguna más, y en cada paso la red no ve nada aparte del vector de esa posición y del estado que trae. La longitud del texto no ha desaparecido del sistema: se ha mudado de las formas al número de vueltas del bucle. Y esa mudanza es la que lo cambia todo, porque una forma queda fija al construir la red y ya no se toca —eso es lo que obligaba a truncar—, mientras que el número de vueltas lo decide cada documento al llegar.

Una sola matriz para todas las posiciones

Suma lo que ocupan las tres piezas de la recurrencia:

dhdhWhh  +  dhdmodelWxh  +  dhbhpesos.\underbrace{d_h \cdot d_h}_{\mathbf{W}_{hh}} \;+\; \underbrace{d_h \cdot d_{\text{model}}}_{\mathbf{W}_{xh}} \;+\; \underbrace{d_h}_{\mathbf{b}_h} \quad \text{pesos}.

Ponla al lado de la de la lección anterior, d1Tmaxdmodeld_1 \cdot T_{\max} \cdot d_{\text{model}}. Lo que hay que mirar no es cuál sale más pequeña —eso es una consecuencia— sino qué letras aparecen en cada una: allí estaba TmaxT_{\max} y aquí no está. Leer documentos diez veces más largos no añade un peso.

Y el ahorro es lo de menos. Lo que cambia es a qué se dedica cada peso. En la concatenación, el trozo de la primera capa que leía la posición 2 y el que leía la posición 8 eran dos matrices distintas, ajustadas con ejemplos distintos, y la segunda se quedaba como salió de la inicialización si el corpus nunca traía nada interesante tan al final. En la recurrencia hay una Wxh\mathbf{W}_{xh}, y la atraviesan los vectores de todas las posiciones de todos los documentos. Lo que aprenda a hacer con el vector de guion lo aplica esté guion donde esté, porque no tiene forma de saber dónde está.

Compartir los pesos es lo que hace que una RNN sea una RNN. Quítalo —una matriz distinta por posición— y lo que queda es el MLP concatenado con otro nombre.

El reverso viene con el gradiente, y conviene enunciarlo aquí aunque esta lección no lo resuelva. Whh\mathbf{W}_{hh} no interviene una vez en el cálculo de hT\mathbf{h}_T: interviene TT veces, como enseña el anidamiento de arriba. Así que la pérdida depende de esa matriz por TT caminos distintos, y su gradiente tendrá que recogerlos todos.

Ejecutando la recurrencia sobre dos frases de distinta longitud

Ocho entradas de vocabulario, embeddings de dmodel=4d_{\text{model}} = 4 y un estado de dh=3d_h = 3 coordenadas. Los pesos me los da un generador con semilla fija y no están entrenados: lo que hay que mirar en esta celda no es qué valores salen, sino cuántas veces se usa cada matriz. paso es la recurrencia transcrita línea por línea, y lee la aplica a una frase imprimiendo el estado en cada posición.

import numpy as np

V = ["la", "el", "actriz", "película", "guion", "salva", "hunde", "y"]
d_model, d_h = 4, 3
rng = np.random.default_rng(0)

E = rng.normal(size=(len(V), d_model)) * 0.5 # (8, 4): una fila por entrada
Wxh = rng.normal(size=(d_h, d_model)) * 0.5 # (3, 4): lee el token de turno
Whh = rng.normal(size=(d_h, d_h)) * 0.5 # (3, 3): lee el estado anterior
bh = np.zeros(d_h)


def paso(h, x):
return np.tanh(Whh @ h + Wxh @ x + bh) # los mismos pesos en cada posición


def lee(frase):
h, estados = np.zeros(d_h), [] # h_0 = 0, antes de leer nada
print(%s»" % frase)
for t, palabra in enumerate(frase.split(), 1):
h = paso(h, E[V.index(palabra)])
estados.append(h)
print(" h_%-2d %s" % (t, np.round(h, 3)))
return estados


corta = lee("el guion hunde la película")
larga = lee("la actriz salva la película y el guion la hunde")

print("\nformas: Whh %s Wxh %s bh %s —ninguna menciona T" % (
Whh.shape, Wxh.shape, bh.shape))
print("«el guion» leído en las posiciones 1-2 y en las 7-8:")
print(" %s %s" % (np.round(corta[0], 3), np.round(corta[1], 3)))
print(" %s %s" % (np.round(larga[6], 3), np.round(larga[7], 3)))
numpy

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

Cuenta las filas: cinco para la primera frase, diez para la segunda, y las mismas Whh y Wxh en las quince. No hay una versión de la red para textos cortos y otra para largos, y no ha hecho falta recortar ni rellenar nada —el bucle da las vueltas que le pida la frase—.

Las dos últimas líneas son la lección anterior contestada. el guion ocupa las posiciones 1 y 2 de la primera frase y las 7 y 8 de la segunda, y los estados que produce son casi los mismos: (0.509, 0.947, 0.072)(0.509,\ 0.947,\ 0.072) frente a (0.481, 0.935, 0.024)(0.481,\ 0.935,\ 0.024). Casi, no exactamente, y las dos mitades cuentan. Se parecen porque los pesos que leen esos dos tokens son los mismos pesos, estén donde estén; en el MLP concatenado no había ninguna razón para que se parecieran, y la lección anterior midió que no se parecían en absoluto. Y difieren porque el estado que traían de antes no era el mismo, que es justo la información de orden que la bolsa de palabras no tenía dónde guardar.

Esa diferencia se puede aislar. Toma dos frases con exactamente los mismos tokens y cámbialos de sitio.

import numpy as np

V = ["la", "el", "actriz", "película", "guion", "salva", "hunde", "y"]
d_model, d_h = 4, 3
rng = np.random.default_rng(0) # la misma red de la celda anterior
E = rng.normal(size=(len(V), d_model)) * 0.5
Wxh = rng.normal(size=(d_h, d_model)) * 0.5
Whh = rng.normal(size=(d_h, d_h)) * 0.5
bh = np.zeros(d_h)


def estado_final(frase):
h = np.zeros(d_h)
for palabra in frase.split():
h = np.tanh(Whh @ h + Wxh @ E[V.index(palabra)] + bh)
return h


def bolsa(frase):
v = np.zeros(len(V), dtype=int)
for palabra in frase.split():
v[V.index(palabra)] += 1
return v


for frase in ["la actriz salva la película", "la película salva la actriz"]:
print(%s»" % frase)
print(" h_T %s" % np.round(estado_final(frase), 3))
print(" bolsa %s" % bolsa(frase))

print("\n T_max=30 T_max=300")
concat = [128 * t * 100 for t in (30, 300)] # d1 · T_max · d_model
recurr = [128 * 128 + 128 * 100 + 128] * 2 # d_h² + d_h · d_model + d_h
print("concat. %-13d %d" % tuple(concat))
print("recurr. %-13d %d" % tuple(recurr))
numpy

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

La bolsa de palabras devuelve el mismo vector para las dos frases, casilla por casilla, porque suma y la suma no tiene orden. El estado final no: (0.117, 0.087, 0.171)(0.117,\ 0.087,\ 0.171) contra (0.141, 0.701, 0.632)(-0.141,\ 0.701,\ -0.632), y no es una diferencia de matiz, porque cambian los tres signos. La recurrencia se ha quedado con lo bueno de las dos representaciones que el bloque lleva comparando: distingue el orden, como la concatenación, y comparte los pesos, como la bolsa.

Las dos últimas líneas ponen el precio en números, con las dimensiones de la lección anterior: dmodel=100d_{\text{model}} = 100 y una capa —o un estado— de 128128 coordenadas. Pasar de leer frases a leer párrafos multiplica por diez la primera capa del MLP concatenado y deja la recurrencia exactamente donde estaba, en 2931229\,312 pesos. Los 38106883\,810\,688 que el MLP añade empiezan todos en su inicialización, y hay que ajustarlos con el mismo corpus de antes.

Comprueba tu intuición

Cuatro preguntas: cuántos pesos pide la recurrencia, por qué una posición que el entrenamiento nunca vio deja de ser un problema, qué dice y qué no dice la fórmula, y qué pasa si se anula la matriz que transporta el estado.

Una RNN con un estado de dh=64d_h = 64 coordenadas sobre embeddings de dmodel=128d_{\text{model}} = 128. ¿Cuántos pesos suman entre Whh\mathbf{W}_{hh}, Wxh\mathbf{W}_{xh} y bh\mathbf{b}_h?

A margin of ±0 is accepted.

En el conjunto de entrenamiento, guion no aparece nunca más allá de la cuarta posición. Llega un documento de prueba que lo trae en la novena, y la RNN lo trata con criterio en vez de ignorarlo. ¿Por qué?

Marca todo lo que sea cierto de la recurrencia 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 h0=0\mathbf{h}_0 = \mathbf{0}.

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

La misma recurrencia, con Whh\mathbf{W}_{hh} puesta a cero a mano. ¿Qué imprime?

import numpy as np
 
Whh, Wxh, bh = np.zeros((2, 2)), np.eye(2), np.zeros(2)
 
def final(xs):
    h = np.zeros(2)
    for x in xs:
        h = np.tanh(Whh @ h + Wxh @ x + bh)
    return h
 
a, b = np.array([1., 0.]), np.array([0., 1.])
print(np.allclose(final([a, b]), final([b, b])),
      np.allclose(final([a, b]), final([b, a])))
 

Y una recurrencia entera escrita por ti, que es la que vas a derivar en la lección siguiente.

Escribe la pasada hacia adelante de esta lección, para una secuencia de longitud cualquiera.

  • paso(h, x, Whh, Wxh, bh) devuelve el estado siguiente a partir del anterior y del vector del token de turno.
  • forward(X, Whh, Wxh, bh) recibe X de forma (T,dmodel)(T, d_{\text{model}}) —una fila por posición— y devuelve los TT estados apilados, de forma (T,dh)(T, d_h), partiendo de h0=0\mathbf{h}_0 = \mathbf{0}.

Ninguna de las dos puede escribir en los arrays que recibe.

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.


La red ya sabe leer y todavía no sabe aprender. Todo lo de arriba es la pasada hacia adelante: pesos puestos por un generador, ni un gradiente, ni un paso de descenso. Y el gradiente es donde compartir los pesos deja de salir gratis. En una red del bloque anterior cada matriz aparecía una vez en el camino de la entrada a la pérdida, y eso es lo que hacía que su gradiente fuera un solo producto; Whh\mathbf{W}_{hh} aparece TT veces.

Qué hacer con esos TT caminos —sumarlos, y recorrer la secuencia entera hacia atrás para obtenerlos— es la siguiente lección, sobre backpropagation a través del tiempo (backpropagation through time, BPTT). Es la regla de la cadena del bloque anterior sin ninguna variación, aplicada a un grafo que en vez de tener LL capas de altura tiene TT pasos de largo, y de ahí sale la cuenta completa de WhhL\nabla_{\mathbf{W}_{hh}}\mathcal{L}, WxhL\nabla_{\mathbf{W}_{xh}}\mathcal{L} y bhL\nabla_{\mathbf{b}_h}\mathcal{L}.

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.

  • Finding Structure in Time
    paperJeffrey L. Elman, 1990Cognitive Science 14EN

    Introdujo la red recurrente simple que montas aquí: un estado que se realimenta como contexto y esta misma recurrencia. Su parte inicial plantea el problema de la lección anterior; luego construye esta red. De 1990 y todavía se lee bien.