LSTM: memoria con compuertas

LSTM: memoria con compuertas

30 min read

La lección anterior, sobre el gradiente que se desvanece, dejó el diagnóstico cerrado: lo que hunde o dispara la señal de vuelta no es el valor de los pesos sino la forma del camino, y mientras cada paso hacia atrás cruce la misma matriz, el resultado será una potencia de la distancia. También dejó apuntada la salida —una vía por la que el estado avance sumando— y el nombre de la red que la lleva, la LSTM (long short-term memory). Nombrarla no es construirla: una memoria que sólo suma tiene un defecto tan grave como el que arregla, y es que nunca borra nada. Falta lo difícil, que es quién decide, paso a paso, qué se conserva, qué se escribe encima y qué de lo guardado sale a la luz.

Que son tres decisiones y no una se ve en una frase corriente. Toma Las cámaras que el ayuntamiento instaló el año pasado en la plaza no funcionan. Para conjugar funcionan hay que arrastrar el plural de cámaras doce posiciones, sin que ayuntamiento ni plaza —los dos en singular, los dos más cerca— lo pisen por el camino. Arrastrarlo intacto es una cosa; meterlo en la memoria cuando aparece cámaras y no en cada palabra que pasa es otra; y tenerlo guardado durante esas doce posiciones sin que gobierne lo que la red responde en ellas es una tercera. La RNN (recurrent neural network) vanilla hace las tres con el mismo gesto, porque sólo tiene un gesto: reescribir entero ht\mathbf{h}_t en cada paso. Por eso no puede hacer ninguna bien.

Separar los tres trabajos empieza por separar los dos estados. Junto al ht\mathbf{h}_t que ya conoces aparece un segundo vector, el estado de celda ctRdh\mathbf{c}_t \in \mathbb{R}^{d_h}, que es la memoria propiamente dicha, y su forma de avanzar es lo único verdaderamente nuevo de esta lección. Todo lo demás de la LSTM existe para gobernar ese vector.

Tres compuertas, cada una por un fallo de la anterior

Empecemos por la vía aditiva sola y vayamos añadiendo lo que le falte. Llamemos candidato al vector que el paso tt propone guardar, una capa corriente a partir del token y del estado anterior:

c~t=tanh(Wxcxt+Whcht1+bc)Rdh,\tilde{\mathbf{c}}_t = \tanh\left(\mathbf{W}_{xc}\mathbf{x}_t + \mathbf{W}_{hc}\mathbf{h}_{t-1} + \mathbf{b}_c\right) \in \mathbb{R}^{d_h},

y que la memoria lo vaya acumulando: ct=ct1+c~t\mathbf{c}_t = \mathbf{c}_{t-1} + \tilde{\mathbf{c}}_t. Para el gradiente esto es perfecto —la derivada de ct\mathbf{c}_t respecto de ct1\mathbf{c}_{t-1} es la identidad, sin matriz que aplicar dd veces—, y para la memoria es inservible por dos motivos. Nunca olvida: lo que entró en la coordenada 3 leyendo cámaras sigue ahí trescientos tokens después, en un párrafo que habla de otra cosa. Y crece sin freno: cada término vive en (1,1)(-1, 1) y se suman TT, así que una coordenada puede acabar valiendo 4040, donde el tanh\tanh que la lea está plano —la saturación que el explorador de activaciones de la lección sobre funciones de activación ya te dejó ver—.

Los dos fallos son el mismo: la suma no tiene forma de decir no. Dásela con un factor por coordenada, entre 00 y 11, que la red calcule en cada paso a partir de lo que está leyendo. Ésa es la compuerta de olvido:

ft=σ(Wxfxt+Whfht1+bf)(0,1)dh,ct=ftct1+c~t.\mathbf{f}_t = \sigma\left(\mathbf{W}_{xf}\mathbf{x}_t + \mathbf{W}_{hf}\mathbf{h}_{t-1} + \mathbf{b}_f\right) \in (0,1)^{d_h}, \qquad \mathbf{c}_t = \mathbf{f}_t \odot \mathbf{c}_{t-1} + \tilde{\mathbf{c}}_t.

La activación es la sigmoide logística y no la tanh\tanh del candidato, por lo mismo que la capa de salida de la lección sobre la RNN vanilla: hace falta una fracción de lo que había, y las fracciones van de 00 a 11. El \odot pide leerse igual de despacio. Multiplicar coordenada a coordenada deja que la 3 se conserve entera mientras la 7 se borra en el mismo paso; una matriz mezclaría las dos, y mezclar es lo que menos le conviene a una memoria.

Queda el otro lado. Esta versión escribe c~t\tilde{\mathbf{c}}_t en cada paso, valga lo que valga el token: que y del escriben tanto como cámaras, y lo escrito sólo se deshace cerrando ft\mathbf{f}_t, que borra de paso lo que esa coordenada guardara de antes. Poner una condición a la escritura pide otro factor igual, la compuerta de entrada:

it=σ(Wxixt+Whiht1+bi),ct=ftct1+itc~t.\mathbf{i}_t = \sigma\left(\mathbf{W}_{xi}\mathbf{x}_t + \mathbf{W}_{hi}\mathbf{h}_{t-1} + \mathbf{b}_i\right), \qquad \mathbf{c}_t = \mathbf{f}_t \odot \mathbf{c}_{t-1} + \mathbf{i}_t \odot \tilde{\mathbf{c}}_t.

Y falta decir qué ve el resto de la red. Si el estado oculto fuera ht=tanh(ct)\mathbf{h}_t = \tanh(\mathbf{c}_t) —toda la memoria, siempre—, el plural de cámaras gobernaría la respuesta en las doce posiciones del inciso, cuando lo que hace falta es que espere. Guardar algo para luego y no usarlo ahora son dos cosas distintas, así que hay una tercera compuerta, la de salida, y con ella el estado oculto es una vista parcial de la memoria:

ot=σ(Wxoxt+Whoht1+bo),ht=ottanh(ct).\mathbf{o}_t = \sigma\left(\mathbf{W}_{xo}\mathbf{x}_t + \mathbf{W}_{ho}\mathbf{h}_{t-1} + \mathbf{b}_o\right), \qquad \mathbf{h}_t = \mathbf{o}_t \odot \tanh\left(\mathbf{c}_t\right).

Ahí está lo que la RNN vanilla no podía hacer: su ht\mathbf{h}_t tenía dos trabajos a la vez, recordar y enseñar; la LSTM los reparte, y ct\mathbf{c}_t recuerda mientras ht\mathbf{h}_t enseña.

Con las cuatro piezas juntas —tres compuertas y un candidato— el paso completo queda así, con h0=c0=0\mathbf{h}_0 = \mathbf{c}_0 = \mathbf{0} por la misma razón que en la RNN:

ft=σ(Wxfxt+Whfht1+bf),it=σ(Wxixt+Whiht1+bi),ot=σ(Wxoxt+Whoht1+bo),c~t=tanh(Wxcxt+Whcht1+bc),ct=ftct1+itc~t,ht=ottanh(ct).\begin{aligned} \mathbf{f}_t &= \sigma\left(\mathbf{W}_{xf}\mathbf{x}_t + \mathbf{W}_{hf}\mathbf{h}_{t-1} + \mathbf{b}_f\right), & \mathbf{i}_t &= \sigma\left(\mathbf{W}_{xi}\mathbf{x}_t + \mathbf{W}_{hi}\mathbf{h}_{t-1} + \mathbf{b}_i\right), \\ \mathbf{o}_t &= \sigma\left(\mathbf{W}_{xo}\mathbf{x}_t + \mathbf{W}_{ho}\mathbf{h}_{t-1} + \mathbf{b}_o\right), & \tilde{\mathbf{c}}_t &= \tanh\left(\mathbf{W}_{xc}\mathbf{x}_t + \mathbf{W}_{hc}\mathbf{h}_{t-1} + \mathbf{b}_c\right), \\ \mathbf{c}_t &= \mathbf{f}_t \odot \mathbf{c}_{t-1} + \mathbf{i}_t \odot \tilde{\mathbf{c}}_t, & \mathbf{h}_t &= \mathbf{o}_t \odot \tanh\left(\mathbf{c}_t\right) \end{aligned}.

Cuatro líneas idénticas y dos que no. Las de arriba son la recurrencia de la lección sobre la RNN vanilla copiada cuatro veces, con WxRdh×dmodel\mathbf{W}_{x\ast} \in \mathbb{R}^{d_h \times d_{\text{model}}}, WhRdh×dh\mathbf{W}_{h\ast} \in \mathbb{R}^{d_h \times d_h} y bRdh\mathbf{b}_{\ast} \in \mathbb{R}^{d_h} en todas ellas, y ahí está el precio: la LSTM cuesta 4(dhdh+dhdmodel+dh)4\left(d_h \cdot d_h + d_h \cdot d_{\text{model}} + d_h\right) pesos, cuatro veces la RNN vanilla. Lo que no cuesta es longitud —tampoco aquí aparece TT en ninguna forma—.

Ya está toda la maquinaria; muévela antes de derivarla. Recorre la secuencia y mira las tres compuertas abrirse y cerrarse en cada posición; luego sube el sesgo bf\mathbf{b}_f de la de olvido y observa la última fila, lo que sobrevive de c1\mathbf{c}_1.

Sube el sesgo de la compuerta de olvido y ésta se abre hacia 1; recorre la secuencia y fíjate en la última fila, lo que sobrevive de c₁. Con el sesgo alto la memoria se conserva casi entera; con el sesgo bajo se filtra hasta casi nada, con los mismos pesos y la misma frase.

Por qué el gradiente sobrevive a la vía aditiva

Un paso de una LSTM. Por arriba lo cruza de lado a lado una línea verde, la del estado de celda, que sólo se encuentra dos círculos por el camino: uno de producto y otro de suma. Debajo, cuatro cajitas rotuladas f, i, c con tilde y o, alimentadas todas por una misma línea de puntos que recoge el estado oculto anterior y el vector del token. Las cajitas f, i y c con tilde suben hasta los dos círculos de la línea de arriba. A la derecha, una rama verde baja de la línea de arriba, pasa por un tanh y se encuentra con la cajita o en un tercer círculo de producto, del que sale la única línea rotulada como estado oculto nuevo. No hay ninguna línea que una el estado oculto anterior con el nuevo.
El estado de celda cruza el paso de lado a lado sin multiplicarse por ninguna matriz: sólo una compuerta lo escala y otra le suma. Ésa es la vía por la que el gradiente podrá volver, y el estado oculto cuelga de ella en vez de sustituirla.

Que la celda «suma» ya se ha dicho tres veces; ahora hay que cobrarlo. Mira la penúltima línea coordenada a coordenada, que es donde se ve lo que el \odot significa:

ct,j=ft,jct1,j+it,jc~t,j.c_{t,j} = f_{t,j}\,c_{t-1,j} + i_{t,j}\,\tilde{c}_{t,j}.

La coordenada jj de la celda nueva depende de la coordenada jj de la vieja y de ninguna otra: por esta vía las coordenadas no se mezclan, porque no hay matriz que las mezcle. Derivando, y dejando quietas las compuertas —volveré sobre esto—, el camino directo de un paso al anterior vale

ct,jct1,j=ft,j,\frac{\partial c_{t,j}}{\partial c_{t-1,j}} = f_{t,j},

un número, no una matriz. Encadénalo desde el final hasta un paso lejano kk, igual que la lección anterior encadenaba la recurrencia hacia atrás, y lo que llega a la celda de kk por esa vía es

ck,j=(t=k+1Tft,j)cT,j+(lo que entra por las compuertas).\frac{\partial \ell}{\partial c_{k,j}} = \left(\prod_{t=k+1}^{T} f_{t,j}\right)\frac{\partial \ell}{\partial c_{T,j}} + \left(\text{lo que entra por las compuertas}\right).

Ponlo al lado del (γσmax)Tk\left(\gamma\,\sigma_{\max}\right)^{T-k} de la lección anterior y cuenta las diferencias, que son tres. La base ya no es una propiedad de una matriz compartida por todo, sino un número por coordenada y por paso que la red calcula a partir de lo que está leyendo. Por esta vía tampoco hay máscara del tanh\tanh que muerda en cada paso —ningún γ1\gamma \le 1 que sólo pueda empeorar las cosas—, porque la celda cruza el paso con un producto y una suma y nada más. Y la tercera remata: ft,j=1f_{t,j} = 1 es alcanzable y además estable, porque una sigmoide saturada se queda en 11 con pendiente casi nula, así que una coordenada que la red aprenda a mantener abierta multiplica por 11 tantas veces como haga falta. El ρ=1\rho = 1 de la RNN era un filo compartido por todas las coordenadas; esto es una decisión de cada una.

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

ft\mathbf{f}_t, it\mathbf{i}_t y c~t\tilde{\mathbf{c}}_t se calculan a partir de ht1\mathbf{h}_{t-1}, y ht1=ot1tanh(ct1)\mathbf{h}_{t-1} = \mathbf{o}_{t-1} \odot \tanh(\mathbf{c}_{t-1}) depende de ct1\mathbf{c}_{t-1}. Así que la derivada completa suma, al ft,jf_{t,j} de arriba, los caminos que pasan por las tres compuertas y por el candidato, y todos ellos cruzan una Wh\mathbf{W}_{h\ast} y una máscara de activación: se comportan como los de la lección anterior, y se desvanecen igual.

Que se desvanezcan no rompe nada, y ésa es la parte que merece pensarse. 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 aditiva, y ése no cruza ninguna matriz. La LSTM no elimina el producto de la lección anterior: lo aparta del camino largo y lo deja en los cortos.

Dos concesiones cierran la cuenta, porque sin ellas esto se lee como una garantía. Si la red aprende ft,jf_{t,j} alrededor de 0.70.7, el producto se hunde igual que antes; la diferencia es que se hunde porque la red lo ha decidido y no porque la arquitectura lo imponga. Y la explosión sigue siendo posible por los caminos del recuadro de arriba, así que el recorte del gradiente de la lección anterior se sigue usando tal cual.

La LSTM, paso a paso, en NumPy

Dos celdas. La primera transcribe el recuadro de seis líneas y lo ejecuta sobre una secuencia corta; la segunda mide lo que sobrevive de un empujón dado al principio de una larga.

Los pesos me los da un generador con semilla fija y no están entrenados, así que las compuertas no significan nada todavía. Lo que sí he puesto a mano es el sesgo de la de olvido, y la celda lo corre con dos valores distintos: es la forma más corta de ver que bf\mathbf{b}_f decide cuánto recuerda la red antes de aprender nada. 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_olvido):
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 [("f", sesgo_olvido), ("i", 0.0), ("o", 0.0), ("c", 0.0)]}


def paso(P, h, c, x):
def entra(k):
Wx, Wh, b = P[k]
return Wx @ x + Wh @ h + b
f, i, o = sigmoide(entra("f")), sigmoide(entra("i")), sigmoide(entra("o"))
c = f * c + i * np.tanh(entra("c")) # la vía aditiva, en una línea
return o * np.tanh(c), c, f, i, o


for sesgo in [1.0, 3.0]:
P = pesos(sesgo)
h, c, olvidos = np.zeros(d_h), np.zeros(d_h), []
print("b_f = %.0f" % sesgo)
print(" t f_t i_t o_t c_t")
for t, x in enumerate(X, 1):
h, c, f, i, o = paso(P, h, c, x)
olvidos.append(f)
print(" %d %s %s %s %s" % (t, np.round(f, 2), np.round(i, 2),
np.round(o, 2), np.round(c, 2)))
print(" de c_1 sobrevive en c_6:", np.round(np.prod(olvidos[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 tres compuertas viven en (0,1)(0,1) y ninguna vale lo mismo en dos coordenadas, que es todo lo que la fórmula prometía. Compara ahora los dos bloques. Con bf=1\mathbf{b}_f = 1 la de olvido se mueve alrededor de 0.70.7 y de c1\mathbf{c}_1 llega a c6\mathbf{c}_6 un 20%20\,\% escaso; con bf=3\mathbf{b}_f = 3 ronda 0.950.95 y llega un 78%78\,\%, con los mismos pesos y la misma frase. Cinco pasos bastan para que un sesgo separe una memoria de un colador, y de ahí la costumbre de inicializar bf\mathbf{b}_f positivo: sale más barato que la red aprenda a olvidar que aprenda a recordar. La primera fila de ct\mathbf{c}_t, en cambio, es idéntica en los dos: en t=1t = 1 la compuerta multiplica c0=0\mathbf{c}_0 = \mathbf{0}.

La segunda celda mide la vía de vuelta sin derivar nada: mueve el estado del primer paso una cantidad diminuta, deja correr treinta y nueve pasos más y compara. Es el sondeo numérico de la lección sobre BPTT (backpropagation through time), aplicado a la distancia en vez de a un peso. Las compuertas de este juguete leen sólo el token —he desconectado el estado a mano— para medir la vía aditiva y nada más.

import numpy as np

T, eps = 40, 1e-6
X = np.random.default_rng(2).normal(size=T) * 0.5 # una coordenada, T pasos


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


def rnn(empujon): # h_t = tanh(w h_{t-1} + u x_t)
h = 0.0
for t, x in enumerate(X):
h = np.tanh(0.9 * h + 0.5 * x)
if t == 0:
h += empujon # movemos el estado del paso 1
return h


def lstm(empujon, b_f): # c_t = f_t c_{t-1} + i_t c~_t
c = 0.0
for t, x in enumerate(X):
c = sigmoide(0.5 * x + b_f) * c + sigmoide(0.5 * x) * np.tanh(0.5 * x)
if t == 0:
c += empujon # movemos la celda del paso 1
return c


print("de un empujon dado en el paso 1, cuanto queda 39 pasos despues:")
print(" RNN, w = 0.9 %.2e" % abs((rnn(eps) - rnn(0.0)) / eps))
for b_f in [0.0, 3.0, 5.0]:
q = abs((lstm(eps, b_f) - lstm(0.0, b_f)) / eps)
print(" LSTM, b_f = %.0f %.2e (f_t alrededor de %.3f)" % (b_f, q, sigmoide(b_f)))
numpy

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

Cuatro números y tres lecturas. De la RNN llega 1.1×1031.1 \times 10^{-3}: 0.9390.0160.9^{39} \approx 0.016 lo pone la matriz y el resto —un orden de magnitud— la máscara del tanh\tanh, ese γ1\gamma \le 1 que la lección anterior dijo que sólo podía empeorar las cosas. La segunda línea es la LSTM saliendo peor que la RNN: con bf=0\mathbf{b}_f = 0 la compuerta ronda 0.50.5, el producto cae a 101110^{-11} y la memoria es peor que la que venía a arreglar. La vía aditiva no regala nada, sólo hace alcanzable lo que en la RNN no lo era. Las dos últimas líneas son ese «alcanzable»: con la compuerta en 0.950.95 sobrevive el 14%14\,\% y con ella en 0.9930.993, el 76%76\,\%, a la distancia donde la RNN llevaba tres ceros tras el punto.

Comprueba tu intuición

Cinco preguntas: por qué la vía aditiva cambia el producto, qué hace cada compuerta en un caso concreto, cuánto cuesta la arquitectura, cómo se ve la diferencia en un número y qué es lo que la LSTM sigue sin garantizar.

¿Por qué el estado de celda de una LSTM no condena al gradiente a la potencia exponencial de la lección anterior?

Una coordenada del estado de celda guarda que el sujeto de la frase era plural. Hay que conservarlo durante toda una oración de relativo, pero sin que gobierne la respuesta de la red en esas posiciones intermedias. ¿Qué hacen las compuertas de esa coordenada mientras dura el relativo?

Una LSTM con un estado de dh=64d_h = 64 coordenadas sobre embeddings de dmodel=128d_{\text{model}} = 128. ¿Cuántos pesos suman sus ocho matrices y sus cuatro sesgos?

A margin of ±0 is accepted.

Cuarenta pasos de distancia, y en cada uno el gradiente se multiplica por el mismo número: en la RNN por un factor que fija Whh\mathbf{W}_{hh}, en la celda de la LSTM por la compuerta de olvido. ¿Qué imprime?

for factor in [0.6, 0.99]:
    print(round(factor ** 40, 3))
 

Marca todo lo que la LSTM no garantiza.

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

Y el paso completo escrito por ti, que es la pieza sobre la que se monta todo lo que viene después.

Escribe un paso de LSTM: paso_lstm(h, c, x, P) recibe el estado oculto y el estado de celda anteriores, el vector del token de turno, y un diccionario P con los pesos de las cuatro piezas —P["f"], P["i"], P["o"] y P["c"], cada una la tripleta (Wx, Wh, b)—. Devuelve la pareja (h_nuevo, c_nuevo), en ese orden.

Las cuatro se calculan igual, Wxxt+Whht1+b\mathbf{W}_{x\ast}\mathbf{x}_t + \mathbf{W}_{h\ast}\mathbf{h}_{t-1} + \mathbf{b}_{\ast}; lo único que cambia es la activación: σ\sigma para las tres compuertas, tanh\tanh para el candidato. No escribas en los arrays que recibes.

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 LSTM resuelve el problema de la lección anterior y cobra un precio que la cuenta de pesos sólo dice a medias. Ocho matrices, dos estados que arrastrar y tres compuertas que ajustar: cada pieza llegó aquí arreglando un fallo concreto, sí, pero ninguna con una prueba de que hiciera falta. Y de al menos una se puede sospechar. La compuerta de entrada decide cuánto se escribe justo cuando la de olvido acaba de decidir cuánto se conserva, y las dos preguntas son la misma vista desde los dos lados: si vas a guardar lo nuevo, es porque lo viejo te importa menos.

Alguien tiró de ese hilo. Fundir esas dos compuertas en una, borrar la de salida junto con el estado de celda entero y quedarse con tres cuartas partes de los pesos —conservando casi todo lo que esta lección acaba de comprar— es la siguiente lección, sobre la GRU (gated recurrent unit): la versión simplificada.

Further reading3 sources · 2 papers, 1 article

Where this lesson comes from, and where to go next. None of it is needed to carry on with the course.

  • Long Short-Term Memory
    paperHochreiter y Schmidhuber, 1997Neural Computation 9(8)EN

    El artículo que la introdujo. Trae el estado de celda y el flujo constante de error que lo recorre —tu vía aditiva—, pero sólo con compuertas de entrada y salida: la de olvido la añaden Gers, Schmidhuber y Cummins en 2000.

  • Understanding LSTM Networks
    articleChristopher Olah, 2015colah.github.ioEN

    Recorre con diagramas las tres compuertas y el estado de celda como «cinta transportadora» —la figura de esta lección—, y de paso presenta la GRU. Dice por qué la vía aditiva ayuda al gradiente; no lo deriva, eso lo has hecho tú.

  • LSTM: A Search Space Odyssey
    paperGreff, Srivastava, Koutník, Steunebrink y Schmidhuber, 2017arXiv:1503.04069EN

    Ocho variantes de la LSTM en 5 400 entrenamientos: lo que de verdad importa es la compuerta de olvido y la activación de salida; acoplar entrada y olvido —lo que hará la GRU— casi no cambia nada. Ninguna variante supera a la estándar.