Backpropagation through time

Backpropagation through time

29 min de lectura

La lección anterior, sobre la RNN vanilla, se detuvo justo antes de entrenarla: la red lee cualquier longitud con dos matrices y un bucle, pero sus pesos siguen siendo los que un generador escribió al construirla, y no se ha movido ni uno. Moverlos es pedir el gradiente de la pérdida respecto de esos pesos, y al pedirlo aparece una particularidad que el bloque anterior no tenía. La red neuronal recurrente (recurrent neural network, RNN) no reparte su trabajo en una matriz por capa: Whh\mathbf{W}_{hh} es una sola matriz que el bucle usa en cada posición, enredada en la secuencia entera.

¿Cómo se deriva respecto de un peso que interviene en los TT pasos a la vez? El bloque anterior tiene las dos piezas para contestar. La regla de la cadena dice qué ocurre cuando un peso influye en la pérdida por varios sitios; backpropagation consigue todas las derivadas de una red con una sola pasada del final al principio, sin sondear un peso cada vez. Falta encajarlas en un bucle, y el resultado tiene nombre propio: backpropagation through time (BPTT). Lo que esta lección hace con él no es enunciarlo, es derivarlo entero y comprobar, casilla por casilla, que sus gradientes son los que un sondeo numérico confirma.

Desplegar el bucle deja la red como una fila de pasos, y sobre esa fila la vuelta se recorre a mano. La pérdida está al final, colgada del último estado; para llegar a un peso, su influencia tiene que volver por donde vino, del último estado al anterior y del anterior al de más atrás, cruzando en cada salto la misma Whh\mathbf{W}_{hh}. No hay otro camino de vuelta porque no hubo otro de ida.

Avanza paso a paso hacia atrás: cada paso deja su aporte en el mismo acumulador de forma d_h por d_h. Sube T y comprueba que se suman más términos sin que cambie la forma del gradiente; en el último paso el aporte al gradiente de Whh es cero porque h0 es cero, y al cambiar el acumulador a Wxh verás que el suyo no lo es.

El error que vuelve, un paso a la vez

Fijemos una secuencia —un documento— y derivemos para ella sola, así que la pérdida es \ell, la de un ejemplo, como en todo el bloque anterior desde la lección sobre funciones de pérdida; la media sobre el batch, L=1Bii\mathcal{L} = \frac{1}{B}\sum_i \ell_i, promedia estos gradientes por documento y es un nivel de más arriba que no toca nada de lo que sigue. La respuesta se lee del último estado, igual que en la lección anterior:

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

de modo que los pesos de la recurrencia llegan a la pérdida por un único sitio: a través de hT\mathbf{h}_T. Todo lo demás de la vuelta es reconstruir, paso a paso, cómo cada estado afecta a ese último.

Dale un nombre a lo que el bloque anterior dejó sin nombrar. La preactivación del paso tt —el argumento del tanh\tanh— es

pt=Whhht1+Wxhxt+bhRdh,ht=tanh(pt),\mathbf{p}_t = \mathbf{W}_{hh}\mathbf{h}_{t-1} + \mathbf{W}_{xh}\mathbf{x}_t + \mathbf{b}_h \in \mathbb{R}^{d_h}, \qquad \mathbf{h}_t = \tanh\left(\mathbf{p}_t\right),

y llamemos error del paso tt, y δtRdh\boldsymbol{\delta}_t \in \mathbb{R}^{d_h}, al gradiente de la pérdida respecto de esa preactivación, δt=pt\boldsymbol{\delta}_t = \nabla_{\mathbf{p}_t}\,\ell. Es exactamente el δ(l)\boldsymbol{\delta}^{(l)} de la lección sobre backpropagation —el error de una capa, en su preactivación—, sólo que el índice ya no cuenta capas sino pasos del tiempo.

Ese error se calcula del final al principio, y en dos movimientos por paso. El tanh\tanh se aplica coordenada a coordenada, así que pasar de ht\mathbf{h}_t a pt\mathbf{p}_t es multiplicar por su pendiente, y esa pendiente, por la identidad tanh=1tanh2\tanh^{\prime} = 1 - \tanh^{2} de la lección sobre funciones de activación, es 1htht\mathbf{1} - \mathbf{h}_t \odot \mathbf{h}_t: se lee directamente del estado que la ida ya guardó, sin conservar pt\mathbf{p}_t. Queda el otro factor, ht\nabla_{\mathbf{h}_t}\ell, y tiene dos casos según de dónde le llegue a ht\mathbf{h}_t la pérdida. En el último paso, sólo de la capa de salida; en los demás, sólo del paso siguiente:

δT=(1hThT)((y^y)Why),\boldsymbol{\delta}_T = \left(\mathbf{1} - \mathbf{h}_T \odot \mathbf{h}_T\right) \odot \left(\left(\hat{y} - y\right)\mathbf{W}_{hy}^{\top}\right), δt=(1htht)(Whhδt+1),t=T1,,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), \qquad t = T-1, \dots, 1.

El arranque, δT\boldsymbol{\delta}_T, es la capa de salida de la lección anterior derivada hacia atrás: la sigmoide con la entropía cruzada entrega y^y\hat{y} - y, y Why\mathbf{W}_{hy}^{\top} lo reparte entre las dhd_h coordenadas del estado. La recurrencia es lo demás, y son los dos movimientos de la lección sobre backpropagation sin ninguna variación: transporta el error con Whh\mathbf{W}_{hh}^{\top} y lo enmascara con la pendiente del tanh\tanh. Lo único distinto es que el transporte usa siempre la misma matriz —guárdalo, porque es la mitad de una lección más adelante, la del gradiente que se desvanece.

Ver el transporte índice a índice

Sin matrices por medio, y para un paso interior t<Tt < T, donde ht\mathbf{h}_t sólo alcanza la pérdida a través del paso siguiente. La coordenada jj de ht\nabla_{\mathbf{h}_t}\ell es la derivada de \ell respecto de ht,jh_{t,j}, y ht,jh_{t,j} entra en la preactivación pt+1\mathbf{p}_{t+1} por cada una de sus dhd_h coordenadas: es la suma sobre caminos de la lección sobre la regla de la cadena, con dhd_h caminos.

ht,j=i=1dhpt+1,ipt+1,iht,j=i=1dhδt+1,i(Whh)ij,\frac{\partial \ell}{\partial h_{t,j}} = \sum_{i=1}^{d_h} \frac{\partial \ell}{\partial p_{t+1,i}}\,\frac{\partial p_{t+1,i}}{\partial h_{t,j}} = \sum_{i=1}^{d_h} \delta_{t+1,i}\,(\mathbf{W}_{hh})_{ij},

donde el primer factor es δt+1,i\delta_{t+1,i} por la definición del error, y el segundo sale de que pt+1,i=k(Whh)ikht,k+p_{t+1,i} = \sum_k (\mathbf{W}_{hh})_{ik}\,h_{t,k} + \dots depende de ht,jh_{t,j} sólo por el sumando k=jk = j, con derivada (Whh)ij(\mathbf{W}_{hh})_{ij}. Esa suma recorre la columna jj de Whh\mathbf{W}_{hh}, que es la fila jj de su transpuesta, o sea la coordenada jj de Whhδt+1\mathbf{W}_{hh}^{\top}\boldsymbol{\delta}_{t+1}. Así que ht=Whhδt+1\nabla_{\mathbf{h}_t}\ell = \mathbf{W}_{hh}^{\top}\boldsymbol{\delta}_{t+1}, y multiplicar por la pendiente del tanh\tanh da δt\boldsymbol{\delta}_t. En el último paso el mismo argumento arranca de la capa de salida en vez del paso siguiente, y de ahí sale el otro renglón.

El gradiente del peso compartido es una suma

Con un error por paso en la mano, los gradientes de los pesos salen sin recorrer nada más —y ahí aparece la suma que la lección anterior anunció—. Mira dónde interviene Whh\mathbf{W}_{hh}: en la preactivación de cada paso, pt=Whhht1+\mathbf{p}_t = \mathbf{W}_{hh}\mathbf{h}_{t-1} + \dots, con tt de 11 a TT. Es un solo peso alcanzado por TT caminos, uno por posición, y la regla de la cadena del bloque anterior ya zanjó qué hacer con eso: entre caminos se suma. Cada camino aporta lo que la lección sobre backpropagation deja para una preactivación —el error por la entrada transpuesta—, y el total es la suma de las TT aportaciones:

Whh=t=1Tδtht1,Wxh=t=1Tδtxt,bh=t=1Tδt.\nabla_{\mathbf{W}_{hh}}\ell = \sum_{t=1}^{T} \boldsymbol{\delta}_t\,\mathbf{h}_{t-1}^{\top}, \qquad \nabla_{\mathbf{W}_{xh}}\ell = \sum_{t=1}^{T} \boldsymbol{\delta}_t\,\mathbf{x}_t^{\top}, \qquad \nabla_{\mathbf{b}_h}\ell = \sum_{t=1}^{T} \boldsymbol{\delta}_t.

Léelas al lado de las del bloque anterior y verás que son las mismas. Allí el gradiente de una capa era δ(l)(h(l1))\boldsymbol{\delta}^{(l)}\left(\mathbf{h}^{(l-1)}\right)^{\top}, un solo producto exterior; aquí es ese producto exterior sumado sobre los pasos, porque allí cada matriz aparecía una vez y aquí Whh\mathbf{W}_{hh} aparece TT veces. Esa suma es la única diferencia entre backpropagation y backpropagation through time, y es el precio exacto de compartir: una matriz, TT trabajos, TT aportaciones que recoger.

Dos comprobaciones cierran la cuenta. La primera es de formas, el reflejo de la lección sobre backpropagation: cada sumando de Whh\nabla_{\mathbf{W}_{hh}}\ell es dh×dhd_h \times d_h, y la suma también, así que el gradiente tiene la forma de la matriz que corrige por más larga que sea la secuencia —TT decide cuántos términos se suman, no el tamaño del resultado—. La segunda mira el primer término: en Whh\nabla_{\mathbf{W}_{hh}}\ell el sumando de t=1t = 1 es δ1h0\boldsymbol{\delta}_1\mathbf{h}_0^{\top}, y como h0=0\mathbf{h}_0 = \mathbf{0} es la matriz nula. Cuadra con la lección anterior, donde h1=tanh(Wxhx1+bh)\mathbf{h}_1 = \tanh\left(\mathbf{W}_{xh}\mathbf{x}_1 + \mathbf{b}_h\right) no llevaba término de Whh\mathbf{W}_{hh}: el primer paso no usó esa matriz, así que no le debe gradiente. En Wxh\nabla_{\mathbf{W}_{xh}}\ell, en cambio, el término de t=1t = 1 sí cuenta, porque x1\mathbf{x}_1 es un token de verdad.

Quedan los pesos de la salida, Why\mathbf{W}_{hy} y by\mathbf{b}_y, y no traen nada nuevo. Aparecen una sola vez —la respuesta se lee una sola vez, al final—, así que su gradiente es un producto exterior sin suma, Why=(y^y)hT\nabla_{\mathbf{W}_{hy}}\ell = \left(\hat{y} - y\right)\mathbf{h}_T^{\top}: la capa de salida del bloque anterior, tal cual. La suma sobre el tiempo es asunto de los tres pesos compartidos, y de nadie más.

BPTT en NumPy, contra el sondeo numérico

La misma comprobación de la lección sobre backpropagation, aplicada aquí: derivar a mano y contrastar cada número contra un sondeo que mueve el peso un poquito y mide cómo cambia la pérdida. Una RNN de juguete con semilla fija, un estado de dh=3d_h = 3 y una secuencia de T=5T = 5. Ejecútala y mira dos cosas: la aportación de cada paso a Whh\nabla_{\mathbf{W}_{hh}} —la última, la de t=1t = 1, sale cero— y la línea final.

import numpy as np

rng = np.random.default_rng(0)
d_model, d_h, T = 4, 3, 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)
Why = rng.normal(size=(1, d_h)) * 0.5 # salida escalar: 1 x d_h
by = np.zeros(1)
X = rng.normal(size=(T, d_model)) * 0.5 # una fila por posición
y = 1.0 # etiqueta del documento


def forward():
hs = [np.zeros(d_h)] # h_0 = 0
for x in X:
hs.append(np.tanh(Whh @ hs[-1] + Wxh @ x + bh))
p = 1.0 / (1.0 + np.exp(-(Why @ hs[-1] + by)[0]))
return hs, p, -(y * np.log(p) + (1 - y) * np.log(1 - p))


hs, p, perdida = forward()
print("ŷ = %.4f ℓ = %.4f" % (p, perdida))

# Hacia atrás: un delta por posición, del final al principio.
gWhh, gWxh, gbh = np.zeros_like(Whh), np.zeros_like(Wxh), np.zeros_like(bh)
dh = Why[0] * (p - y) # ∂ℓ/∂h_T, de la capa de salida
aporte = []
for t in range(T, 0, -1):
delta = (1 - hs[t] ** 2) * dh # δ_t = (1 - h_t²) ⊙ (∂ℓ/∂h_t)
gWhh += np.outer(delta, hs[t - 1]) # δ_t h_{t-1}ᵀ : lo que aporta el paso t
gWxh += np.outer(delta, X[t - 1])
gbh += delta
aporte.append(np.linalg.norm(np.outer(delta, hs[t - 1])))
dh = Whh.T @ delta # ∂ℓ/∂h_{t-1} : transporta con Whhᵀ
gWhy, gby = np.outer([p - y], hs[T]), np.array([p - y])

print("aporte de cada paso a ∇Whh, de t=%d a t=1:" % T)
print(" ", np.round(aporte, 4), " <- el ultimo, t=1, es 0 porque h_0 = 0")

# El sondeo: mover cada peso a mano y ver cuánto cambia la pérdida.
peor = 0.0
for P, G in [(Whh, gWhh), (Wxh, gWxh), (bh, gbh), (Why, gWhy), (by, gby)]:
for k in np.ndindex(P.shape):
v = P[k]
P[k] = v + 1e-6; mas = forward()[2]
P[k] = v - 1e-6; menos = forward()[2]
P[k] = v
peor = max(peor, abs((mas - menos) / 2e-6 - G[k]))
print("mayor diferencia con el sondeo numerico: %.1e" % peor)
numpy

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

Lee la salida contra la derivación. La lista de aportes tiene cinco números, uno por paso, y el último es cero exacto: el paso t=1t = 1 multiplica por h0=0\mathbf{h}_0 = \mathbf{0} y no le pide nada a Whh\mathbf{W}_{hh}, tal como decía la cuenta. Y la última línea es la lección: los gradientes de los cinco grupos de pesos coinciden con el sondeo hasta 101010^{-10}, que es el error del sondeo y no del algoritmo. La suma sobre el tiempo no es una aproximación ni un apaño; es el gradiente exacto, y el sondeo lo confirma número a número.

Comprueba tu intuición

Cinco preguntas: qué forma tiene el gradiente del peso compartido, de dónde sale su suma, cómo es la recurrencia hacia atrás, qué acumula el sesgo y qué obliga a guardar la vuelta.

Una RNN con un estado de dh=64d_h = 64 coordenadas lee un documento de T=40T = 40 tokens. ¿Cuántos números tiene Whh\nabla_{\mathbf{W}_{hh}}\ell?

Se acepta un margen de ±0.

En una red del bloque anterior, el gradiente de cada matriz de pesos era un solo producto exterior. En la RNN, Whh\nabla_{\mathbf{W}_{hh}}\ell es una suma de TT de ellos. ¿De dónde sale esa suma?

Marca todo lo que sea cierto de la recurrencia hacia atrás δ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).

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

La pasada hacia atrás acumula bh=tδt\nabla_{\mathbf{b}_h}\ell = \sum_t \boldsymbol{\delta}_t paso a paso. Con estos tres errores, ¿qué imprime?

import numpy as np
 
deltas = [
    np.array([1.0, -2.0]),
    np.array([0.5, 0.5]),
    np.array([-1.0, 1.0]),
]
g = np.zeros(2)
for d in deltas:
    g += d
print(g)
 

La pasada hacia atrás recorre los pasos de TT a 11 y en cada uno usa ht\mathbf{h}_t y ht1\mathbf{h}_{t-1}. ¿Qué obliga eso a hacer durante la pasada hacia adelante?

Y la vuelta entera escrita por ti, que es la pieza que faltaba para entrenar la RNN de la lección anterior.

Escribe la pasada hacia atrás de la recurrencia —los gradientes de los pesos compartidos—, a partir de lo que la de ida dejó guardado.

  • bptt(hs, X, Whh, dhT) recibe la lista de estados hs =[h0,,hT]= [\mathbf{h}_0, \dots, \mathbf{h}_T] con h0=0\mathbf{h}_0 = \mathbf{0}; la entrada X de forma (T,dmodel)(T, d_{\text{model}}), cuya fila t1t-1 es xt\mathbf{x}_t; la matriz recurrente Whh; y dhT =hT= \nabla_{\mathbf{h}_T}\ell, el gradiente que la capa de salida entrega al último estado.
  • Devuelve la terna (gWhh, gWxh, gbh), cada una sumada sobre los TT pasos.

La función no puede escribir en la lista ni en los arrays que recibe.

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.


El gradiente ya está, exacto y comprobado, y trae dentro una advertencia que aún no hemos leído. La recurrencia hacia atrás multiplica por Whh\mathbf{W}_{hh}^{\top} en cada paso —la misma matriz, una vez por posición—, y además enmascara con las pendientes 1htht1 - \mathbf{h}_t \odot \mathbf{h}_t, esas que la lección sobre funciones de activación ya te enseñó a apagarse cuando el estado se satura. Un error que sale de la posición TT y tiene que llegar hasta la 11 cruza ese producto decenas de veces, y multiplicar tantas veces por una misma matriz hace una de dos cosas: encoge hacia nada, o se dispara.

Cualquiera de las dos rompe justo lo que la RNN venía a arreglar. Si el gradiente que conecta la primera posición con la pérdida llega desvanecido, la red no puede aprender que algo del principio del documento importaba para el final —la dependencia larga, que era el motivo de tener memoria—. Cuánto encoge o crece, y por qué depende de una sola propiedad de Whh\mathbf{W}_{hh}, es la siguiente lección, sobre el gradiente que se desvanece: la derivamos con cuidado y la vemos, en una gráfica, caer en picado con la distancia.

Para profundizar1 fuente · 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.

  • Deep Learning, cap. 10: Sequence Modeling: Recurrent and Recursive Nets
    libroGoodfellow, Bengio y Courville, 2016deeplearningbook.orgEN

    Su §10.2.2 es esta misma derivación: en su ecuación 10.13 reaparecen el transporte por la matriz recurrente y la máscara 1 - h². Más corta que la tuya y sin comprobación numérica; llama W, U y V a tus W_hh, W_xh y W_hy.