Clasificación, entropía cruzada y cómo no engañarte a ti mismo
Crea un clasificador logístico y descubre por qué un 98 % de accuracy puede no encontrar nada.
En esta página
Un modelo que responde esta pieza está bien sobre cada pieza que sale de la cinta acierta el 98,15 % de las veces. También es inútil: de las 74 piezas defectuosas del conjunto de prueba, no detecta ninguna.
Ambas frases describen el mismo modelo. La distancia entre ellas es este capítulo.
La primera mitad construye el clasificador. No necesita casi nada nuevo: el Capítulo 2 dio la receta para convertir una suposición sobre cómo se producen los datos en una función de pérdida, y el Capítulo 3 dio la maquinaria para bajar la pendiente de cualquier pérdida que esa receta te entregue. Aplica ambas cosas a una pregunta de sí/no y aparece la regresión logística, más una idea nueva —un logit— que volverá a cobrarse en el Capítulo 17.
La segunda mitad es la difícil. Todo lo que viene después en el curso se juzga por un número que alguien ha medido, y si no sabes distinguir una mejora real de un artefacto de medición, todos los capítulos siguientes son decoración. Así que: la matriz de confusión, precisión y recall, las tres particiones, la fuga de datos y la pregunta que casi nadie responde con honestidad: ¿cuántos ejemplos de prueba necesito de verdad?
La aritmética aquí recorre 20.000 filas, así que está vectorizada de principio a fin: NumPy lleva haciendo el trabajo desde el Capítulo 2, y a partir de aquí deja de merecer la pena señalarlo.
La cinta, con una pregunta más rara
Enlace a la sección: La cinta, con una pregunta más raraLa misma fábrica que en el Capítulo 1, pero una pregunta más difícil. En lugar de aceptar o rechazar, la pregunta es ¿esta pieza es defectuosa?; y los defectos son raros, lo que hace difícil la mitad de medición de este capítulo y engañosamente fácil la mitad de modelado.
import numpy as np
rng = np.random.default_rng(4)
N = 20_000
width = rng.normal(22.0, 0.9, N) # millimetres
weight = rng.normal(57.0, 3.0, N) # grams
z_true = -5.90 + 1.90 * (width - 22.0) + 0.42 * (weight - 57.0)
y = (rng.random(N) < 1 / (1 + np.exp(-z_true))).astype(float)
perm = rng.permutation(N)
train, val, test = perm[:12_000], perm[12_000:16_000], perm[16_000:]N = 20000 defects = 337 base rate = 0.0169
defects per split = 203 60 74Tres particiones, no dos. La razón merece su propia sección y la tendrá más abajo; por ahora, entrena con la primera, ajusta con la segunda y no mires la tercera.
Las features se estandarizan —se resta la media y se divide por la desviación típica— usando solo las estadísticas de entrenamiento, por la razón que el Capítulo 1 demostró con la cota de convergencia del perceptrón: los datos no centrados vuelven hostil la geometría. De qué filas puedes calcular esa media se convierte en una pregunta viva más adelante en este capítulo.
De un veredicto a una probabilidad
Enlace a la sección: De un veredicto a una probabilidadEl perceptrón devolvía un signo. Un signo no puede distinguir rechazar de rechazar, pero por poco, y esa diferencia es exactamente lo que una fábrica necesita para decidir qué piezas debería reinspeccionar primero una persona.
Así que sigue literalmente la receta del Capítulo 2. Escribe lo que afirmas sobre cómo se produce una etiqueta, toma la verosimilitud, toma el logaritmo, cámbiale el signo, y tienes una pérdida. Para un resultado de sí/no, la afirmación es una distribución Bernoulli: hay una probabilidad de que la pieza sea defectuosa, y
que es solo una forma compacta de escribir « si , y si ». Toma el logaritmo de eso y cámbiale el signo, y la pérdida para un ejemplo es
Esto es entropía cruzada binaria. No se eligió porque sea cómoda; es la log-verosimilitud negativa de la única distribución que puede tener el lanzamiento de una moneda. No había nada más disponible.
Lo que aún falta es de dónde sale . El modelo calcula una suma ponderada , que es un número real y recorre toda la recta, y una probabilidad tiene que vivir en . La función que se mueve entre ambas es la sigmoide logística:
logit -4.0 -> p = 0.0180 loss when y=1 and p=0.9 : 0.1054
logit -1.0 -> p = 0.2689 loss when y=1 and p=0.5 : 0.6931
logit 0.0 -> p = 0.5000 loss when y=1 and p=0.01 : 4.6052
logit 4.0 -> p = 0.9820Lee la columna de la derecha como una lista de precios. Acertar con un 90 % de confianza cuesta 0,105. Negarse a comprometerse cuesta 0,693 —que es , el precio de encogerse de hombros—. Equivocarse con seguridad cuesta 4,6, cuarenta y cuatro veces más, y el precio sube sin límite a medida que el modelo se vuelve más seguro de un error. La entropía cruzada no se limita a contar errores: cobra por la arrogancia.
El gradiente es predicción menos verdad
Enlace a la sección: El gradiente es predicción menos verdadEl Capítulo 3 decía: para entrenar cualquier cosa, obtén la derivada de la pérdida respecto a cada parámetro. Hazlo para un ejemplo. Con y :
Mostrar detalles
Las dos líneas que hacen que el lío se cancele. La sigmoide tiene una derivada inusualmente agradable, . Y la pérdida se diferencia como
Multiplica ambas por la regla de la cadena y el aparece una vez arriba y una vez abajo. Se cancela exactamente, y es lo que sobrevive. Esa cancelación no es una coincidencia: es lo que ocurre siempre que la pérdida es la log-verosimilitud negativa de una distribución y la función de salida es la que esa distribución usa de forma natural. Ese emparejamiento tiene un nombre —un modelo lineal generalizado— y el gradiente limpio es su huella.1
Así que la actualización es predicción menos verdad, por la entrada. Nada más. Aquí está el entrenador completo, que es el descenso del Capítulo 3 con una línea cambiada:
def sigmoid(z):
return np.where(z >= 0, 1.0 / (1.0 + np.exp(-z)),
np.exp(np.minimum(z, 0)) / (1.0 + np.exp(np.minimum(z, 0))))
def fit_logistic(X, y, lr=0.5, epochs=4000):
w, b = np.zeros(X.shape[1]), 0.0
for _ in range(epochs):
p = sigmoid(X @ w + b)
g = p - y
w -= lr * (X.T @ g) / len(y)
b -= lr * g.sum() / len(y)
return w, bEl np.where en sigmoid no es cosmético. Calcular directamente desborda con negativos grandes; la rama elige la forma algebraicamente idéntica que mantiene negativo el exponente. Esta es la caja de coma flotante del Capítulo 2 cobrando su primera deuda, y cobrará una mayor dentro de dos secciones.
Por qué no error cuadrático, y por qué la respuesta va del gradiente
Enlace a la sección: Por qué no error cuadrático, y por qué la respuesta va del gradienteLa explicación estándar para preferir la entropía cruzada al error cuadrático es el argumento de verosimilitud anterior: el error cuadrático es lo que obtienes al suponer ruido gaussiano; las etiquetas no son gaussianas; por tanto, no lo hagas. Es correcto y no convence a nadie, porque puedes escribir sobre una sigmoide y entrenará.
El argumento que sí cala va del gradiente. Pon error cuadrático encima de una sigmoide y la regla de la cadena da
Ese extra es el que se cancelaba antes. Ahora no lo hace, y se va a cero siempre que el modelo tiene confianza, incluso cuando el modelo está equivocado con confianza. Evalúa ambos en unas cuantas puntuaciones, para un ejemplo cuya etiqueta verdadera es 1:
| puntuación | entropía cruzada | error cuadrático | ratio | |
|---|---|---|---|---|
| 0,000335 | 1.491 | |||
| 0,017986 | 28,3 | |||
| 0,119203 | 4,8 | |||
| 0,500000 | 2,0 | |||
| 0,880797 | 4,8 |
En el modelo está tan equivocado como puede estarlo, y el error cuadrático responde con un gradiente 1.491 veces menor que el de la entropía cruzada. Cuanto peor es el error, menos aprende el modelo de él. El gradiente de la entropía cruzada, mientras tanto, se satura en : estar máximamente equivocado produce una señal máximamente grande, y no mayor.
Haz la carrera. Dos mil puntos equilibrados, pesos iniciales idénticos elegidos para estar equivocadamente seguros (), tasa de aprendizaje idéntica; solo cambia la pérdida. Ambas ejecuciones se puntúan con entropía cruzada para que las columnas sean comparables.
| época | pérdida de entropía cruzada | accuracy | pérdida de error cuadrático | accuracy |
|---|---|---|---|---|
| 1 | 5,4865 | 0,2300 | 5,9499 | 0,2290 |
| 10 | 1,5525 | 0,2460 | 5,9042 | 0,2290 |
| 50 | 0,4642 | 0,7780 | 5,6913 | 0,2320 |
| 100 | 0,4639 | 0,7770 | 5,3955 | 0,2410 |
| 200 | 0,4639 | 0,7770 | 4,6311 | 0,2745 |
| 500 | 0,4639 | 0,7770 | 0,5291 | 0,7660 |
| 1.000 | 0,4639 | 0,7770 | 0,4640 | 0,7765 |
La entropía cruzada ha terminado en la época 50. El error cuadrático sigue en un 24 % de accuracy en la época 100 —y no se había movido del 23 % en la época 10—, peor que adivinar, porque empezó equivocadamente seguro y el gradiente que lo habría rescatado se ha multiplicado por 0,0007. Escapa alrededor de la época 500 y aterriza en el mismo sitio. Así que el resumen honesto es que el error cuadrático sobre una sigmoide no es incorrecto; es lento exactamente donde más importa la velocidad. En un modelo de dos parámetros pierdes 450 épocas. En una red con cien capas, donde alguna unidad en algún lugar siempre está equivocadamente segura, pierdes la ejecución de entrenamiento.
Entropía, entropía cruzada y KL, en una página
Enlace a la sección: Entropía, entropía cruzada y KL, en una páginaTres cantidades, necesarias de verdad en el Capítulo 8 para la perplejidad y en el Capítulo 11 para la penalización que mantiene una política con fine-tuning cerca de su referencia. Son más fáciles de lo que dice su fama.2
Entropía es el número medio de bits que debes gastar para comunicar una muestra de una distribución, si usas el mejor código posible para ella:
Entropía cruzada es lo que gastas cuando usas un código construido para sobre datos que en realidad vienen de :
Divergencia KL es el exceso —el desperdicio, en bits— causado por creer cuando la verdad es :
Comprueba las tres en la cinta:
test defect rate = 0.0185
entropy of that coin = 0.1329 bits
cross-entropy of the constant predictor on test = 0.1330 bits
KL(test coin || fair coin) = 0.8671 bits
H + KL = 1.0000 bits
cross-entropy of the p=0.5 predictor on test = 1.0000 bitsAhí se ven dos cosas. Primero, un modelo que se limita a informar de la tasa base de entrenamiento, 1,69 %, consigue una entropía cruzada de 0,1330 bits, casi exactamente la entropía de las etiquetas de prueba, como debe ser, ya que tiene la distribución correcta y ninguna otra información. La entropía es el suelo que te compra la ignorancia sobre el individuo. Segundo, un modelo que se encoge de hombros y dice 0,5 paga exactamente 1 bit, y la brecha entre ambos, 0,8671 bits, es precisamente la divergencia KL. no es una identidad para memorizar; es una factura que puedes ver sumarse.
Y la conexión de vuelta con el entrenamiento: cuando la etiqueta es una única clase conocida, la distribución «verdadera» es one-hot, su entropía es cero y la entropía cruzada equivale a la divergencia KL. Minimizar la entropía cruzada y acercar la distribución del modelo a la verdad son el mismo acto.
Más de dos respuestas: softmax, y el desplazamiento que no cuesta nada
Enlace a la sección: Más de dos respuestas: softmax, y el desplazamiento que no cuesta nadaDefectuoso no es una sola cosa. En moldeo, una pieza puede salir como un short shot (material insuficiente), flash (demasiado material, expulsado del molde) o quemada. Cuatro resultados, así que cuatro logits, y deben convertirse en cuatro probabilidades que sumen uno. Eso es softmax:
Tiene una propiedad que parece un accidente y en realidad es toda la implementación:
para cualquier constante , porque y el se cancelan arriba y abajo. Solo las diferencias entre logits significan algo. El nivel absoluto no es información.
Por suerte, porque el nivel absoluto es lo que rompe el ordenador:
logits = [800. 801. 799.]
naive softmax = [nan nan nan]
shifted by -max = [0.2447 0.6652 0.09 ]
same softmax after adding 1000 to every logit: True desborda un float de 64 bits, la suma se convierte en infinito, e infinito dividido por infinito es nan: no un error, no un cuelgue, solo un agujero silencioso donde antes había tres probabilidades. Restar el logit máximo no cambia nada matemáticamente y lo cambia todo numéricamente, porque el exponente mayor se vuelve exactamente . Es el truco logsumexp del Capítulo 2 con ropa de trabajo, y toda implementación seria lo hace:
def softmax(Z):
Z = Z - Z.max(axis=1, keepdims=True)
E = np.exp(Z)
return E / E.sum(axis=1, keepdims=True)
def fit_softmax(X, Y, lr=1.0, epochs=6000):
W, b = np.zeros((X.shape[1], Y.shape[1])), np.zeros(Y.shape[1])
for _ in range(epochs):
G = (softmax(X @ W + b) - Y) / len(X)
W -= lr * (X.T @ G)
b -= lr * G.sum(0)
return W, bEl gradiente vuelve a ser predicción menos verdad, ahora con one-hot. El caso binario fue un caso especial desde el principio.
Entrenado con 3.000 piezas y probado con 1.000, con tres mediciones cada una (anchura, peso, temperatura de fusión), alcanza 94,00 % de accuracy. Esto es lo que ese número esconde:
| verdad ↓ / predicción → | ok | short shot | flash | quemada | recall |
|---|---|---|---|---|---|
| ok | 850 | 5 | 9 | 0 | 0,984 |
| short shot | 22 | 21 | 0 | 0 | 0,488 |
| flash | 20 | 0 | 30 | 1 | 0,588 |
| quemada | 3 | 0 | 0 | 39 | 0,929 |
| precision | 0,950 | 0,808 | 0,769 | 0,975 |
El modelo encuentra menos de la mitad de los short shots. La accuracy no puede verlo, porque el 86 % de las piezas están bien y acertarlas basta para sostener la media. Macro F1 —la media de las puntuaciones F1 por clase, que pondera una clase rara igual que una común— es 0,7983, frente a un micro F1 de 0,9400 que, por definición, es idéntico a la accuracy. Cuando alguien comunique un único número F1, pregunta cuál.
Esto es lo último del modelado. El resto del capítulo va de los números.
Tres modelos, una accuracy
Enlace a la sección: Tres modelos, una accuracyToma el modelo binario entrenado y crea dos variantes multiplicando cada logit por una constante: 0,35 para una versión dubitativa, 4 para una sobreconfiada. Multiplicar por un número positivo no puede cambiar ningún signo, así que los tres modelos predicen exactamente la misma etiqueta para las 4.000 piezas de prueba. La accuracy no puede distinguirlos. La entropía cruzada no tiene ningún problema:
| modelo | accuracy | entropía cruzada | pérdida media al acertar | pérdida media al fallar | peor pérdida individual |
|---|---|---|---|---|---|
| dubitativo (logits × 0,35) | 0,9830 | 0,1549 | 0,1369 | 1,1990 | 2,80 |
| tal como fue entrenado | 0,9830 | 0,0564 | 0,0147 | 2,4689 | 7,82 |
| sobreconfiado (logits × 4) | 0,9830 | 0,1563 | 0,0009 | 9,1427 | 27,63 |
El modelo dubitativo paga un pequeño impuesto por cada pieza, incluidas las miles que acierta. El sobreconfiado es casi gratis cuando acierta y catastrófico cuando falla: una pieza de ese conjunto de prueba le cuesta por sí sola 27,63 nats. Los dos llegan a casi el mismo total por rutas opuestas, y el modelo entrenado, cuyas probabilidades están calibradas con los datos, queda tres veces por debajo de ambos.
Esta es la forma más precisa de expresar la diferencia entre una pérdida y una métrica. La pérdida es lo que optimizas: debe ser diferenciable, y ve todo lo que dijo el modelo, incluida su seguridad. La métrica es aquello por lo que te juzgan: puede ser una función escalón, una regla de negocio, un recuento de defectos no detectados. No son el mismo objeto y no siempre coinciden; por eso defines ambas antes de empezar, y nunca dejas que la pérdida sustituya a la métrica solo porque esté en pantalla.
El baseline tonto va primero
Enlace a la sección: El baseline tonto va primeroAntes de cualquier modelo, el requisito: ¿qué puntuación obtiene la respuesta más perezosa posible? En esta cinta, decir siempre que está bien:
always-say-fine baseline: accuracy = 0.9815
confusion (tn, fp, fn, tp) = (3926, 0, 74, 0)98,15 %. Ahora el modelo logístico entrenado, con el umbral por defecto de 0,5:
logistic @0.5: accuracy=0.9830 precision=0.8000 recall=0.1081 F1=0.1905
confusion (tn, fp, fn, tp) = (3924, 2, 66, 8)98,30 %. Ha superado al baseline por 0,15 puntos porcentuales, y cualquier informe que se detenga en la accuracy lo llamará victoria. La matriz de confusión dice lo que ocurrió en realidad:
| predicho bien | predicho defectuoso | |
|---|---|---|
| realmente bien | 3.924 | 2 |
| realmente defectuoso | 66 | 8 |
Ha encontrado 8 piezas defectuosas de 74 y ha dejado pasar 66. Tres números nombran las tres formas de leer esa tabla:
- Precision . De las piezas que marcó, cuántas eran realmente defectuosas. Este es el coste de las inspecciones desperdiciadas.
- Recall . De las piezas defectuosas, cuántas capturó. Este es el coste de enviar una pieza defectuosa a un cliente.
- F1 , su media armónica, que se mantiene cerca de la menor de las dos y por tanto se niega a dejarse halagar por una sola.
Cuál importa depende de la fábrica, no de las matemáticas: una inspección cuesta unos segundos y un defecto enviado cuesta una notificación de retirada, así que aquí domina el recall y 0,108 es un fracaso.
Pero el problema no es el modelo. Es el umbral, y el umbral no forma parte del modelo: es una decisión de negocio aplicada después a una probabilidad. Barrámoslo:
| umbral | TP | FP | FN | accuracy | precision | recall | F1 |
|---|---|---|---|---|---|---|---|
| 0,500 | 8 | 2 | 66 | 0,9830 | 0,800 | 0,108 | 0,190 |
| 0,200 | 27 | 28 | 47 | 0,9812 | 0,491 | 0,365 | 0,419 |
| 0,100 | 42 | 118 | 32 | 0,9625 | 0,263 | 0,568 | 0,359 |
| 0,050 | 54 | 236 | 20 | 0,9360 | 0,186 | 0,730 | 0,297 |
| 0,020 | 67 | 570 | 7 | 0,8558 | 0,105 | 0,905 | 0,188 |
| 0,005 | 71 | 1.360 | 3 | 0,6593 | 0,050 | 0,959 | 0,094 |
Lee la columna de accuracy hacia abajo. Cae de principio a fin —del 98,30 % al 65,93 %— mientras el modelo pasa de detectar 8 defectos a detectar 71 de 74. Todo lo útil que puede hacer este modelo empeora su accuracy. Un equipo que optimizara el número del titular enviaría la versión que no encuentra nada.
Mostrar detalles
Ponderar clases no crea señal, mueve el punto de operación. El reflejo habitual con clases desbalanceadas es ponderar la clase rara en la pérdida. Al hacerlo, con pesos de 1, 10 y 60 en los positivos:
| peso en positivos | accuracy | precision | recall | F1 | AUC |
|---|---|---|---|---|---|
| 1 | 0,9830 | 0,800 | 0,108 | 0,190 | 0,9363 |
| 10 | 0,9605 | 0,253 | 0,581 | 0,352 | 0,9361 |
| 60 | 0,8290 | 0,091 | 0,919 | 0,166 | 0,9361 |
Precision y recall se mueven mucho. El AUC —la probabilidad de que el modelo coloque una pieza defectuosa aleatoria por encima de una buena aleatoria, ignorando por completo el umbral— se mueve 0,0002, es decir, nada. Reponderar deslizó el mismo modelo por la misma curva de compromiso. A menudo eso es lo que quieres, y nunca es información nueva: si el ranking es malo, ningún esquema de ponderación lo salvará.
Tres particiones, y la fuga que estás a punto de encontrar
Enlace a la sección: Tres particiones, y la fuga que estás a punto de encontrar¿Por qué tres particiones y no dos? Porque en el momento en que usas un conjunto de ejemplos para elegir algo —un umbral, una tasa de aprendizaje, cuál de seis modelos enviar— ese conjunto se ha usado para ajustar, y su puntuación deja de ser insesgada.3 Medido en esta cinta: barrer el umbral en el conjunto de validación elige 0,196, y el modelo luego obtiene F1 = 0,4122 en el conjunto de prueba intacto. Si el barrido se hubiera ejecutado directamente en el conjunto de prueba, el mejor valor alcanzable allí era 0,4186, un número que nadie tiene derecho a comunicar.
La brecha es pequeña aquí, 0,006, porque es un hiperparámetro barrido una vez contra 4.000 ejemplos de validación. Crece con cada decisión extra y con cada reducción del conjunto de validación. Observa también que la dirección no está garantizada en una ejecución concreta: el umbral elegido puntuó 0,3902 en validación y 0,4122 en prueba, así que esta vez la validación lo subestimó. El sesgo es sistemático a través de muchas decisiones, no visible en una sola.4
Ahora el ejercicio. El registro de la cinta llega con una tercera columna, station_seconds: cuánto tiempo pasó cada pieza en la estación de inspección. Añadirla es un cambio de una línea en el preprocesamiento. Esto es lo que hace:
| modelo | accuracy | precision | recall | F1 | entropía cruzada | AUC |
|---|---|---|---|---|---|---|
| anchura + peso | 0,9830 | 0,800 | 0,108 | 0,190 | 0,0564 | 0,9363 |
| + station_seconds | 0,9920 | 0,792 | 0,770 | 0,781 | 0,0236 | 0,9970 |
El recall pasa del 10,8 % al 77,0 %. F1 se multiplica por más de cuatro. Y fíjate en lo que hizo la accuracy: 98,30 % → 99,20 %, una ganancia de nueve décimas de punto, que es el tipo de número que en una diapositiva de resumen se redondea a «aproximadamente 99 % en ambos casos». La accuracy no vio el fallo antes y ahora no ve el fraude.
Antes de seguir leyendo: el modelo está haciendo trampas. Averigua cómo.
Cómo cazar una fuga, en el orden que la encuentra más rápido.
-
Compara entrenamiento y prueba. El overfitting aparece como una brecha grande. Aquí: modelo honesto 0,9838 entrenamiento / 0,9830 prueba; modelo con fuga 0,9936 entrenamiento / 0,9920 prueba. Ambas brechas están por debajo de 0,2 puntos. Una fuga no parece overfitting: la feature con fuga está igual de disponible en el momento de prueba, así que el modelo generaliza de maravilla a un mundo que no existe.
-
Entrena un modelo por feature, sola. Cualquier cosa que lleve la respuesta se anunciará:
feature sola accuracy recall F1 AUC anchura 0,9815 0,014 0,026 0,8691 peso 0,9815 0,000 0,000 0,7914 station_seconds0,9850 0,405 0,500 0,9960 Una columna, por sí sola, ordena defectos con AUC 0,9960. Dos mediciones tomadas con un calibre y una báscula consiguen 0,87 y 0,79. Esa asimetría es la alarma.
-
Pregunta cuándo se anotó cada número. Tiempo medio de permanencia: 2,23 segundos para las piezas que pasaron, 15,56 segundos para las que fallaron. Claro que sí. Una pieza permanece en la estación porque un inspector la retiró de la cinta, lo que ocurre después, y solo porque alguien decidió que era defectuosa. La columna no es una medición de la pieza. Es una medición del veredicto.
station = 1.8 + rng.exponential(0.35, N) # a part just passing through
audited = rng.random(N) < 0.006 # random spot checks
station[audited] += rng.uniform(6.0, 26.0, audited.sum())
station[y == 1] = 9.0 + rng.exponential(7.0, (y == 1).sum()) La línea resaltada es la fuga: el tiempo de permanencia de una pieza defectuosa se extrae de una distribución distinta, porque una persona la retiró de la cinta. Este es el bug grave más común en machine learning aplicado, y tiene nombre: fuga de objetivo —información en las features de entrenamiento que no estaría disponible en el momento en que hay que hacer la predicción—.5 No lanza ninguna excepción. Produce un número mejor. Todos los incentivos de un proyecto apuntan a conservarla.
La defensa es una pregunta, hecha a cada columna: en el instante en que necesito esta predicción, ¿existe ya este valor? En una cinta en vivo, station_seconds se desconoce hasta después de que la pieza haya sido inspeccionada, que es lo que se suponía que el modelo debía sustituir.
¿Cuántos ejemplos de prueba necesito?
Enlace a la sección: ¿Cuántos ejemplos de prueba necesito?Supón que puntúas un modelo con 20 ejemplos y acierta 17. Informas de un 85 %.
17 correct out of 20 -> accuracy 0.8500
Wilson 95% CI : [0.6396, 0.9476]
bootstrap 95% CI : [0.7000, 1.0000]
P(a 65% model scores 17 or more out of 20) = 0.0444
P(an 85% model scores 17 or more out of 20) = 0.6477La lectura honesta de 17/20 es en algún lugar entre el 64 % y el 95 %. Un modelo realmente del 65 % produce este resultado el 4,4 % de las veces —una ejecución de cada veintitrés—, y si probaste un puñado de prompts y comunicaste el mejor, fabricaste tú mismo esa ejecución. Diecisiete de veinte no puede distinguir un modelo del 85 % de uno del 65 %.
Dos formas de poner un intervalo a una tasa, y ambas pertenecen a tu caja de herramientas:
def wilson(k, n, z=1.959963985):
"""95% interval for k successes in n trials. Correct at small n; no simulation."""
ph, d = k / n, 1 + z * z / n
centre = (ph + z * z / (2 * n)) / d
half = z * (ph * (1 - ph) / n + z * z / (4 * n * n)) ** 0.5 / d
return centre - half, centre + half
def bootstrap_ci(correct, n_resamples=10_000, alpha=0.05, seed=0):
"""95% interval for the mean of any per-example score array. Works on F1 too."""
rng = np.random.default_rng(seed)
correct = np.asarray(correct, dtype=float)
draws = correct[rng.integers(0, len(correct), size=(n_resamples, len(correct)))]
lo, hi = np.quantile(draws.mean(axis=1), [alpha / 2, 1 - alpha / 2])
return float(correct.mean()), float(lo), float(hi)Usa Wilson6 para una tasa de éxito simple; se comporta bien con cualquier y no necesita aleatoriedad. Observa arriba que con el extremo superior del bootstrap es 1,0000: al remuestrear 20 puntos es fácil extraer 20 correctos, así que no puede representar un intervalo más estrecho que su propia granularidad. Usa el bootstrap7 cuando no exista fórmula, que es la mayoría de los casos interesantes: F1, macro-medias, BLEU, pass@1, la puntuación de un juez basado en rúbrica. En esta cinta, el F1 de 0,4122 del modelo ajustado lleva un intervalo bootstrap de [0,3009, 0,5156]; ese es el número que debería aparecer en el informe, porque la estimación puntual por sí sola invita a una comparación que no puede sostener.
Una medición más, porque cambia cómo deberías comparar dos modelos. Dos modelos puntuados sobre los mismos 500 ejemplos:
model A: 0.8580 95% CI [0.8260, 0.8880]
model B: 0.8120 95% CI [0.7780, 0.8460]
the two intervals overlap: True
paired difference A-B: 0.0460 95% CI [0.0260, 0.0680]
they disagree on 31 of 500 examples (A right 27, B right 4)Sus intervalos se solapan, y la regla popular —barras de error solapadas significa que no hay diferencia significativa— llamaría inconclusa la comparación. No lo es. Los dos modelos se ejecutaron sobre los mismos ejemplos, así que la cantidad correcta es la diferencia por ejemplo, cuyo intervalo es [0.0260, 0.0680], cómodamente por encima de cero. Solo discrepan en 31 de 500 elementos, y A gana 27 de esos desacuerdos; los ejemplos compartidos, fáciles y difíciles por igual, se cancelan en lugar de añadir ruido. Compara modelos de forma pareada y llegarás a la misma conclusión con una fracción de los datos.
Hacia dónde va esto ahora
Enlace a la sección: Hacia dónde va esto ahoraAhora tienes un modelo que emite probabilidades calibradas, una pérdida derivada de una afirmación sobre los datos en lugar de elegida por comodidad, un gradiente que es literalmente predicción menos verdad y —más importante— la maquinaria para averiguar si algo de eso funciona. El intervalo de Wilson de diez líneas de arriba se reutiliza literalmente: sostiene las variantes de prompt en el Capítulo 15, las tablas de recuperación en el Capítulo 19 y el conjunto dorado en el Capítulo 29. El bootstrap es a lo que recurres cuando no existe fórmula.
Pero el modelo sigue teniendo una sola capa. Dibuja una línea, y el Capítulo 1 demostró con cuatro filas de XOR que una línea no basta. La solución es apilar: una primera capa que curva el espacio, una segunda que dibuja la línea en el espacio curvado.
Ahí es donde se agota el gradiente limpio de este capítulo. Todo lo anterior funcionó porque podía escribirse a mano, una vez, para un modelo con una capa entre la entrada y la pérdida. Pon una segunda capa en medio y la pregunta cambia de forma: ¿cuál es la derivada de la pérdida respecto a un peso que no toca la salida en absoluto, uno cuya influencia llega solo a través de otra capa, quizá por varios caminos a la vez?
Esa derivada existe. Calcularla a mano es inviable para cualquier cosa mayor que un juguete, y calcularla parámetro a parámetro es inviable a otra escala. Lo que se necesita es un procedimiento que obtenga todas las derivadas de la red a partir de una única pasada hacia atrás sobre el mismo grafo que la pasada hacia delante acaba de recorrer.
Eso es el Capítulo 5, y es el motor sobre el que corre el resto de este curso.
Fuentes y método
Enlace a la sección: Fuentes y métodoTambién merece la pena leer junto a este capítulo: Bishop, Pattern Recognition and Machine Learning §1.2, §1.5, §1.6 y §4.3, que cubre probabilidad, teoría de la decisión, teoría de la información y clasificación lineal en el orden que sigue este capítulo; Murphy, Probabilistic Machine Learning: An Introduction, capítulos 6 y 10; Prince, Understanding Deep Learning §5.4–5.7; y Saito y Rehmsmeier, The Precision-Recall Plot Is More Informative than the ROC Plot When Evaluating Binary Classifiers on Imbalanced Datasets (PLOS ONE, 2015): por qué el AUC citado arriba no debería ser el único número independiente del umbral que miras cuando el 1,7 % de las piezas son defectuosas.
Referencias
Enlace a la sección: Referencias-
Ma, T. y Ng, A. CS229 Lecture Notes, Stanford University, capítulos 2 y 3. Donde la cancelación que produce deja de parecer suerte: elige la distribución de la familia exponencial que encaja con tu salida, usa su enlace canónico, y el gradiente siempre es predicción menos verdad. ↩
-
Olah, C. Visual Information Theory (2015),
colah.github.io/posts/2015-09-Visual-Information. La explicación más clara disponible de entropía, entropía cruzada y divergencia KL como costes en bits en lugar de como fórmulas. ↩ -
Abu-Mostafa, Y. S., Magdon-Ismail, M. y Lin, H.-T. Learning From Data (AMLBook, 2012), clases 13 y 17 del curso de Caltech. La clase 13 es validación; la clase 17, sobre los tres principios del aprendizaje, es donde se nombra el data snooping. Entre ambas son la fuente de la disciplina de este capítulo: cada mirada a un conjunto de datos es una decisión de ajuste, hayas ejecutado o no un optimizador. ↩
-
James, G., Witten, D., Hastie, T. y Tibshirani, R. An Introduction to Statistical Learning, 2.ª edición (Springer, 2021), capítulos 2 y 5, para la descomposición sesgo-varianza y para el remuestreo. El volumen complementario es donde la trampa de selección se formula sin rodeos: Hastie, Tibshirani y Friedman, The Elements of Statistical Learning, 2.ª edición, §7.10.2, The Wrong and Right Way to Do Cross-validation. ↩
-
Kaufman, S., Rosset, S., Perlich, C. y Stitelman, O. Leakage in Data Mining: Formulation, Detection, and Avoidance. ACM Transactions on Knowledge Discovery from Data 6(4), 2012. Un tratamiento formal del fallo demostrado arriba, con estudios de caso de competiciones ganadas por un modelo que había aprendido un artefacto de cómo se ensamblaron los datos. ↩
-
Wilson, E. B. Probable Inference, the Law of Succession, and Statistical Inference. Journal of the American Statistical Association 22(158), pp. 209–212 (1927). El intervalo de puntuación usado en
wilson()arriba, todavía el valor por defecto correcto para una proporción. El intervalo de manual es el que hay que evitar: da sinsentidos cerca de 0 y 1, y subcubre gravemente con pequeños. ↩ -
Efron, B. Bootstrap Methods: Another Look at the Jackknife. The Annals of Statistics 7(1), pp. 1–26 (1979). La idea que te permite poner un intervalo a cualquier estadístico que puedas calcular, incluidos los que no tienen teoría muestral. ↩