Saltar ao contido
4/30Capítulo 4 de 30

Clasificación, entropía cruzada e como non enganarte a ti mesmo

Constrúe un clasificador loxístico e descubre por que un 98 % de exactitude pode non atopar nada.

Nesta páxina

Un modelo que responde esta peza está ben sobre cada peza que sae da cinta acerta o 98.15 % das veces. E tamén non vale para nada: das 74 pezas defectuosas do conxunto de proba, non colle ningunha.

As dúas frases describen o mesmo modelo. A distancia entre elas é este capítulo.

A primeira metade constrúe o clasificador. Non necesita case nada novo: o Capítulo 2 deu a receita para converter unha suposición sobre como se producen os datos nunha función de perda, e o Capítulo 3 deu a maquinaria para baixar pola pendente de calquera perda que esa receita che entregue. Aplica ambas a unha pregunta de si/non e aparece a regresión loxística, máis unha idea nova — un logit — que volverá pasar factura no Capítulo 17.

A segunda metade é a difícil. Todo o que vén despois neste curso xúlgase por un número que alguén mediu, e se non sabes distinguir unha mellora real dun artefacto de medición, cada capítulo seguinte é decoración. Así que: a matriz de confusión, precisión e recall, as tres particións, a leakage, e a pregunta que case ninguén responde con honestidade — cantos exemplos de proba necesito realmente?

A aritmética aquí percorre 20,000 filas, así que está vectorizada de principio a fin — NumPy leva facendo o traballo desde o Capítulo 2, e a partir de aquí deixa de merecer a pena comentalo.

A mesma fábrica ca no Capítulo 1, unha pregunta máis difícil. En vez de aceptar ou rexeitar, a pregunta é esta peza é defectuosa — e os defectos son raros, o que fai difícil a metade de medición deste capítulo e enganadoramente doada a metade de modelado.

belt.pyPYTHON
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:]
TEXT
N = 20000  defects = 337  base rate = 0.0169
defects per split = 203 60 74

Tres particións, non dúas. A razón merece a súa propia sección e tena máis abaixo; polo de agora, adestra coa primeira, axusta coa segunda e non mires a terceira.

As features están estandarizadas — réstase a media e divídese pola desviación estándar — usando só as estatísticas de adestramento, pola razón que o Capítulo 1 demostrou coa cota de converxencia do perceptrón: os datos non centrados fan hostil a xeometría. De que filas tes permiso para calcular esa media convértese nunha pregunta viva máis adiante neste capítulo.

O perceptrón devolvía un signo. Un signo non pode distinguir rexeitar de rexeitar, pero por moi pouco, e esa diferenza é exactamente o que unha fábrica necesita para decidir que pezas debería reinspeccionar antes unha persoa.

Así que segue literalmente a receita do Capítulo 2. Escribe o que afirmas sobre como se produce unha etiqueta, toma a verosimilitude, toma o log, négano, e tes unha perda. Para un resultado si/non, a afirmación é unha distribución Bernoulli: hai unha probabilidade pp de que a peza sexa defectuosa, e

P(yp)=py(1p)1yP(y \mid p) = p^{\,y}\,(1-p)^{\,1-y}

que é só unha forma compacta de escribir «pp se y=1y = 1, e 1p1-p se y=0y = 0». Toma o log diso e négano, e a perda para un exemplo é

L=[ylogp+(1y)log(1p)]L = -\big[\,y \log p + (1 - y)\log(1 - p)\,\big]

Isto é entropía cruzada binaria. Non se escolleu porque sexa cómoda; é a log-verosimilitude negativa da única distribución que pode ter o lanzamento dunha moeda. Non había outra cousa dispoñible.

O que aínda falta é de onde sae pp. O modelo calcula unha suma ponderada s=wx+bs = \mathbf{w}\cdot\mathbf{x} + b, que é un número real e percorre toda a recta, e unha probabilidade ten que vivir en (0,1)(0,1). A función que se move entre ambos é a sigmoide loxística:

σ(s)=11+es\sigma(s) = \frac{1}{1 + e^{-s}}
TEXT
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.9820

Le a columna da dereita como unha lista de prezos. Ter razón cun 90 % de confianza custa 0.105. Negarse a comprometerse custa 0.693 — que é log2\log 2, o prezo dun encollerse de ombros. Estar equivocadamente seguro custa 4.6, corenta e catro veces máis, e o prezo sobe sen límite a medida que o modelo está máis seguro dun erro. A entropía cruzada non conta simplemente erros: cobra pola arrogancia.

O Capítulo 3 dixo: para adestrar calquera cousa, obtén a derivada da perda respecto de cada parámetro. Faino para un exemplo. Con s=wx+bs = \mathbf{w}\cdot\mathbf{x} + b e p=σ(s)p = \sigma(s):

Ls=py,Lw=(py)x,Lb=py\frac{\partial L}{\partial s} = p - y, \qquad \frac{\partial L}{\partial \mathbf{w}} = (p - y)\,\mathbf{x}, \qquad \frac{\partial L}{\partial b} = p - y
Mostrar detalles

As dúas liñas que fan que a lea se cancele. A sigmoide ten unha derivada inusualmente agradable, σ(s)=σ(s)(1σ(s))=p(1p)\sigma'(s) = \sigma(s)\,(1 - \sigma(s)) = p(1-p). E a perda deriva a

Lp=yp+1y1p=pyp(1p)\frac{\partial L}{\partial p} = -\frac{y}{p} + \frac{1-y}{1-p} = \frac{p - y}{p\,(1-p)}

Multiplica as dúas pola regra da cadea e o p(1p)p(1-p) aparece unha vez arriba e unha vez abaixo. Cancélase exactamente, e pyp - y é o que sobrevive. Esa cancelación non é unha coincidencia — é o que pasa sempre que a perda é a log-verosimilitude negativa dunha distribución e a función de saída é a que esa distribución usa de forma natural. Ese emparellamento ten nome — un modelo lineal xeneralizado — e o gradiente limpo é a súa pegada.1

Así que a actualización é predición menos verdade, multiplicado pola entrada. Nada máis. Aquí está todo o adestrador, que é o descenso do Capítulo 3 cunha liña cambiada:

logistic.pyPYTHON
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, b

O np.where en sigmoid non é cosmético. Calcular 1/(1+es)1/(1+e^{-s}) directamente desborda para ss negativos grandes; a rama escolle a forma alxebraicamente idéntica que mantén o expoñente negativo. É a caixa de punto flotante do Capítulo 2 cobrando a súa primeira débeda, e cobrará unha maior dentro de dúas seccións.

Por que non erro cadrático, e por que a resposta vai sobre o gradiente

Ligazón á sección: Por que non erro cadrático, e por que a resposta vai sobre o gradiente

A explicación estándar para preferir a entropía cruzada ao erro cadrático é o argumento da verosimilitude de arriba: o erro cadrático é o que obtés ao asumir ruído gaussiano, as etiquetas non son gaussianas, polo tanto non o fagas. É correcto e non convence a ninguén, porque podes escribir L=(py)2L = (p - y)^2 sobre unha sigmoide e adestrará.

O argumento que chega é sobre o gradiente. Pon erro cadrático enriba dunha sigmoide e a regra da cadea dá

Ls=2(py)p(1p)\frac{\partial L}{\partial s} = 2\,(p - y)\,p\,(1-p)

Ese p(1p)p(1-p) extra é o que antes se cancelaba. Agora non, e vai a cero sempre que o modelo está confiado — incluído cando o modelo está confiada e equivocadamente. Avalía ambos nalgúns scores, para un exemplo cuxa etiqueta verdadeira é 1:

score ssppentropía cruzada L/s\partial L/\partial serro cadrático L/s\partial L/\partial sratio
8-80.0003350.999665-0.9996650.000670-0.0006701,491
4-40.0179860.982014-0.9820140.034690-0.03469028.3
2-20.1192030.880797-0.8807970.184956-0.1849564.8
000.5000000.500000-0.5000000.250000-0.2500002.0
+2+20.8807970.119203-0.1192030.025031-0.0250314.8

En s=8s = -8 o modelo está tan equivocado como é posible estar, e o erro cadrático responde cun gradiente 1,491 veces menor ca o da entropía cruzada. Canto peor é o erro, menos aprende o modelo del. O gradiente da entropía cruzada, pola contra, satura en 1-1: estar maximalmente equivocado produce un sinal maximalmente grande, e non maior.

Corre a carreira. Dous mil puntos equilibrados, pesos iniciais idénticos escollidos para estar equivocadamente confiados (w=[6,6]\mathbf{w} = [-6, -6]), learning rate idéntico, só cambia a perda. Ambas execucións puntúanse con entropía cruzada para que as columnas sexan comparables.

epochperda de entropía cruzadaexactitudeperda con erro cadráticoexactitude
15.48650.23005.94990.2290
101.55250.24605.90420.2290
500.46420.77805.69130.2320
1000.46390.77705.39550.2410
2000.46390.77704.63110.2745
5000.46390.77700.52910.7660
1,0000.46390.77700.46400.7765

A entropía cruzada remata no epoch 50. O erro cadrático segue no 24 % de exactitude no epoch 100 — e non se movera do 23 % no epoch 10 — peor ca adiviñar, porque empezou equivocadamente confiado e o gradiente que o rescataría foi multiplicado por 0.0007. Escapa arredor do epoch 500 e chega ao mesmo sitio. Así que o resumo honesto é que o erro cadrático sobre unha sigmoide non é incorrecto; é lento exactamente onde máis importa a velocidade. Nun modelo de dous parámetros perdes 450 epochs. Nunha rede con cen capas, onde algunha unidade nalgún sitio sempre está equivocadamente confiada, perdes a execución de adestramento.

Entropía, entropía cruzada e KL, nunha páxina

Ligazón á sección: Entropía, entropía cruzada e KL, nunha páxina

Tres cantidades, necesarias de verdade no Capítulo 8 para a perplexity e no Capítulo 11 para a penalización que mantén unha política con fine-tuning preto da súa referencia. Son máis doadas do que di a súa fama.2

Entropía é o número medio de bits que debes gastar para comunicar unha mostra dunha distribución, se usas o mellor código posible para ela:

H(p)=ipilog2piH(p) = -\sum_i p_i \log_2 p_i

Entropía cruzada é o que gastas cando usas un código construído para qq sobre datos que en realidade veñen de pp:

H(p,q)=ipilog2qiH(p, q) = -\sum_i p_i \log_2 q_i

Diverxencia KL é o exceso — o desperdicio, en bits, causado por crer qq cando a verdade é pp:

DKL(pq)=H(p,q)H(p)D_{\mathrm{KL}}(p \parallel q) = H(p,q) - H(p)

Comproba as tres na cinta:

TEXT
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 bits

Aí vense dúas cousas. Primeiro, un modelo que simplemente informa da taxa base de adestramento, 1.69 %, acada unha entropía cruzada de 0.1330 bits, case exactamente a entropía das etiquetas de proba — como debe ser, xa que ten a distribución correcta e ningunha outra información. A entropía é o chan que che compra a ignorancia sobre o individuo. Segundo, un modelo que se encolle de ombros e di 0.5 paga exactamente 1 bit, e a fenda entre os dous, 0.8671 bits, é precisamente a diverxencia KL. H+DKL=H(p,q)H + D_{\mathrm{KL}} = H(p,q) non é unha identidade para memorizar; é unha factura que podes ver sumarse.

E a conexión de volta co adestramento: cando a etiqueta é unha única clase coñecida, a distribución «verdadeira» é one-hot, a súa entropía é cero, e a entropía cruzada iguala a diverxencia KL. Minimizar a entropía cruzada e empurrar a distribución do modelo cara á verdade son o mesmo acto.

Máis de dúas respostas: softmax, e o desprazamento que non custa nada

Ligazón á sección: Máis de dúas respostas: softmax, e o desprazamento que non custa nada

Defectuoso non é unha soa cousa. No moldeado, unha peza pode saír como short shot (material insuficiente), flash (demasiado, expulsado do molde), ou queimadura. Catro resultados, así que catro logits, e deben converterse en catro probabilidades que sumen un. Iso é softmax:

softmax(z)i=ezijezj\operatorname{softmax}(\mathbf{z})_i = \frac{e^{z_i}}{\sum_j e^{z_j}}

Ten unha propiedade que parece un accidente e en realidade é toda a implementación:

softmax(z+c)=softmax(z)\operatorname{softmax}(\mathbf{z} + c) = \operatorname{softmax}(\mathbf{z})

para calquera constante cc, porque ezi+c=ecezie^{z_i + c} = e^{c} e^{z_i} e o ece^c cancélanse arriba e abaixo. Só as diferenzas entre logits significan algo. O nivel absoluto non é información.

Por sorte é así, porque o nivel absoluto é o que rompe o computador:

TEXT
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

e800e^{800} desborda un float de 64 bits, a suma convértese en infinito, e infinito dividido por infinito é nan — non un erro, non un crash, só un burato silencioso onde antes había tres probabilidades. Restar o logit máximo non cambia nada matematicamente e cámbiao todo numericamente, porque o maior expoñente convértese exactamente en e0=1e^0 = 1. É o truco logsumexp do Capítulo 2 co mono de traballo posto, e toda implementación seria faino:

softmax.pyPYTHON
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, b

O gradiente volve ser predición menos verdade, agora con YY one-hot. O caso binario foi un caso especial desde o principio.

Adestrado con 3,000 pezas e probado con 1,000, con tres medicións cada unha (anchura, peso, temperatura de fusión), alcanza 94.00 % de exactitude. Isto é o que ese número agocha:

verdade ↓ / previsto →okshort shotflashburnrecall
ok8505900.984
short shot2221000.488
flash2003010.588
burn300390.929
precisión0.9500.8080.7690.975

O modelo atopa menos da metade dos short shots. A exactitude non pode velo, porque o 86 % das pezas están ben e acertalas abonda para soster a media. Macro F1 — a media dos F1 por clase, que pesa unha clase rara igual ca unha común — é 0.7983, fronte a un micro F1 de 0.9400 que por definición é idéntico á exactitude. Sempre que alguén informe dun único número F1, pregunta cal.

Ese é o final do modelado. O resto do capítulo vai sobre os números.

Colle o modelo binario adestrado e fai dúas variantes multiplicando cada logit por unha constante: 0.35 para unha versión dubidosa, 4 para unha sobreconfiada. Multiplicar por un número positivo non pode cambiar ningún signo, así que os tres modelos predín exactamente a mesma etiqueta para as 4,000 pezas de proba. A exactitude non pode distinguilos. A entropía cruzada non ten ningún problema:

modeloexactitudeentropía cruzadaperda media cando acertaperda media cando fallapeor perda única
dubidoso (logits × 0.35)0.98300.15490.13691.19902.80
tal como se adestrou0.98300.05640.01472.46897.82
sobreconfiado (logits × 4)0.98300.15630.00099.142727.63

O modelo dubidoso paga un pequeno imposto en cada peza, incluídas as miles que acerta. O sobreconfiado é case gratis cando acerta e catastrófico cando falla — unha peza dese conxunto de proba cústalle 27.63 nats ela soa. Os dous chegan case ao mesmo total por camiños opostos, e o modelo adestrado, cuxas probabilidades están calibradas cos datos, queda tres veces por baixo de ambos.

Esta é a forma máis nítida de expresar a diferenza entre unha perda e unha métrica. A perda é o que optimizas: debe ser diferenciable, e ve todo o que dixo o modelo, incluído o seguro que estaba. A métrica é aquilo polo que te xulgan: pode ser unha función chanzo, unha regra de negocio, un reconto de defectos perdidos. Non son o mesmo obxecto e non sempre están de acordo — por iso defines ambas antes de empezar, e nunca deixas que a perda substitúa a métrica só porque apareza na pantalla.

Antes de calquera modelo, o requisito: que puntuación obtén a resposta máis preguiceira posible? Nesta cinta, dicir sempre que está ben:

TEXT
always-say-fine baseline: accuracy = 0.9815
confusion (tn, fp, fn, tp) = (3926, 0, 74, 0)

98.15 %. Agora o modelo loxístico adestrado, co limiar por defecto de 0.5:

TEXT
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 %. Superou a baseline por 0.15 puntos porcentuais, e calquera informe que pare na exactitude chamará a iso unha vitoria. A matriz de confusión di o que pasou de verdade:

previsto benprevisto defectuoso
realmente ben3,9242
realmente defectuoso668

Atopou 8 pezas defectuosas de 74 e deixou pasar 66. Tres números nomean as tres formas de ler esa táboa:

  • Precisión =TP/(TP+FP)=8/10=0.800= \mathrm{TP}/(\mathrm{TP}+\mathrm{FP}) = 8/10 = 0.800. Das pezas que marcou, cantas eran realmente defectuosas. Este é o custo das inspeccións desperdiciadas.
  • Recall =TP/(TP+FN)=8/74=0.108= \mathrm{TP}/(\mathrm{TP}+\mathrm{FN}) = 8/74 = 0.108. Das pezas defectuosas, cantas capturou. Este é o custo de enviar unha peza mala a un cliente.
  • F1 =2PR/(P+R)=0.190= 2PR/(P+R) = 0.190, a súa media harmónica, que queda preto do menor dos dous e polo tanto rexeita deixarse afagar por un deles só.

Cal importa depende da fábrica, non das matemáticas: unha inspección custa uns segundos e un defecto enviado custa un aviso de retirada, así que aquí domina o recall e 0.108 é un fracaso.

Pero o problema non é o modelo. É o limiar, e o limiar non forma parte do modelo — é unha decisión de negocio aplicada despois a unha probabilidade. Varreo:

limiarTPFPFNexactitudeprecisiónrecallF1
0.50082660.98300.8000.1080.190
0.2002728470.98120.4910.3650.419
0.10042118320.96250.2630.5680.359
0.05054236200.93600.1860.7300.297
0.0206757070.85580.1050.9050.188
0.005711,36030.65930.0500.9590.094

Le a columna de exactitude cara abaixo. Cae todo o tempo — de 98.30 % a 65.93 % — mentres o modelo pasa de capturar 8 defectos a capturar 71 de 74. Cada cousa útil que este modelo pode facer empeora a súa exactitude. Un equipo que optimizase o número do titular enviaría a versión que non atopa nada.

Mostrar detalles

A ponderación de clases non crea sinal, move o punto de operación. O primeiro reflexo habitual con clases desbalanceadas é ponderar a clase rara na perda. Facéndoo, con pesos de 1, 10 e 60 nos positivos:

peso nos positivosexactitudeprecisiónrecallF1AUC
10.98300.8000.1080.1900.9363
100.96050.2530.5810.3520.9361
600.82900.0910.9190.1660.9361

Precisión e recall móvense moito. A AUC — a probabilidade de que o modelo ordene unha peza defectuosa aleatoria por riba dunha boa aleatoria, ignorando por completo o limiar — móvese 0.0002, que non é nada. A reponderación desprazou o mesmo modelo pola mesma curva de intercambio. Iso é a miúdo o que queres, e nunca é información nova: se o ranking é malo, ningún esquema de ponderación o salvará.

Tres particións, e a leak que estás a piques de atopar

Ligazón á sección: Tres particións, e a leak que estás a piques de atopar

Por que tres particións e non dúas? Porque no momento en que usas un conxunto de exemplos para escoller algo — un limiar, un learning rate, cal dos seis modelos enviar — ese conxunto usouse para axustar, e a súa puntuación deixa de ser non nesgada.3 Medido nesta cinta: varrer o limiar no conxunto de validación escolle 0.196, e o modelo logo puntúa F1 = 0.4122 no conxunto de proba intacto. Se o varrido se fixese directamente no conxunto de proba, o mellor acadable alí era 0.4186 — un número que ninguén ten dereito a informar.

A fenda é pequena aquí, 0.006, porque é un hiperparámetro varrido unha vez contra 4,000 exemplos de validación. Medra con cada decisión extra e con cada redución do conxunto de validación. Observa tamén que a dirección non está garantida nunha soa execución: o limiar escollido puntuou 0.3902 en validación e 0.4122 en proba, así que a validación infravalorouno esta vez. O nesgo é sistemático ao longo de moitas decisións, non visible nunha.4

Agora o exercicio. O rexistro da cinta chega cunha terceira columna, station_seconds: canto tempo pasou cada peza na estación de inspección. Engadila é un cambio dunha liña no preprocesamento. Isto é o que fai:

modeloexactitudeprecisiónrecallF1entropía cruzadaAUC
anchura + peso0.98300.8000.1080.1900.05640.9363
+ station_seconds0.99200.7920.7700.7810.02360.9970

O recall pasa do 10.8 % ao 77.0 %. O F1 máis ca cuadruplica. E fíxate no que fixo a exactitude: 98.30 % → 99.20 %, unha ganancia de nove décimas de punto, que é o tipo de número que nunha diapositiva de resumo se redondea a «arredor do 99 % en calquera caso». A exactitude non viu antes o fracaso e agora non ve a fraude.

Antes de seguir lendo: o modelo está facendo trampas. Descubre como.

Como cazar unha leak, na orde que a atopa máis rápido.

  1. Compara train e test. O overfitting aparece como unha fenda grande. Aquí: modelo honesto 0.9838 train / 0.9830 test; modelo con leak 0.9936 train / 0.9920 test. Ambas fendas están por baixo de 0.2 puntos. Unha leak non se parece ao overfitting — a feature con leak está igual de dispoñible no test, así que o modelo xeneraliza de marabilla a un mundo que non existe.

  2. Adestra un modelo por feature, soa. Calquera cousa que leve a resposta anunciarase:

    feature soaexactituderecallF1AUC
    anchura0.98150.0140.0260.8691
    peso0.98150.0000.0000.7914
    station_seconds0.98500.4050.5000.9960

    Unha columna, por si soa, ordena defectos cun AUC de 0.9960. Dúas medicións tomadas cun calibre e unha báscula conseguen 0.87 e 0.79. Esa asimetría é a alarma.

  3. Pregunta cando se anotou cada número. Tempo medio de permanencia: 2.23 segundos para pezas que pasaron, 15.56 segundos para pezas que fallaron. Claro que si. Unha peza permanece na estación porque un inspector a retirou da cinta — o que ocorre despois, e só porque, alguén decidiu que era defectuosa. A columna non é unha medición da peza. É unha medición do veredicto.

the planted leakPYTHON
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())  

A liña destacada é a leak: o tempo de permanencia dunha peza defectuosa extráese dunha distribución diferente, porque unha persoa a retirou da cinta. Este é o erro serio máis común no machine learning aplicado, e ten nome: target leakage — información nas features de adestramento que non estaría dispoñible no momento no que hai que facer a predición.5 Non lanza ningunha excepción. Produce un número mellor. Todos os incentivos dun proxecto apuntan a mantelo.

A defensa é unha pregunta, feita a cada columna: no instante en que necesito esta predición, este valor existe xa? Nunha cinta en vivo, station_seconds é descoñecido ata despois de que a peza fose inspeccionada — que é precisamente o que se supuña que o modelo ía substituír.

Supoñamos que puntúas un modelo en 20 exemplos e acerta 17. Informas 85 %.

TEXT
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.6477

A lectura honesta de 17/20 é nalgún lugar entre 64 % e 95 %. Un modelo realmente do 65 % produce este resultado o 4.4 % das veces — unha execución de cada vinte e tres — e se probaches un feixe de prompts e informaches o mellor, fabricaches ti mesmo esa execución. Dezasete de vinte non distinguen un modelo do 85 % dun do 65 %.

Dúas formas de poñer un intervalo nunha taxa, e ambas pertencen á túa caixa de ferramentas:

uncertainty.pyPYTHON
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 unha taxa de éxito simple; compórtase ben con calquera nn e non necesita aleatoriedade. Observa arriba que en n=20n = 20 o extremo superior do bootstrap é 1.0000 — ao remostrar 20 puntos é doado extraer 20 correctos, así que non pode representar un intervalo máis estreito ca a súa propia granularidade. Usa o bootstrap7 onde non existe unha fórmula, que é a maioría dos casos interesantes: F1, medias macro, BLEU, pass@1, a puntuación dun xuíz baseado nunha rúbrica. Nesta cinta, o F1 de 0.4122 do modelo axustado leva un intervalo bootstrap de [0.3009, 0.5156] — que é o número que debería aparecer no informe, porque a estimación puntual soa invita a unha comparación que non pode soster.

Unha medición máis, porque cambia como deberías comparar dous modelos. Dous modelos puntuados nos mesmos 500 exemplos:

TEXT
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)

Os seus intervalos solápanse, e a regra popular — barras de erro solapadas significa que non hai diferenza significativa — chamaría inconclusiva a comparación. Non o é. Os dous modelos correron sobre os mesmos exemplos, así que a cantidade correcta é a diferenza por exemplo, cuxo intervalo é [0.0260, 0.0680], comodamente por riba de cero. Só discrepan en 31 de 500 ítems, e A gaña 27 deses desacordos; os exemplos compartidos, doados e difíciles por igual, cancélanse en vez de engadir ruído. Compara modelos emparellados, e chegas á mesma conclusión cunha fracción dos datos.

Agora tes un modelo que emite probabilidades calibradas, unha perda derivada dunha afirmación sobre os datos e non escollida por comodidade, un gradiente que é literalmente predición menos verdade, e — máis importante — a maquinaria para descubrir se algo diso funciona. O intervalo Wilson de dez liñas de arriba reutilízase literalmente: sostén as variantes de prompt no Capítulo 15, as táboas de retrieval no Capítulo 19, e o golden set no Capítulo 29. O bootstrap é ao que recorres cando non existe unha fórmula.

Pero o modelo segue tendo unha soa capa. Debuxa unha liña, e o Capítulo 1 probou con catro filas de XOR que unha liña non abonda. A solución é apilar: unha primeira capa que dobra o espazo, unha segunda que debuxa a liña no espazo dobrado.

Aí é onde se esgota o gradiente ordenado deste capítulo. Todo o anterior funcionou porque L/s=py\partial L/\partial s = p - y podía escribirse á man, unha vez, para un modelo cunha capa entre a entrada e a perda. Pon unha segunda capa no medio e a pregunta cambia de forma: cal é a derivada da perda respecto dun peso que non toca a saída en absoluto — un peso cuxa influencia chega só a través doutra capa, posiblemente por varios camiños á vez?

Esa derivada existe. Calculala á man é imposible para calquera cousa máis grande ca un xoguete, e calculala parámetro a parámetro é imposible noutra escala. O que se necesita é un procedemento que obteña cada derivada da rede cunha soa pasada cara atrás sobre o mesmo grafo que a pasada cara adiante acaba de percorrer.

Iso é o Capítulo 5, e é o motor no que funciona o resto deste curso.


Tamén paga a pena ler xunto con este capítulo: Bishop, Pattern Recognition and Machine Learning §1.2, §1.5, §1.6 e §4.3, que cobre probabilidade, teoría da decisión, teoría da información e clasificación lineal na orde que segue este capítulo; Murphy, Probabilistic Machine Learning: An Introduction, capítulos 6 e 10; Prince, Understanding Deep Learning §5.4–5.7; e Saito e Rehmsmeier, The Precision-Recall Plot Is More Informative than the ROC Plot When Evaluating Binary Classifiers on Imbalanced Datasets (PLOS ONE, 2015) — por que a AUC citada arriba non debería ser o único número independente do limiar que mires cando o 1.7 % das pezas son defectuosas.

  1. Ma, T. e Ng, A. CS229 Lecture Notes, Stanford University, capítulos 2 e 3. Onde a cancelación que produce pyp - y deixa de parecer sorte: escolle a distribución da familia exponencial que coincide coa túa saída, usa o seu enlace canónico, e o gradiente sempre é predición menos verdade.

  2. Olah, C. Visual Information Theory (2015), colah.github.io/posts/2015-09-Visual-Information. A explicación máis clara dispoñible de entropía, entropía cruzada e diverxencia KL como custos en bits máis que como fórmulas.

  3. Abu-Mostafa, Y. S., Magdon-Ismail, M. e Lin, H.-T. Learning From Data (AMLBook, 2012), conferencias 13 e 17 do curso de Caltech. A conferencia 13 é validación; a conferencia 17, sobre os tres principios da aprendizaxe, é onde se nomea o data snooping. Entre ambas son a fonte da disciplina deste capítulo: cada ollada a un conxunto de datos é unha decisión de axuste, executases ou non un optimizador.

  4. James, G., Witten, D., Hastie, T. e Tibshirani, R. An Introduction to Statistical Learning, 2ª edición (Springer, 2021), capítulos 2 e 5, pola descomposición nesgo–varianza e pola remostraxe. O volume compañeiro é onde a trampa da selección se expresa sen rodeos: Hastie, Tibshirani e Friedman, The Elements of Statistical Learning, 2ª edición, §7.10.2, The Wrong and Right Way to Do Cross-validation.

  5. Kaufman, S., Rosset, S., Perlich, C. e Stitelman, O. Leakage in Data Mining: Formulation, Detection, and Avoidance. ACM Transactions on Knowledge Discovery from Data 6(4), 2012. Un tratamento formal do fallo demostrado arriba, con estudos de caso de competicións gañadas por un modelo que aprendera un artefacto de como se montaran os datos.

  6. Wilson, E. B. Probable Inference, the Law of Succession, and Statistical Inference. Journal of the American Statistical Association 22(158), pp. 209–212 (1927). O intervalo de score usado en wilson() arriba, aínda o valor por defecto correcto para unha proporción. O intervalo de manual p^±zp^(1p^)/n\hat{p} \pm z\sqrt{\hat{p}(1-\hat{p})/n} é o que hai que evitar: dá disparates preto de 0 e 1, e cobre por baixo gravemente con nn pequeno.

  7. Efron, B. Bootstrap Methods: Another Look at the Jackknife. The Annals of Statistics 7(1), pp. 1–26 (1979). A idea que che permite poñer un intervalo sobre calquera estatística que poidas calcular, incluídas as que non teñen teoría de mostraxe.

Listo para deixar que LIA escolla por ti?

Crea con todos os modelos de IA nun só sitio: empeza gratis hoxe mesmo.