Funciones de pérdida: MSE y entropía cruzada
30 min de lectura
Medir no es contestar, y la red sólo sabe contestar. La lección anterior, sobre el forward pass en forma matricial, cerró la mitad del recorrido que va de a : una cadena de productos que atraviesa las capas sin preguntarse ni una vez si lo que entrega se parece a lo que debía. La otra mitad empieza aquí, fijando una regla que compare con y resuma en un único número lo mal que va el batch. Esa regla se llama función de pérdida, y hay que elegirla.
Reglas que sirvan hay muchas, y la elección pesa. Toma divertida y la recomiendo, que habla bien de la película: su etiqueta es , y la red contesta con un número entre y que vale más cuanto más se incline por esa clase. Supón que entrega un . Falla, porque tenía que haber dado algo pegado al , y falla con toda la seguridad del mundo, porque más cerca del no podía quedarse. Esta lección compara dos candidatas, el error cuadrático medio y la entropía cruzada, y las dos coinciden en que ahí la pérdida es alta: por ahí no se distinguen. Lo que las separa es cuánto cambia esa pérdida cuando un peso cambia, y en un ejemplo así la primera de las dos, que es la que se le ocurre a cualquiera, apenas se inmuta.
Antes de escribir ninguna fórmula, mira la forma de esa diferencia. Lo dibujado no son las dos pérdidas sino la corrección que entrega cada una: cuánto se mueve por cada unidad que se mueva la preactivación de la neurona de salida, que es de donde tiran los pesos. En verde la entropía cruzada; en gris el error cuadrático medio (mean squared error, MSE). Con la etiqueta valiendo , la red acierta hacia la derecha del dibujo y falla hacia la izquierda.
De predicciones a un solo número
La red entrega , y los ejemplos traen sus etiquetas en un array de la misma forma, , fila con fila. Una pérdida resume ese par en un número, y se construye en dos niveles: una función que juzga un ejemplo comparando su predicción con su etiqueta, y la media de las que resulta,
donde son la predicción y la etiqueta del ejemplo : las filas de y de , escritas en columna. Es la media y no la suma para que no crezca con .
De se exigen tres cosas y no pesan igual: que devuelva un número, que baje al acercarse la predicción a la etiqueta, y que tenga derivada respecto de cada parámetro, porque de esas derivadas está hecho el método que corrige los pesos —la siguiente lección sobre el descenso de gradiente—. Las dos primeras las cumple casi cualquier cosa; la tercera es la que descarta candidatas, y es el mismo criterio con el que la lección sobre funciones de activación y no linealidad eligió : allí el escalón se cayó por tener derivada nula, no por medir mal.
La primera de las dos mide la distancia al cuadrado. El error cuadrático medio es
con el subíndice nombrando cuál de las candidatas es, y con la segunda forma escrita para —la red de reseñas—, donde e son escalares y los escribo sin negrita. El cuadrado hace dos cosas: quita el signo, de modo que pasarse y quedarse corto cuestan lo mismo, y castiga los errores grandes más que proporcionalmente, porque fallar por cuesta dieciséis veces lo que fallar por , no cuatro.
Ese es el criterio correcto cuando la red predice una cantidad: si hay que adivinar cuántas estrellas le puso su autor, la diferencia son estrellas y su cuadrado mide lo que se busca. Al derivar sale , la primera propiedad buena de esta pérdida: la corrección es proporcional al error. Sobre ese factor hay una convención que no adopto —mucha gente pone un delante para que se cancele al derivar—, y no cambia nada medible: multiplicar todas las correcciones por una constante se absorbe en el tamaño del paso, que es lo que ajusta la lección siguiente.
Por qué se apaga justo cuando más falta hace
Clasificar no es predecir una cantidad, y ahí la cuenta cambia. La etiqueta vale o , la salida sale de una sigmoide, , y la derivada que hace falta no es la de sino la de , que es lo que los pesos controlan. Encadenando, y con de la lección sobre funciones de activación:
Mira el producto de los dos últimos factores: es la derivada de la sigmoide, y la lección sobre funciones de activación ya dijo dónde se apaga: cuando se acerca a o a , o sea, cuando la red está segura. La saturación que allí era un defecto de la activación reaparece aquí multiplicando a la corrección. Con y —la red dice que no, con convicción, y se equivoca— el primer factor vale casi , pero vale , y el producto se queda en . El peor error del batch pide la corrección más pequeña del batch.
No es un caso raro elegido a mano: es la mitad izquierda entera de la curva gris de arriba. Y ningún ajuste del tamaño del paso arregla una pérdida que confunde sus dos colas, porque las escala a las dos por igual.
La entropía cruzada, leyendo la salida como probabilidad
El arreglo no es parchear la fórmula, sino cambiar lo que se supone que la red está diciendo. Hasta aquí era un número entre y ; tómalo ahora como la probabilidad que la red le da a la clase positiva, . Con eso, la probabilidad que le asigna a la etiqueta que traía el ejemplo se escribe de una vez para los dos casos,
porque con el segundo factor vale y queda , y con es el primero el que desaparece. Suponiendo los ejemplos independientes, la probabilidad de observar el conjunto entero es el producto de las , y a esa cantidad, como función de los parámetros, se la llama la verosimilitud. Hay que hacerla grande. Multiplicar números menores que uno lleva el producto a cero muy deprisa, así que tomamos logaritmos —que conservan el orden, luego el máximo no se mueve—, le damos la vuelta al signo para tener algo que minimizar y dividimos entre . Sale la entropía cruzada:
Léela por ejemplos. Uno con aporta : vale si la red le da probabilidad y crece sin techo según esa probabilidad baja hacia cero, y uno con mira a . Que no haya techo es la diferencia de fondo con el MSE, que nunca pasaba de .
Ahora la derivada, que es lo que decide. Derivamos respecto de , ponemos los dos términos sobre denominador común y encadenamos con la sigmoide:
En la primera línea el numerador se simplifica porque los dos productos se cancelan. En la segunda, el denominador es exactamente el que la sigmoide aporta al encadenar, y los dos se van.
La corrección es la diferencia entre lo que la red dijo y lo que debía decir. Nada más. En el ejemplo de antes vale en lugar de los del MSE. La saturación no ha desaparecido, la sigue teniendo la sigmoide; lo que pasa es que la pérdida trae en su derivada el factor inverso exacto que la deshace, y por eso la fórmula final no menciona a .
Más de dos clases: softmax
Con tres respuestas en lugar de dos —positiva, negativa y neutra— la capa de salida tiene neuronas y la sigmoide ya no sirve: aplasta cada coordenada por su cuenta, y tres números en no forman una distribución. La función que sí lo hace es el softmax, el mismo de la lección sobre Word2Vec, que repartía probabilidad entre las entradas del vocabulario:
positivo por ser exponencial y sumando por el denominador. La etiqueta se escribe one-hot, con un único en la posición de la clase correcta , tal como la lección sobre one-hot codificaba una entrada del vocabulario. La pérdida es la de antes, escrita para clases:
donde la segunda igualdad sale de que la one-hot anula todos los sumandos menos uno. Sólo cuenta la probabilidad que la red le dio a la clase verdadera; lo que reparta entre las otras importa porque el denominador las ata a todas.
La derivada, con las dos clases de índice que hay que distinguir
Derivamos respecto de cada . El softmax mete todas las coordenadas en cada una —mover mueve el denominador, y con él las probabilidades—, así que hacen falta dos cuentas, según sea o no la clase correcta. Llamemos , con .
Cuando , el numerador también depende de , y es una derivada de un cociente:
Cuando , el numerador no depende de y sólo se mueve el denominador:
Con , la regla de la cadena da los dos casos:
y las dos son la misma expresión, porque y . Coordenada a coordenada:
Con esto reproduce el caso binario: la fórmula corta de la sección anterior no era un accidente de tener una sola salida.
Las dos derivaciones llegan al mismo sitio, y por eso merece la pena hacerlas una vez: de aquí a la lección sobre backpropagation, la capa de salida no vuelve a derivarse, entra con y ya está. El softmax está aquí porque el resto del curso predice sobre vocabularios enteros, donde las clases son decenas de miles.
Las dos pérdidas sobre las diez reseñas
La celda evalúa las dos candidatas sobre las diez reseñas de siempre, con las preactivaciones de la neurona de pesos puestos a mano de la lección sobre la neurona artificial. Nada va al azar. Ejecútala y mira tres cosas: qué ordena cada columna, dónde se apaga cada corrección, y qué pasa con un logaritmo de cero.
# --- Las diez reseñas de siempre, con sus etiquetas y sus preactivaciones.
y = np.array([1., 0., 1., 0., 1., 0., 0., 1., 0., 1.])
z = np.array([2., -2., 2., -2., 2., -1., -2., 1., 0., 2.])
sigmoide = lambda z: 1.0 / (1.0 + np.exp(-z))
def mse(y_hat, y):
return float(np.mean((y_hat - y) ** 2))
def entropia_cruzada(y_hat, y):
p = np.clip(y_hat, 1e-12, 1.0 - 1e-12) # sin esto, un 0 se lleva la media entera
return float(-np.mean(y * np.log(p) + (1.0 - y) * np.log(1.0 - p)))
# 1. Tres redes: la de pesos a mano, la misma sin convicción, y una que contesta 0.5 a todo.
convencida = sigmoide(z)
tibia = 0.5 + (convencida - 0.5) * 0.1
print("%-25s %8s %8s %8s" % ("", "aciertos", "MSE", "EC"))
for nombre, p in [("neurona de pesos a mano", convencida),
("la misma, sin convicción", tibia),
("contesta 0.5 a todo", np.full(10, 0.5))]:
aciertos = int(((p >= 0.5) == (y == 1.0)).sum())
print("%-25s %8d %8.4f %8.4f" % (nombre, aciertos, mse(p, y), entropia_cruzada(p, y)))
print()
# 2. La corrección de cada pérdida, para un ejemplo cuya etiqueta vale 1.
print("%6s %6s %11s %11s %9s" % ("z", "ŷ", "|dMSE/dz|", "|dEC/dz|", "cociente"))
for zz in [-6.0, -3.0, 0.0, 3.0, 6.0]:
s = sigmoide(zz)
d_mse = abs(2.0 * (s - 1.0) * s * (1.0 - s))
d_ec = abs(s - 1.0)
print("%6.1f %6.3f %11.6f %11.6f %9.1f" % (zz, s, d_mse, d_ec, d_ec / d_mse))
print()
# 3. El logaritmo de cero, y lo que cuesta taparlo.
p = np.array([0.9, 0.0, 0.5])
recortada = np.clip(p, 1e-12, 1.0)
with np.errstate(divide="ignore"):
print("sin recortar:", -np.log(p), " media:", float(-np.mean(np.log(p))))
print("recortando: ", np.round(-np.log(recortada), 2),
" media:", round(float(-np.mean(np.log(recortada))), 2))
La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador.
La primera tabla justifica la lección. El recuento de aciertos empata las dos primeras en nueve de diez: no distingue entre decir y decir , porque las dos caen del mismo lado. Las dos pérdidas las separan — contra la convencida, contra la tibia—, y dejan en último lugar a la que contesta a todo, con y , que es . Ordenan igual. Lo que no hacen igual es corregir.
Eso es la segunda tabla. En , con la etiqueta valiendo , la entropía cruzada entrega y el MSE : un cociente de . En , con la red sin opinión, es . La ventaja crece hacia los extremos, y uno de ellos es donde la red se equivoca. La última fila repite el con la red acertando, y ahí las dos correcciones son minúsculas porque deben serlo.
La tercera prueba enseña el detalle que rompe una implementación. es infinito, y basta un ejemplo al que la red le dé probabilidad cero para que la media del batch entero valga infinito y deje de decir nada. El recorte a lo evita a un precio que digo en voz alta: ese no lo ha medido nadie.
Comprueba tu intuición
Cinco preguntas —una entropía cruzada a mano, el MSE al clasificar, la corrección limpia, el recorte del logaritmo y el máximo del softmax— y un desafío con la capa de salida.
Una red con sigmoide en la salida le asigna a una reseña cuya etiqueta es . ¿Cuánto vale la entropía cruzada de ese ejemplo?
Se acepta un margen de ±0.01.
¿Qué le pasa al error cuadrático medio cuando la salida de la red es una sigmoide y la etiqueta vale o ?
Marca todo lo que sea cierto sobre el resultado de esta lección.
Marca todas las opciones correctas. Se corrige todo o nada: no hay puntuación parcial.
Así recorta la celda las probabilidades antes de tomar el logaritmo. ¿Qué imprime?
import numpy as np
p = np.array([0.9, 1e-15, 0.5])
print(np.round(-np.log(np.clip(p, 1e-12, 1.0)), 2))
softmax(z) y softmax(z + 100) devuelven exactamente lo mismo. ¿Por qué, y qué se gana restando el máximo de cada fila antes de exponenciar?
Escribe la capa de salida de una red de varias clases, sobre un batch entero y sin bucles de Python:
softmax(Z)recibeZde forma y devuelve las probabilidades fila a fila. Cada fila tiene que sumar , y la función tiene que sobrevivir a preactivaciones grandes:np.exp(1000)desborda.entropia_cruzada(Y_hat, Y)devuelve la media sobre el batch de , conYcodificada one-hot. Un que valga cero no puede devolverinf.
Ninguna de las dos puede escribir 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.
Ya hay un número que dice lo mal que va la red, y en toda la lección no se ha movido ni un peso. Las tres redes llegaron con sus números puestos, la pérdida las ha ordenado, que es lo que se le pedía, y ninguna ha mejorado por haber sido medida.
Entrenar es hacer pequeño ese número, y lo único que se puede tocar para conseguirlo son los pesos: cambiarlos cambia , y con ella . De ese movimiento, la derivada da el sentido —si es positivo, subir sube la pérdida, luego hay que bajarlo— y no da la distancia: describe la pendiente en el punto donde estás y deja de valer en cuanto te mueves. Así que los pesos no se corrigen de una vez, sino a pasos, y el tamaño de cada paso es una decisión: demasiado corto no llega nunca, demasiado largo se pasa. Cómo se elige, qué ocurre cuando se elige mal y por qué repetirlo acaba dando con el mínimo sin verlo, es la siguiente lección, sobre descenso de gradiente.
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. 6: Deep Feedforward Networks
Su §6.2 saca la pérdida de la verosimilitud —que es la entropía cruzada— y explica por qué una salida que satura no va con el error cuadrático medio, y cómo el logaritmo lo arregla. El argumento de la lección, con menos álgebra.