GRU: la versión simplificada

GRU: la versión simplificada

26 min de lectura

Hasta aquí el bloque ha avanzado preguntando qué le falta a cada red: a la RNN (recurrent neural network) vanilla le faltaba una memoria que sobreviviera a la distancia, y la LSTM (long short-term memory) de la lección anterior, sobre la memoria con compuertas, se la dio. Pero la dio cara —el cuádruple de los pesos de una RNN vanilla y un segundo estado que arrastrar en cada paso—, y la propia lección se despidió sospechando de que una de sus compuertas sobraba. Cuando una arquitectura funciona, la pregunta que queda no es qué le falta, sino qué le sobra.

Simplificar de verdad no es quitar piezas hasta que algo se rompe, sino quitarlas conservando lo que compraban. Y lo que la LSTM compró —lo único que justificaba todo el aparato— era la vía por la que el gradiente vuelve intacto a través de muchos pasos. Cualquier recorte tiene que pasar por ahí: si al fundir compuertas o borrar estados esa vía se pierde, la simplificación ha tirado justo lo que importaba. La GRU (gated recurrent unit) es el recorte que pasa la prueba, y por eso es la celda recurrente más usada en la práctica: tres cuartas partes de los pesos de la LSTM, un solo estado, y —casi siempre, no siempre— el mismo resultado.

El recorte cabe en dos ideas. La primera funde las dos compuertas sospechosas en una: en vez de decidir por separado cuánto se conserva y cuánto se escribe, la GRU usa un solo mando —la compuerta de actualización— que dice cuánto reemplazar, y lo que no reemplaza es lo que conserva. El estado deja de ser una suma que puede crecer sin freno y pasa a ser una mezcla, un punto intermedio entre lo que había y lo que se propone escribir; y una mezcla de cosas acotadas se queda acotada, así que ya no hace falta la compuerta de salida que en la LSTM recortaba el estado antes de enseñarlo.

La segunda idea es una compuerta pequeña y nueva, la compuerta de reset, que no toca la memoria sino la propuesta: deja que el candidato ignore el estado anterior cuando lo que hay que escribir ahora no arrastra nada de lo de antes —al empezar una oración nueva, por ejemplo—. Con esas dos, y con un único estado que hace de memoria y de salida a la vez, desaparece el segundo vector: ya no hay un estado de celda escondido detrás del que se enseña.

Antes de las fórmulas, vuelve un momento a la LSTM que la GRU simplifica. Recórrela y mira las dos primeras filas —la compuerta de olvido y la de entrada— moverse por separado: ésas son las dos que la GRU va a atar en una.

El paso de la LSTM de la lección anterior, para mirarlo con otros ojos. Fíjate en las dos primeras filas, la compuerta de olvido y la de entrada: se mueven por separado, y ésa es justo la libertad que la GRU ata en una sola compuerta. La última fila, lo que sobrevive de c₁, es la garantía que la simplificación no puede perder.

Conservar y escribir con una sola compuerta

Empieza por la línea que lo cambia todo. La celda de la LSTM sumaba dos términos con compuertas independientes, ct=ftct1+itc~t\mathbf{c}_t = \mathbf{f}_t \odot \mathbf{c}_{t-1} + \mathbf{i}_t \odot \tilde{\mathbf{c}}_t; la GRU trabaja directamente sobre el estado oculto htRdh\mathbf{h}_t \in \mathbb{R}^{d_h} —no hay segundo vector— y ata las dos compuertas en una:

ht=(1zt)ht1+zth~t,\mathbf{h}_t = (1 - \mathbf{z}_t) \odot \mathbf{h}_{t-1} + \mathbf{z}_t \odot \tilde{\mathbf{h}}_t,

con zt(0,1)dh\mathbf{z}_t \in (0,1)^{d_h} la compuerta de actualización y h~t(1,1)dh\tilde{\mathbf{h}}_t \in (-1,1)^{d_h} el candidato, el estado que el paso propone. Coordenada a coordenada es una media ponderada entre lo viejo y lo nuevo: zt,j=0z_{t,j} = 0 deja ht,jh_{t,j} igual a ht1,jh_{t-1,j} —conserva— y zt,j=1z_{t,j} = 1 lo reemplaza entero por el candidato —escribe—. No hay forma de hacer las dos cosas a la vez, y ésa es justo la libertad que la LSTM tenía y la GRU entrega: allí ft\mathbf{f}_t e it\mathbf{i}_t podían abrirse las dos y acumular. A cambio, una media de números acotados no se desborda, y por eso la GRU enseña el estado tal cual, sin la compuerta de salida ni el tanh\tanh que en la LSTM lo recortaban antes de sacarlo.

La compuerta de actualización se calcula como todas las de la lección anterior, una capa sigmoide sobre el token y el estado previo:

zt=σ(Wxzxt+Whzht1+bz),\mathbf{z}_t = \sigma\left(\mathbf{W}_{xz}\mathbf{x}_t + \mathbf{W}_{hz}\mathbf{h}_{t-1} + \mathbf{b}_z\right),

con WxzRdh×dmodel\mathbf{W}_{xz} \in \mathbb{R}^{d_h \times d_{\text{model}}}, WhzRdh×dh\mathbf{W}_{hz} \in \mathbb{R}^{d_h \times d_h} y bzRdh\mathbf{b}_z \in \mathbb{R}^{d_h}: las mismas formas de la recurrencia de la lección sobre la RNN vanilla.

Falta el candidato, y aquí entra la segunda compuerta. En la LSTM el candidato leía siempre el estado anterior completo; la GRU le pone delante un filtro, la compuerta de reset rt\mathbf{r}_t, que decide cuánto de ese estado entra en la propuesta:

rt=σ(Wxrxt+Whrht1+br),h~t=tanh(Wxhxt+Whh(rtht1)+bh).\mathbf{r}_t = \sigma\left(\mathbf{W}_{xr}\mathbf{x}_t + \mathbf{W}_{hr}\mathbf{h}_{t-1} + \mathbf{b}_r\right), \qquad \tilde{\mathbf{h}}_t = \tanh\left(\mathbf{W}_{xh}\mathbf{x}_t + \mathbf{W}_{hh}\left(\mathbf{r}_t \odot \mathbf{h}_{t-1}\right) + \mathbf{b}_h\right).

El rtht1\mathbf{r}_t \odot \mathbf{h}_{t-1} es lo único que distingue este candidato del de una RNN vanilla: con rt=1\mathbf{r}_t = \mathbf{1} el candidato es la recurrencia de la lección sobre la RNN vanilla tal cual, y con rt=0\mathbf{r}_t = \mathbf{0} en una coordenada esa coordenada del candidato se calcula sólo con el token, como si la secuencia empezara ahí. Conviene no confundir las dos compuertas: zt\mathbf{z}_t decide si el estado se actualiza, rt\mathbf{r}_t decide qué mira el candidato para proponer. Una gobierna la memoria; la otra, la propuesta.

Las tres piezas juntas, con h0=0\mathbf{h}_0 = \mathbf{0} por lo mismo que en la RNN y la LSTM:

zt=σ(Wxzxt+Whzht1+bz),rt=σ(Wxrxt+Whrht1+br),h~t=tanh(Wxhxt+Whh(rtht1)+bh),ht=(1zt)ht1+zth~t.\begin{aligned} \mathbf{z}_t &= \sigma\left(\mathbf{W}_{xz}\mathbf{x}_t + \mathbf{W}_{hz}\mathbf{h}_{t-1} + \mathbf{b}_z\right), & \mathbf{r}_t &= \sigma\left(\mathbf{W}_{xr}\mathbf{x}_t + \mathbf{W}_{hr}\mathbf{h}_{t-1} + \mathbf{b}_r\right), \\ \tilde{\mathbf{h}}_t &= \tanh\left(\mathbf{W}_{xh}\mathbf{x}_t + \mathbf{W}_{hh}\left(\mathbf{r}_t \odot \mathbf{h}_{t-1}\right) + \mathbf{b}_h\right), & \mathbf{h}_t &= (1 - \mathbf{z}_t) \odot \mathbf{h}_{t-1} + \mathbf{z}_t \odot \tilde{\mathbf{h}}_t \end{aligned}.

Tres piezas —dos compuertas y un candidato—, cada una con la forma exacta de la recurrencia de la lección sobre la RNN vanilla: seis matrices y tres sesgos, 3(dhdh+dhdmodel+dh)3\left(d_h \cdot d_h + d_h \cdot d_{\text{model}} + d_h\right) pesos. Frente a los cuatro grupos de la LSTM, tres cuartas partes; y como allí, TT no aparece en ninguna de las formas, así que la GRU lee cualquier longitud con el mismo juego de pesos.

Por qué la vía de vuelta sobrevive al recorte

Toda la lección pende de una comprobación: que la vía por la que el gradiente vuelve —la que la LSTM compró y la lección sobre el gradiente que se desvanece explicó por qué hacía falta— siga en pie con un estado en lugar de dos. Mira la última línea coordenada a coordenada, que es donde se ve lo que el \odot significa:

ht,j=(1zt,j)ht1,j+zt,jh~t,j.h_{t,j} = (1 - z_{t,j})\,h_{t-1,j} + z_{t,j}\,\tilde{h}_{t,j}.

La coordenada jj del estado nuevo depende de la jj del viejo y de ninguna otra por esta vía: no hay matriz que las mezcle. Derivando por el camino directo, y dejando quietas las compuertas y el candidato —volveré sobre eso—,

ht,jht1,j=1zt,j,\frac{\partial h_{t,j}}{\partial h_{t-1,j}} = 1 - z_{t,j},

un número, no una matriz. Encadénalo hacia atrás hasta un paso lejano kk, igual que la lección sobre el gradiente que se desvanece encadenaba la recurrencia, y lo que llega por la vía directa es un producto de esos números:

hk,j=(t=k+1T(1zt,j))hT,j+(lo que entra por las compuertas).\frac{\partial \ell}{\partial h_{k,j}} = \left(\prod_{t=k+1}^{T}\left(1 - z_{t,j}\right)\right)\frac{\partial \ell}{\partial h_{T,j}} + \left(\text{lo que entra por las compuertas}\right).

Es, letra por letra, el ft,j\prod f_{t,j} de la LSTM con 1zt,j1 - z_{t,j} en el sitio de ft,jf_{t,j}: la compuerta de actualización cerrada (zt,j0z_{t,j} \to 0) hace de compuerta de olvido abierta. Y vale lo mismo que allí —1zt,j=11 - z_{t,j} = 1 es alcanzable y estable, porque una sigmoide saturada se queda plana en el extremo—, así que una coordenada que la red aprenda a no actualizar multiplica por 11 tantos pasos como haga falta. Frente al (γσmax)Tk\left(\gamma\,\sigma_{\max}\right)^{T-k} de la lección sobre el gradiente que se desvanece, la base ha dejado de ser una propiedad de una matriz compartida por todo para ser un número por coordenada y por paso que la red calcula.

Ver qué se ha dejado fuera al «dejar quietas las compuertas»

zt\mathbf{z}_t, rt\mathbf{r}_t y h~t\tilde{\mathbf{h}}_t se calculan a partir de ht1\mathbf{h}_{t-1}, así que la derivada completa suma, al 1zt,j1 - z_{t,j} del camino directo, los caminos que pasan por las dos compuertas y por el candidato. Todos ellos cruzan una Wh\mathbf{W}_{h\ast} y una máscara de activación, y se desvanecen igual que los de la lección sobre el gradiente que se desvanece.

Que se desvanezcan no rompe nada, y es la misma lectura que en la LSTM. Esos caminos ajustan cómo se decide en cada paso, y una decisión de este paso se toma con lo que hay cerca; el que tiene que llegar de lejos es el de la vía directa, y ése no cruza ninguna matriz. El recorte no elimina el producto de la lección sobre el gradiente que se desvanece: lo aparta del camino largo y lo deja en los cortos.

Dos concesiones cierran la cuenta, las mismas que en la LSTM. Si la red aprende zt,jz_{t,j} alrededor de 0.30.3, el producto se hunde igual que en la RNN, sólo que porque la red lo decide y no porque la arquitectura lo imponga. Y la explosión sigue siendo posible por los caminos del recuadro, así que la GRU se entrena con recorte del gradiente igual que todo lo anterior. La vía aditiva no regala una garantía; hace alcanzable una opción que antes no lo era.

La GRU, paso a paso, en NumPy

Una celda. Transcribe el recuadro de cuatro líneas y lo corre sobre una secuencia corta con dos sesgos distintos de la compuerta de actualización, para ver cuánto de h1\mathbf{h}_1 llega a h6\mathbf{h}_6 por la vía directa.

Los pesos me los da un generador con semilla fija y no están entrenados, así que las compuertas todavía no significan nada. Lo que sí he puesto a mano es el sesgo bz\mathbf{b}_z: recuerda que conservar es z0z \to 0, así que un sesgo negativo es lo que empuja a la red a recordar —el reflejo del bf\mathbf{b}_f positivo de la LSTM—. Mira la última línea de cada bloque.

import numpy as np

d_model, d_h, T = 4, 3, 6
X = np.random.default_rng(1).normal(size=(T, d_model)) * 0.5


def sigmoide(z):
return 1.0 / (1.0 + np.exp(-z))


def pesos(sesgo_z):
rng = np.random.default_rng(0) # los mismos pesos en los dos casos
return {k: (rng.normal(size=(d_h, d_model)) * 0.5,
rng.normal(size=(d_h, d_h)) * 0.5,
np.full(d_h, b))
for k, b in [("z", sesgo_z), ("r", 0.0), ("h", 0.0)]}


def paso(P, h, x):
def entra(k, estado):
Wx, Wh, b = P[k]
return Wx @ x + Wh @ estado + b
z = sigmoide(entra("z", h))
r = sigmoide(entra("r", h))
cand = np.tanh(entra("h", r * h)) # el candidato lee r (*) h_{t-1}
h = (1 - z) * h + z * cand # la via aditiva, atada
return h, z, r


for sesgo in [0.0, -3.0]:
P = pesos(sesgo)
h, guardados = np.zeros(d_h), []
print("b_z = %.0f" % sesgo)
print(" t z_t r_t h_t")
for t, x in enumerate(X, 1):
h, z, r = paso(P, h, x)
guardados.append(1 - z)
print(" %d %s %s %s" % (t, np.round(z, 2), np.round(r, 2), np.round(h, 2)))
print(" de h_1 sobrevive en h_6:", np.round(np.prod(guardados[1:], axis=0), 3))
numpy

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

Las dos compuertas viven en (0,1)(0,1) y ninguna vale lo mismo en dos coordenadas. Compara los dos bloques por la última línea. Con bz=0\mathbf{b}_z = 0 la de actualización ronda 0.50.5, se reemplaza medio estado en cada paso y de h1\mathbf{h}_1 llega a h6\mathbf{h}_6 un 3%3\,\% escaso; con bz=3\mathbf{b}_z = -3 ronda 0.050.05, casi no se actualiza nada, y sobrevive un 77%77\,\%, con los mismos pesos y la misma secuencia. De ahí la costumbre de inicializar bz\mathbf{b}_z negativo, por lo mismo que bf\mathbf{b}_f positivo en la LSTM: sale más barato que la red aprenda a escribir que aprenda a recordar. Y fíjate en el precio del atajo: con bz=3\mathbf{b}_z = -3 el estado casi no se mueve —ht\mathbf{h}_t se queda en centésimas—, porque conservar y escribir son el mismo mando, y pedirle que conserve es pedirle que apenas escriba.

Comprueba tu intuición

Cuatro preguntas: qué ata y qué suelta la compuerta de actualización, qué hace la de reset sin tocar la memoria, cuánto cuesta la arquitectura, y qué se lleva la GRU de la LSTM y qué deja por el camino.

¿Qué gana y qué pierde la GRU al sustituir las compuertas de olvido y de entrada de la LSTM por una sola, la de actualización zt\mathbf{z}_t?

La compuerta de reset rt\mathbf{r}_t multiplica al estado anterior dentro del candidato: h~t=tanh(Wxhxt+Whh(rtht1)+bh)\tilde{\mathbf{h}}_t = \tanh\left(\mathbf{W}_{xh}\mathbf{x}_t + \mathbf{W}_{hh}\left(\mathbf{r}_t \odot \mathbf{h}_{t-1}\right) + \mathbf{b}_h\right). Con rt0\mathbf{r}_t \approx \mathbf{0} en una coordenada, ¿qué ocurre?

Una GRU con un estado de dh=64d_h = 64 coordenadas sobre embeddings de dmodel=128d_{\text{model}} = 128, las mismas dimensiones de la LSTM de la lección anterior. ¿Cuántos pesos suman sus seis matrices y sus tres sesgos?

Se acepta un margen de ±0.

Marca todo lo que sea cierto de la GRU frente a la LSTM.

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

Y el paso completo escrito por ti, que es la recurrencia de la RNN vanilla envuelta en dos compuertas.

Escribe un paso de GRU: paso_gru(h, x, P) recibe el estado anterior, el vector del token de turno, y un diccionario P con los pesos de las tres piezas —P["z"], P["r"] y P["h"], cada una la tripleta (Wx, Wh, b)—. Devuelve h_nuevo.

Las dos compuertas son sigmoides sobre el token y el estado anterior; el candidato es un tanh\tanh que lee rtht1\mathbf{r}_t \odot \mathbf{h}_{t-1} en lugar de ht1\mathbf{h}_{t-1} a secas, y el estado nuevo es la mezcla (1zt)ht1+zth~t(1-\mathbf{z}_t)\odot\mathbf{h}_{t-1} + \mathbf{z}_t\odot\tilde{\mathbf{h}}_t. Calcula las dos compuertas con el estado que entra, no con el que vas a devolver. No escribas en los arrays que recibes.

La primera comprobación descarga el intérprete de Python (~15 MB); después queda en la caché del navegador. Este desafío se resuelve mejor con un teclado físico: en el móvil puedes leerlo y volver luego.


Con la GRU y la LSTM, el bloque tiene ya dos maneras de darle memoria a una red, y las dos descansan en el mismo hallazgo: una vía por la que el estado avanza sumándose, sin cruzar una matriz, que es lo único que deja al gradiente volver desde lejos. Pero saber construir la celda no es lo mismo que haber construido algo con ella. En todo el bloque no ha corrido todavía una red que lea un texto de principio a fin y produzca lenguaje propio: se ha derivado la maquinaria sin encenderla.

Encenderla es la siguiente lección, el proyecto de un modelo de lenguaje a nivel de carácter. Junta la RNN, la BPTT (backpropagation through time) y una de estas celdas en una red que se entrena sobre unos pocos miles de caracteres y después genera, carácter a carácter, un texto nuevo. Saldrá malo —lo diré sin adornos cuando salga—, correrá en el navegador en unos segundos, y será la primera vez en el curso que algo escribe.

Para profundizar2 fuentes · 2 papers

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.