Descenso de gradiente

Descenso de gradiente

30 min de lectura

Ya se puede escribir entera la pregunta del entrenamiento, y escribirla no la resuelve. La lección anterior, sobre funciones de pérdida, dejó puesta la última pieza: una regla que compara lo que la red contesta con lo que debía contestar y lo resume en un número. Buscar los parámetros que lo hacen mínimo ya es una pregunta bien planteada, y sigue sin haber forma de contestarla probando. La red de la lección sobre el forward pass guarda 4141 números; con un vocabulario de veinte mil entradas y una capa oculta de 128128 neuronas son 25602572\,560\,257, y cada uno de ellos es un número real.

Hay, en cambio, una pregunta mucho más pequeña que sí se contesta. Coge uno solo de esos 4141 parámetros, súbelo una centésima, deja los otros 4040 quietos y vuelve a evaluar la pérdida: si ha subido, ese parámetro iba en la dirección equivocada, y si ha bajado, iba bien. Es una respuesta local —no dice dónde está el mínimo, dice hacia dónde queda el suelo desde aquí— y, llevada al límite, tiene nombre desde el siglo XVII: es la derivada parcial de la pérdida respecto de ese parámetro. Sondear los 4141 cuesta 4141 evaluaciones. Derivar los da sin sondear ninguno, y para una neurona la cuenta cabe en una línea; conseguirla con capas ocultas de por medio ocupa el resto del bloque.

Antes de derivar nada, mira el terreno. Fija todos los parámetros de una red menos dos: la pérdida pasa a ser una función de esos dos números, y una función de dos variables se dibuja como una superficie con valles y laderas. La del explorable no sale de ninguna red: me la he inventado para que quepa en un plano y tenga dos fondos en vez de uno.

Mueve el punto de partida a un lado y a otro de la cresta central: la trayectoria termina siempre en el fondo que le queda debajo, y ninguna cruza al otro. La segunda barra es la tasa de aprendizaje, que fija cuánto avanza en cada paso; déjala quieta de momento.

El descenso no compara los dos valles, porque no los ve. Desde donde está tiene acceso a una sola cosa —cuánto cambia la pérdida al moverse un poco en cada dirección— y con eso decide el paso siguiente.

El gradiente apunta cuesta arriba, y por eso se resta

Llamemos θ\theta a todos los parámetros de la red a la vez —los pesos y los sesgos de todas las capas, apilados en un solo vector de PP coordenadas— y L(θ)\mathcal{L}(\theta) a la pérdida sobre el batch. El gradiente es el vector que reúne una derivada parcial por parámetro, θLRP\nabla_{\theta}\mathcal{L} \in \mathbb{R}^{P}. Cada una de sus coordenadas dice a qué ritmo cambia la pérdida cuando se mueve un parámetro y los demás se quedan quietos: lo que medía el sondeo de una centésima, ahora exacto y sin evaluar nada dos veces. Para la neurona de la lección sobre la neurona artificial, con sus pesos wRd\mathbf{w} \in \mathbb{R}^{d} y su sesgo bb, son d+1d + 1 números:

wL=(Lw1,  ,  Lwd)Rd,LbR.\nabla_{\mathbf{w}}\mathcal{L} = \left(\frac{\partial \mathcal{L}}{\partial w_1}, \;\dots, \;\frac{\partial \mathcal{L}}{\partial w_d}\right)^{\top} \in \mathbb{R}^{d}, \qquad \frac{\partial \mathcal{L}}{\partial b} \in \mathbb{R}.

Lo primero que hay que ver es que θL\nabla_{\theta}\mathcal{L} tiene la misma forma que θ\theta. No es casualidad: es lo que permite restar el uno del otro.

Falta justificar por qué se resta ése y no otra cosa. Toma una dirección cualquiera, un vector uRP\mathbf{u} \in \mathbb{R}^{P} de longitud 11, y pregunta cuánto cambia la pérdida al avanzar un poco en ella. Eso es la derivada direccional, que para una función con derivadas parciales es un producto escalar,

DuL(θ)=(θL)u,D_{\mathbf{u}}\mathcal{L}(\theta) = \left(\nabla_{\theta}\mathcal{L}\right)^{\top}\mathbf{u},

que está acotado. La desigualdad de Cauchy-Schwarz da (θL)uθLu=θL\left\lvert \left(\nabla_{\theta}\mathcal{L}\right)^{\top}\mathbf{u} \right\rvert \leq \lVert \nabla_{\theta}\mathcal{L} \rVert \, \lVert \mathbf{u} \rVert = \lVert \nabla_{\theta}\mathcal{L} \rVert, con igualdad únicamente cuando u\mathbf{u} es paralela al gradiente. De todas las direcciones, entonces, la del gradiente es la que más sube la pérdida y la contraria la que más la baja. Ninguna otra baja más por unidad de paso.

De ahí sale la regla entera. Partiendo de unos parámetros iniciales θ0\theta_0, cada paso resta el gradiente escalado:

θt+1=θtηθL(θt),\theta_{t+1} = \theta_t - \eta\,\nabla_{\theta}\mathcal{L}(\theta_t),

donde el subíndice cuenta pasos —θ5\theta_5 son los parámetros después de cinco pasos, no la quinta coordenada de nada— y η>0\eta > 0 es la tasa de aprendizaje (learning rate, que es como la vas a encontrar escrita en cualquier librería), el único número de la fórmula que eliges tú. Eso es el descenso de gradiente, y no tiene más piezas.

Que repetirlo baje la pérdida se ve en la aproximación de primer orden. Cerca de θt\theta_t la pérdida se parece a su plano tangente, L(θt+Δ)L(θt)+(θL)Δ\mathcal{L}(\theta_t + \Delta) \approx \mathcal{L}(\theta_t) + \left(\nabla_{\theta}\mathcal{L}\right)^{\top}\Delta, y el paso del descenso es precisamente Δ=ηθL\Delta = -\eta\,\nabla_{\theta}\mathcal{L}, así que

L(θt+1)L(θt)ηθL2.\mathcal{L}(\theta_{t+1}) \approx \mathcal{L}(\theta_t) - \eta \left\lVert \nabla_{\theta}\mathcal{L} \right\rVert^{2}.

Lo que se resta es una norma al cuadrado, luego no es negativo: mientras el gradiente no sea nulo, el paso baja. Y toda la deuda está en el signo de esa ecuación, que es un \approx: la igualdad es del plano tangente, no de la pérdida, y sólo vale «cerca». Cuánto es cerca lo decide η\eta.

Una concesión, dicha en voz alta: el batch de esta lección son las diez reseñas enteras, así que cada paso mira todos los ejemplos antes de mover nada. Partirlo en trozos y dar un paso por trozo hace falta cuando los conjuntos son grandes, y el del curso tiene diez ejemplos.

Un paso demasiado largo sube la pérdida

Para ver qué hace η\eta conviene una pérdida cuya cuenta salga entera, así que tomemos la más pequeña que existe: un solo parámetro y L(w)=w2\mathcal{L}(w) = w^{2}, con el mínimo en w=0w = 0 y dLdw=2w\dfrac{d\mathcal{L}}{dw} = 2w. La regla queda

wt+1=wtη2wt=(12η)wt,w_{t+1} = w_t - \eta\,2w_t = \left(1 - 2\eta\right)w_t,

una multiplicación por la misma constante en cada paso. Iterándola desde w0w_0,

wt=(12η)tw0,w_t = \left(1 - 2\eta\right)^{t} w_0,

y una potencia tiende a cero si y sólo si la base mide menos que 11. La condición completa es 12η<1\lvert 1 - 2\eta \rvert < 1, es decir 0<η<10 < \eta < 1, y el propio factor cuenta lo que pasa en cada tramo:

η\eta12η1 - 2\etaqué hace wtw_t
0<η<1/20 < \eta < 1/2entre 00 y 11se encoge sin cambiar de signo
η=1/2\eta = 1/200cae en el mínimo en un solo paso
1/2<η<11/2 < \eta < 1entre 1-1 y 00alterna de signo y se encoge
η=1\eta = 11-1alterna entre w0w_0 y w0-w_0 para siempre
η>1\eta > 1menor que 1-1alterna y crece: diverge
Tres gráficas apiladas, las tres con la misma parábola: la pérdida en el eje vertical y el peso en el horizontal, con el mínimo en el centro. En las tres el primer punto está en el mismo sitio, a la derecha, donde el peso vale uno. Arriba, con tasa de aprendizaje 0.10, los seis puntos bajan por el brazo derecho acercándose al centro cada vez menos, y el último se queda a medio camino. En medio, con 0.40, el segundo punto ya está casi en el fondo y el tercero encima de él. Abajo, con 1.05, los puntos saltan de un brazo al otro y cada uno queda más alto que el anterior: el quinto llega casi al borde superior del dibujo.
Los dos fallos son opuestos: el paso corto baja siempre y no termina de llegar, y el paso largo cruza el fondo y sale más arriba de donde entró. No hay un tamaño de paso correcto por naturaleza; hay uno que depende de cuánto se curva la pérdida.

Ese 11 del límite es de esta parábola y no de todas. Con L(w)=aw2\mathcal{L}(w) = a\,w^{2} la derivada pasa a ser 2aw2aw, el factor a 12ηa1 - 2\eta a y la condición a η<1/a\eta < 1/a: cuanto más se curva la pérdida, más corto tiene que ser el paso. Lo incómodo es que una red se curva de manera distinta en cada dirección, y η\eta es un solo número para todas.

El explorable arranca justo en ese problema, L=x2+20y2\mathcal{L} = x^{2} + 20y^{2}: casi plano a lo largo de xx y muy empinado a lo largo de yy. La cuenta de arriba, coordenada a coordenada, pide η<1\eta < 1 para la primera y η<1/20=0.05\eta < 1/20 = 0.05 para la segunda, y manda la segunda.

Sube la tasa de aprendizaje del cañón de cinco en cinco milésimas y mira dónde se rompe: hasta 0.045 la trayectoria desciende, en 0.05 rebota entre las dos paredes sin acercarse al fondo y por encima se escapa del gráfico. El aviso rojo salta cuando la trayectoria abandona el dibujo, no cuando empieza a ir mal.

Con η=0.05\eta = 0.05 —el valor con el que arranca— el factor de la coordenada empinada vale 1-1 exacto, y la trayectoria hace en dos dimensiones lo que la cuarta fila de la tabla hace en una: alternar sin avanzar. El cuenco del primer botón es convexo —un solo fondo, así que baje quien baje llega al mismo sitio— y aguanta el deslizador entero: su dirección más empinada sólo pide η<1/3\eta < 1/3. Que un mismo paso sirva para una dirección y no para la otra tiene arreglos con nombre propio —momentum, Adam—, y quedan fuera de este curso.

Con una sola capa, el gradiente sale de una línea

Falta bajar de θ\theta a un parámetro concreto, y con la neurona de las reseñas la cuenta está medio hecha. Es una neurona con sigmoide en la salida y entropía cruzada como pérdida, que es justo el par que la lección anterior derivó: para un ejemplo, /z=y^y\partial \ell / \partial z = \hat{y} - y. Lo único que falta es cómo depende zz de cada peso, y eso lo da la definición misma de la neurona, z=wx+b=j=1dwjxj+bz = \mathbf{w}^{\top}\mathbf{x} + b = \sum_{j=1}^{d} w_j x_j + b, donde cada wjw_j aparece una sola vez, multiplicado por xjx_j:

zwj=xj,zb=1.\frac{\partial z}{\partial w_j} = x_j, \qquad \frac{\partial z}{\partial b} = 1.

Encadenando las dos derivadas,

wj=zzwj=(y^y)xj,b=y^y,\frac{\partial \ell}{\partial w_j} = \frac{\partial \ell}{\partial z} \cdot \frac{\partial z}{\partial w_j} = \left(\hat{y} - y\right)x_j, \qquad \frac{\partial \ell}{\partial b} = \hat{y} - y,

que en forma vectorial es w=(y^y)x\nabla_{\mathbf{w}}\,\ell = \left(\hat{y} - y\right)\mathbf{x}. Esa fórmula dice dos cosas y conviene separarlas. El error y^y\hat{y} - y es un único número que multiplica a todas las coordenadas por igual: fija cuánto se corrige. Y x\mathbf{x} decide el reparto: una entrada del vocabulario que no aparece en la reseña tiene xj=0x_j = 0 y su peso no se mueve ni una milésima, y una que aparece dos veces tira el doble. La neurona sólo aprende sobre lo que ha visto.

La pérdida del batch es la media de las BB pérdidas por ejemplo, así que su gradiente es la media de los BB gradientes:

wL=1Bi=1B(y^iyi)xi=1BX(Y^Y),\nabla_{\mathbf{w}}\mathcal{L} = \frac{1}{B}\sum_{i=1}^{B}\left(\hat{y}_i - y_i\right)\mathbf{x}_i = \frac{1}{B}\,\mathbf{X}^{\top}\left(\hat{\mathbf{Y}} - \mathbf{Y}\right),

donde la segunda forma es la primera escrita como producto: con dL=1d_L = 1 las predicciones y las etiquetas son columnas de BB números, X\mathbf{X}^{\top} tiene forma d0×Bd_0 \times B, y el resultado es la columna de d0d_0 coordenadas que hace falta. El sesgo sale de la misma suma sin el xi\mathbf{x}_i: es L/b=1Bi(y^iyi)\partial \mathcal{L} / \partial b = \frac{1}{B}\sum_i \left(\hat{y}_i - y_i\right).

Todo esto ha costado dos derivadas porque zz depende de w\mathbf{w} directamente: entre un peso y la pérdida hay una suma y una sigmoide, y nada más.

Descenso a mano, y una neurona que aprende sola

La primera celda es la parábola de arriba, sin red y sin NumPy: cinco pasos desde w0=1w_0 = 1 con las tres tasas de la figura, y al lado la forma cerrada. Ejecútala y compara los dos bloques.

# La pérdida más pequeña que existe: un parámetro y L(w) = w², con dL/dw = 2w.
def perdida(w):
return w * w


def desciende(w0, eta, pasos):
w = w0
trayectoria = [w]
for _ in range(pasos):
w = w - eta * 2.0 * w # w <- w - eta * dL/dw
trayectoria.append(w)
return trayectoria


w0 = 1.0
print("%6s %7s %s" % ("eta", "factor", " ".join(("w%d" % t).rjust(9) for t in range(6))))
for eta in [0.10, 0.40, 1.05]:
print("%6.2f %+7.2f %s"
% (eta, 1 - 2 * eta, " ".join("%9.5f" % w for w in desciende(w0, eta, 5))))
print()

# La forma cerrada da los mismos números sin dar un solo paso.
for eta in [0.10, 0.40, 1.05]:
bucle = desciende(w0, eta, 5)
cerrada = [(1 - 2 * eta) ** t * w0 for t in range(6)]
iguales = all(abs(a - b) < 1e-12 for a, b in zip(bucle, cerrada))
print("eta = %.2f bucle y formula coinciden: %-5s perdida final %10.7f"
% (eta, iguales, perdida(bucle[-1])))

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

Las tres filas son las tres de la figura. Con η=0.10\eta = 0.10 el factor vale 0.80.8 y cinco pasos dejan ww en 0.327680.32768: va bien y va despacio. Con 0.400.40 el factor es 0.20.2 y en dos pasos ww ya vale 0.040.04. Con 1.051.05 el factor es 1.1-1.1, los signos alternan, las magnitudes crecen y la pérdida final vale 2.592.59: dos veces y media la del punto de partida, después de cinco pasos dedicados a bajarla.

La segunda celda entrena de verdad. La neurona es la de la lección sobre la neurona artificial, con las mismas diez reseñas y el mismo vocabulario, y esta vez sus pesos empiezan en cero: con w=0\mathbf{w} = \mathbf{0} y b=0b = 0 la sigmoide contesta 0.50.5 a las diez, que es la peor de las tres redes que ordenó la lección anterior. No hay un solo número al azar, así que verás estos valores y no otros. Ejecútala y mira caer la pérdida.

import numpy as np

# --- Las diez reseñas de siempre, su vocabulario y sus etiquetas.
resenas = [
"la película es divertida y la recomiendo",
"la película es lenta y aburrida",
"buena película la recomiendo",
"mala película muy lenta",
"divertida y buena",
"la película es mala",
"aburrida y lenta",
"la recomiendo",
"buena película pero lenta",
"divertida y la recomiendo",
]
y = np.array([1., 0., 1., 0., 1., 0., 0., 1., 0., 1.])
V = ["aburrida", "buena", "divertida", "la", "lenta", "mala", "película", "recomiendo"]
X = np.array([[t.split().count(e) for e in V] for t in resenas], dtype=float)


def entropia_cruzada(y_hat):
p = np.clip(y_hat, 1e-12, 1.0 - 1e-12)
return -np.mean(y * np.log(p) + (1.0 - y) * np.log(1.0 - p))


w = np.zeros(len(V)) # todos los pesos a cero
b, eta, B = 0.0, 1.0, len(y) # sesgo, tasa de aprendizaje, tamaño del batch

for paso in range(201):
y_hat = 1.0 / (1.0 + np.exp(-(X @ w + b))) # forward pass
if paso % 25 == 0:
print("paso %3d pérdida %.4f aciertos %2d de 10"
% (paso, entropia_cruzada(y_hat), int(((y_hat >= 0.5) == (y == 1.)).sum())))
if paso == 200:
break
error = y_hat - y # (B,)
w -= eta * (X.T @ error) / B
b -= eta * float(error.sum()) / B

a_mano = np.array([-1., 1., 1., 0., -1., -1., 0., 1.])
print("\ncon los pesos que puse a mano:",
"%.4f" % entropia_cruzada(1.0 / (1.0 + np.exp(-(X @ a_mano)))))
for entrada, peso in sorted(zip(V, w), key=lambda par: -par[1]):
print(" %-11s %+6.3f" % (entrada, peso))
print(" %-11s %+6.3f" % ("(sesgo)", b))
numpy

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

La primera línea es el punto de partida: pérdida 0.69310.6931, que es log2\log 2, y cinco aciertos de diez, que es lo que saca cualquier red que conteste lo mismo a las diez reseñas. Veinticinco pasos después la pérdida vale 0.08960.0896 y los aciertos son diez, y en el paso 200200 queda en 0.01300.0130. Compárala con el 0.22080.2208 de los ocho pesos que escribí yo mirando las reseñas: el descenso pasa por delante de mi versión antes del paso 2525 y sigue bajando. Nadie le ha dicho a la neurona qué palabra significa qué.

Los pesos finales lo confirman. El mayor, +3.289+3.289, se lo lleva recomiendo, y el menor, 3.023-3.023, es el de lenta. Pero mira los dos que enseñan el límite de un conjunto de diez ejemplos: película acaba en 2.238-2.238 y la en +1.221+1.221, y ninguna de las dos opina de nada. La primera sale en cuatro reseñas negativas y en dos positivas, la segunda al contrario, y el descenso hace lo que se le pidió, que es bajar la pérdida sobre estas diez reseñas y no sobre el idioma. Con diez ejemplos, una neurona no distingue una opinión de una coincidencia.

Comprueba tu intuición

Cuatro preguntas: un descenso a mano, por qué el paso va en esa dirección, qué garantiza el método y qué no, y una actualización en NumPy.

El descenso corre sobre L(w)=w2\mathcal{L}(w) = w^{2} partiendo de w0=8w_0 = 8, con η=0.25\eta = 0.25. ¿Cuánto vale w3w_3?

Se acepta un margen de ±0.01.

De todas las direcciones en las que se pueden mover los parámetros, ¿por qué el paso va en la de θL-\nabla_{\theta}\mathcal{L}?

Marca todo lo que sea cierto sobre lo que el descenso de gradiente garantiza y lo que no.

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

Un paso de descenso sobre una neurona de dos entradas, con los pesos empezando en cero y sin sesgo. ¿Qué imprime?

import numpy as np
 
X = np.array([[1., 0.], [0., 2.]])
y = np.array([1., 0.])
w = np.zeros(2)
 
y_hat = 1.0 / (1.0 + np.exp(-(X @ w)))
w = w - 1.0 * (X.T @ (y_hat - y)) / len(y)
print(np.round(w, 3).tolist())
 

El método está completo y la neurona que acaba de aprender tiene una sola capa. Eso no es un detalle del ejemplo, es la condición que ha hecho posible la última sección: entre un peso y la pérdida hay dos pasos, y dos derivadas encadenadas cierran la cuenta. Pon una capa oculta en medio y el camino se ramifica. Un peso de la primera capa mueve la salida de su neurona oculta; esa salida entra en todas las neuronas de la capa siguiente, y cada una contribuye a la pérdida por su cuenta. La regla del descenso no cambia ni una letra. Lo que deja de estar disponible es el gradiente, y sin él no hay nada que restar.

La regla de la cadena que hemos usado dos veces, la de una variable, no cubre ese caso: está escrita para una cantidad que llega a la siguiente por un único camino. Cómo se deriva una composición cuando las variables intermedias son vectores, y qué objeto sustituye a la derivada cuando hay muchas entradas y muchas salidas a la vez, es la siguiente lección, sobre la regla de la cadena. Convertirla después en un procedimiento que recorra una red entera sin repetir ninguna cuenta ocupa el resto del bloque.

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.

  • Deep Learning, cap. 4: Numerical Computation
    libroGoodfellow, Bengio y Courville, 2016deeplearningbook.orgEN

    Su §4.3 explica por qué el paso va en contra del gradiente (la derivada direccional) y cómo la curvatura —el número de condición del hessiano— fija el paso más largo que no diverge: el x² + 20y² de la lección dicho en general.

  • An overview of gradient descent optimization algorithms
    paperSebastian Ruder, 2016arXiv:1609.04747EN

    Sigue donde la lección lo deja: momentum, Adam y los demás métodos que hacen que una sola tasa de aprendizaje sirva para todas las direcciones, y el reparto del batch en trozos. Es un repaso; no desarrolla los métodos desde cero.