El gradiente que se desvanece

El gradiente que se desvanece

28 min de lectura

La lección anterior, sobre backpropagation through time (BPTT), cerró con un gradiente exacto y un cabo suelto. Exacto porque el sondeo numérico lo confirmó casilla por casilla; suelto porque la forma de calcularlo repite un mismo gesto sin que nadie mirara sus consecuencias: para arrastrar el error de una posición a la de más atrás hay que multiplicarlo por la matriz recurrente Whh\mathbf{W}_{hh}, y en una secuencia de cuarenta tokens ese producto se repite casi cuarenta veces. Esta lección mide qué le hace al gradiente repetir la misma multiplicación tantas veces seguidas, y encuentra que todo se decide en un único número asociado a Whh\mathbf{W}_{hh}.

Y no es una curiosidad numérica: en juego está la única razón para preferir una RNN a la red del principio del bloque. Toma una reseña que arranca con Dudo que a alguien le sirva y, treinta palabras más tarde, cierra con su valoración. Ese Dudo inicial invierte el signo de todo lo que viene después, pero la red solo puede aprender esa conexión si el gradiente calculado en la última palabra consigue volver hasta la primera. Si por el camino se apaga, la red jamás registrará que el principio importaba, y la memoria que la definía se quedará, en la práctica, en memoria de lo reciente.

Antes de mirar la matriz entera, quítale coordenadas hasta que quede una sola. Con dh=1d_h = 1 el estado es un número, Whh\mathbf{W}_{hh} es un número ww, y arrastrar el error un paso hacia atrás es multiplicarlo por ww. Repetir el paso dd veces lo multiplica por wdw^{d}.

Y wdw^{d} no tiene término medio. Si w<1\lvert w \rvert < 1 se hunde: con w=0.5w = 0.5, cuarenta pasos dejan el error en 0.5409×10130.5^{40} \approx 9 \times 10^{-13}, doce órdenes de magnitud por debajo de donde salió. Si w>1\lvert w \rvert > 1 se dispara: 1.5401.1×1071.5^{40} \approx 1.1 \times 10^{7}. Solo w=1\lvert w \rvert = 1 mantiene el error en su tamaño, y es un filo de un solo punto —basta desviarse un pelo para caer a un lado o al otro—. Con más de una coordenada, ese papel de ww lo hereda una propiedad de la matriz; el explorable la llama radio espectral, y la formalización la define enseguida. Deslízala a un lado y a otro de 11 y verás las dos mitades del fenómeno: por debajo, la curva se desploma; por encima, se va por arriba.

El gradiente, en órdenes de magnitud, según cuántos pasos tenga que retroceder. Por debajo de un radio espectral de 1 se desvanece; por encima, explota; y cuanto más larga es la secuencia, más extremo el resultado —la recta se inclina, nunca se dobla—.

El producto que conecta dos posiciones

Volvamos a la recurrencia hacia atrás de la lección anterior y sigámosla más de un paso. Allí, el error de un paso salía del siguiente con dos gestos —transportar y enmascarar—:

δt=(1htht)(Whhδt+1).\boldsymbol{\delta}_t = \left(\mathbf{1} - \mathbf{h}_t \odot \mathbf{h}_t\right) \odot \left(\mathbf{W}_{hh}^{\top}\boldsymbol{\delta}_{t+1}\right).

El \odot con la pendiente del tanh\tanh es multiplicar por una matriz diagonal —diag ⁣(1htht)\text{diag}\!\left(\mathbf{1} - \mathbf{h}_t \odot \mathbf{h}_t\right), con esas pendientes en la diagonal y ceros fuera—, así que un paso de la vuelta multiplica δt+1\boldsymbol{\delta}_{t+1} por dos matrices: primero Whh\mathbf{W}_{hh}^{\top}, luego esa diagonal. Encadena la recurrencia desde el último paso hasta uno lejano kk y las dos se repiten, una pareja por posición:

δk=(t=kT1diag ⁣(1htht)Whh)δTRdh.\boldsymbol{\delta}_k = \left(\prod_{t=k}^{T-1} \text{diag}\!\left(\mathbf{1} - \mathbf{h}_t \odot \mathbf{h}_t\right)\mathbf{W}_{hh}^{\top}\right)\boldsymbol{\delta}_T \in \mathbb{R}^{d_h}.

El producto tiene TkT - k parejas, una por cada paso entre kk y TT —la distancia que el error recorre—. Y dentro de cada pareja, la misma Whh\mathbf{W}_{hh}^{\top} que en todas las demás: esto no es una capa distinta por posición, como lo eran los bloques Wt(1)\mathbf{W}^{(1)}_t de la primera capa concatenada de la lección sobre por qué el MLP falla con secuencias, sino una sola matriz elevada, en la práctica, a la potencia de la distancia. Ese es el gesto que la lección anterior repetía sin mirarlo.

Por qué decae o explota

Un producto de matrices es difícil de mirar; su tamaño, no. Toma normas a los dos lados y usa que la norma de un producto no supera el producto de las normas:

δkδTt=kT1diag ⁣(1htht)Whh.\lVert \boldsymbol{\delta}_k \rVert \le \lVert \boldsymbol{\delta}_T \rVert \prod_{t=k}^{T-1} \left\lVert \text{diag}\!\left(\mathbf{1} - \mathbf{h}_t \odot \mathbf{h}_t\right)\right\rVert \left\lVert \mathbf{W}_{hh}^{\top}\right\rVert.

Los dos factores de cada pareja tienen un techo. El de la diagonal es la pendiente del tanh\tanh, que nunca pasa de 11 y en cuanto el estado se satura se acerca a 00 —son las bandas rojas que el explorador de activaciones de la lección sobre funciones de activación ya te dejó ver—; llámalo γ1\gamma \le 1. El de Whh\mathbf{W}_{hh}^{\top} es el único hecho nuevo que necesitas aquí: una matriz estira la longitud de un vector como mucho por su mayor valor singular σmax\sigma_{\max}, así que multiplicar por ella dd veces la estira, a lo sumo, por σmaxd\sigma_{\max}^{d}. Junta las dos cotas sobre las TkT - k parejas:

δk(γσmax)TkδT.\lVert \boldsymbol{\delta}_k \rVert \le \left(\gamma\,\sigma_{\max}\right)^{T-k}\lVert \boldsymbol{\delta}_T \rVert.

Ahí está el exponente. El error que llega a la posición kk es, como mucho, el del final multiplicado por un número elevado a la distancia. Si ese número es menor que 11, el gradiente se desvanece —cae en picado, y sin remedio, porque γ1\gamma \le 1 solo puede empujarlo más abajo, nunca rescatarlo—. Si es mayor que 11, explota.

Cuál de los dos lados te toca lo decide una propiedad de Whh\mathbf{W}_{hh}: su radio espectral ρ(Whh)\rho(\mathbf{W}_{hh}) —el mayor de los módulos de sus valores propios—, que marca el ritmo del producto a la larga. Por debajo de 11 el gradiente se hunde, por encima se dispara, y entre las dos cosas no hay más que ese filo. No necesitas calcular valores propios en esta lección; lo único que hace falta es saber que ese número existe, que gobierna el exponente y que es el que el explorable te dejó mover.

Ver por qué el ritmo lo marca el radio espectral y no el valor singular

La cota (γσmax)Tk\left(\gamma\,\sigma_{\max}\right)^{T-k} es un techo, y suele sobrar: un vector concreto rara vez se estira el máximo posible en cada paso, y σmax\sigma_{\max} puede ser bastante mayor que ρ\rho. Lo que sí se cumple a la larga es que (Whh)d1/d\left\lVert \left(\mathbf{W}_{hh}^{\top}\right)^{d}\right\rVert^{1/d} tiende al radio espectral ρ(Whh)\rho(\mathbf{W}_{hh}) cuando dd crece. Por eso ρ<1\rho < 1 garantiza el desvanecimiento y ρ>1\rho > 1 la explosión, aunque para una distancia corta la cota con σmax\sigma_{\max} vaya por delante del ritmo real.

Verlo y domarlo en NumPy

Dos celdas. La primera reduce todo al caso de una coordenada y mira el exponente crudo; la segunda ataca la mitad del problema que sí tiene un parche barato.

El caso escalar no es un juguete: es la fórmula de arriba con dh=1d_h = 1, donde σmax\sigma_{\max}, ρ\rho y w\lvert w \rvert son el mismo número. Ejecútala y lee la fila del medio —w=1w = 1— contra las otras dos.

# W_hh de una sola coordenada es un numero w. Arrastrar el error d pasos hacia
# atras lo multiplica por w**d. La mascara del tanh (pendiente <= 1) solo lo
# encoge mas, asi que w = 1 ya es el caso mas favorable.
print(" w d=10 d=20 d=40")
for w in [0.5, 1.0, 1.5]:
print("%.1f %9.2e %9.2e %9.2e" % (w, w ** 10, w ** 20, w ** 40))

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

La primera fila se hunde, la última se dispara y la del medio no se mueve, tal como decía la intuición. Con w=0.5w = 0.5 el error llega a los cuarenta pasos convertido en 9×10139 \times 10^{-13} de lo que era, y ningún descenso de gradiente aprende de una señal tan pequeña. Y esto es el mejor caso: la máscara del tanh\tanh, en cuanto el estado se satura, lo empuja todavía más abajo.

La explosión, en cambio, tiene una respuesta directa. Si el problema es que la norma del gradiente crece sin freno, ponle un techo: fija un umbral θ\theta y, cuando g\lVert \mathbf{g} \rVert lo pase, reescala el gradiente para devolverlo justo a ese tamaño. Es el recorte del gradiente (gradient clipping):

ggmin ⁣(1, θg).\mathbf{g} \leftarrow \mathbf{g}\cdot\min\!\left(1,\ \frac{\theta}{\lVert \mathbf{g} \rVert}\right).

El min\min con 11 es la clave: si el gradiente ya es pequeño, el factor vale 11 y no lo toca; solo actúa sobre el que se ha pasado, y lo hace multiplicándolo por un número positivo, que cambia su longitud pero no su dirección. Compruébalo.

import numpy as np

rng = np.random.default_rng(1)
theta = 5.0 # el techo que le ponemos a la norma


def recorta(g):
n = np.linalg.norm(g)
return g * min(1.0, theta / n), n # factor <= 1: encoge o deja igual, no gira


for nombre, g in [("normal ", rng.normal(size=100) * 0.3),
("explotado", rng.normal(size=100) * 3.0)]:
g_rec, n = recorta(g)
coseno = float(g @ g_rec) / (n * np.linalg.norm(g_rec))
print("%s ||g|| = %6.2f -> ||recortado|| = %.2f (coseno con g: %.4f)"
% (nombre, n, np.linalg.norm(g_rec), coseno))
numpy

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

El gradiente normal pasa intacto; el explotado, con norma casi 3030, sale recortado a 5.005.00 exacto, y el coseno con el original es 11 —misma dirección, otro tamaño—. El recorte salva el entrenamiento cuando el gradiente explota y no cuesta casi nada. Lo que no hace —y conviene decirlo en voz alta— es nada por el desvanecimiento: un gradiente que ya llegó valiendo 101310^{-13} tiene norma minúscula, el min\min vale 11, y se queda como estaba.

Comprueba tu intuición

Cinco preguntas: qué desvanece o explota el gradiente, qué arregla el recorte y qué no, cómo se ve el exponente en un número, si la máscara del tanh\tanh puede salvar la situación y qué se pierde cuando no.

¿Qué hace que el gradiente de una RNN se desvanezca o explote al retroceder por una secuencia larga?

El recorte del gradiente reescala ggmin ⁣(1, θ/g)\mathbf{g} \leftarrow \mathbf{g}\cdot\min\!\left(1,\ \theta/\lVert\mathbf{g}\rVert\right) cuando su norma pasa de un umbral θ\theta. ¿Qué problema resuelve y cuál no?

El caso de una sola coordenada: Whh\mathbf{W}_{hh} es un número ww, y a distancia dd el error se ha multiplicado por wdw^{d}. ¿Qué imprime?

w = 0.5
for d in [0, 4, 8]:
    print(round(w ** d, 4))
 

La máscara 1htht\mathbf{1} - \mathbf{h}_t \odot \mathbf{h}_t que aparece en cada paso de la vuelta es la pendiente del tanh\tanh. ¿Puede rescatar un gradiente que se desvanece?

¿Qué pierde una RNN cuando el gradiente que conecta el final de la secuencia con el principio se ha desvanecido?


El recorte le pone un techo a la explosión, pero deja el otro lado del problema donde estaba. Un gradiente que se ha multiplicado camino a cero no se recupera multiplicándolo por 11, y ninguna elección de Whh\mathbf{W}_{hh} arregla las dos cosas a la vez: bajar su radio espectral para no explotar es justo lo que garantiza el desvanecimiento. El problema no está en los pesos, sino en el camino —cualquier ruta que multiplique por la misma matriz en cada paso acaba en uno de los dos extremos—.

La salida es cambiar el camino. Hace falta una vía por la que el estado avance de un paso al siguiente sumando, en lugar de volver a pasar por Whh\mathbf{W}_{hh}: sobre una suma el gradiente fluye sin multiplicarse una vez por posición, y así ni se hunde ni se dispara. Esa vía aditiva es el estado de celda de la LSTM (long short-term memory), y construirla —una compuerta cada vez, cada una arreglando un fallo concreto de la RNN vanilla— es la siguiente lección, sobre la LSTM: memoria con compuertas.

Para profundizar2 fuentes · 1 paper, 1 libro

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.

  • On the difficulty of training Recurrent Neural Networks
    paperPascanu, Mikolov y Bengio, 2013arXiv:1211.5063EN

    De aquí salen las dos piezas de la lección: el radio espectral de W_hh como frontera entre desvanecerse y explotar, y el recorte que limita la norma del gradiente a un umbral θ. Lo justifica con más aparato; las compuertas no las toca.

  • Gradient Flow in Recurrent Nets: the Difficulty of Learning Long-Term Dependencies
    libroHochreiter, Bengio, Frasconi y Schmidhuber, 2001A Field Guide to Dynamical Recurrent NetworksEN

    Su §2 demuestra lo que aquí ves en una gráfica: el error que retrocede por la secuencia decae o crece de forma exponencial en el número de pasos. Reúne las pruebas de 1991 y 1994; se queda en el diagnóstico, sin proponer arreglo.