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: 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 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 . No hay otro camino de vuelta porque no hubo otro de ida.
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 , la de un ejemplo, como en todo el bloque anterior desde la lección sobre funciones de pérdida; la media sobre el batch, , 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:
de modo que los pesos de la recurrencia llegan a la pérdida por un único sitio: a través de . 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 —el argumento del — es
y llamemos error del paso , y , al gradiente de la pérdida respecto de esa preactivación, . Es exactamente el 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 se aplica coordenada a coordenada, así que pasar de a es multiplicar por su pendiente, y esa pendiente, por la identidad de la lección sobre funciones de activación, es : se lee directamente del estado que la ida ya guardó, sin conservar . Queda el otro factor, , y tiene dos casos según de dónde le llegue a la pérdida. En el último paso, sólo de la capa de salida; en los demás, sólo del paso siguiente:
El arranque, , es la capa de salida de la lección anterior derivada hacia atrás: la sigmoide con la entropía cruzada entrega , y lo reparte entre las 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 y lo enmascara con la pendiente del . 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 , donde sólo alcanza la pérdida a través del paso siguiente. La coordenada de es la derivada de respecto de , y entra en la preactivación por cada una de sus coordenadas: es la suma sobre caminos de la lección sobre la regla de la cadena, con caminos.
donde el primer factor es por la definición del error, y el segundo sale de que depende de sólo por el sumando , con derivada . Esa suma recorre la columna de , que es la fila de su transpuesta, o sea la coordenada de . Así que , y multiplicar por la pendiente del da . 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 : en la preactivación de cada paso, , con de a . Es un solo peso alcanzado por 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 aportaciones:
Léelas al lado de las del bloque anterior y verás que son las mismas. Allí el gradiente de una capa era , un solo producto exterior; aquí es ese producto exterior sumado sobre los pasos, porque allí cada matriz aparecía una vez y aquí aparece veces. Esa suma es la única diferencia entre backpropagation y backpropagation through time, y es el precio exacto de compartir: una matriz, trabajos, aportaciones que recoger.
Dos comprobaciones cierran la cuenta. La primera es de formas, el reflejo de la lección sobre backpropagation: cada sumando de es , 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 — decide cuántos términos se suman, no el tamaño del resultado—. La segunda mira el primer término: en el sumando de es , y como es la matriz nula. Cuadra con la lección anterior, donde no llevaba término de : el primer paso no usó esa matriz, así que no le debe gradiente. En , en cambio, el término de sí cuenta, porque es un token de verdad.
Quedan los pesos de la salida, 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, : 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 y una secuencia de . Ejecútala y mira dos cosas: la aportación de cada paso a —la última, la de , sale cero— y la línea final.
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)
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 multiplica por y no le pide nada a , 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 , 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 coordenadas lee un documento de tokens. ¿Cuántos números tiene ?
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, es una suma de de ellos. ¿De dónde sale esa suma?
Marca todo lo que sea cierto de la recurrencia hacia atrás .
Marca todas las opciones correctas. Se corrige todo o nada: no hay puntuación parcial.
La pasada hacia atrás acumula 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 a y en cada uno usa y . ¿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 estadoshscon ; la entradaXde forma , cuya fila es ; la matriz recurrenteWhh; ydhT, el gradiente que la capa de salida entrega al último estado.- Devuelve la terna
(gWhh, gWxh, gbh), cada una sumada sobre los 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 en cada paso —la misma matriz, una vez por posición—, y además enmascara con las pendientes , 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 y tiene que llegar hasta la 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 , 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
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.