Classificazione, cross-entropy e come non ingannarti da solo
Costruisci un classificatore logistico e scopri perché il 98% di accuracy può indicare un modello che non trova nulla.
In questa pagina
Un modello che risponde questo pezzo va bene per ogni pezzo che esce dal nastro ha ragione il 98,15% delle volte. Ed è anche inutile: dei 74 pezzi difettosi nel test set, non ne intercetta nessuno.
Entrambe le frasi descrivono lo stesso modello. La distanza tra loro è questo capitolo.
La prima metà costruisce il classificatore. Serve quasi nulla di nuovo: il Capitolo 2 ha dato la ricetta per trasformare un’ipotesi su come vengono prodotti i dati in una funzione di loss, e il Capitolo 3 ha dato il meccanismo per scendere lungo qualunque loss quella ricetta produca. Applica entrambe a una domanda sì/no e ne esce la regressione logistica, più una nuova idea — un logit — che tornerà a farsi pagare nel Capitolo 17.
La seconda metà è quella più difficile. Da questo punto in poi, nel corso, tutto viene giudicato da un numero che qualcuno ha misurato, e se non sai distinguere un miglioramento reale da un artefatto di misura, ogni capitolo successivo è decorazione. Quindi: matrice di confusione, precision e recall, i tre split, leakage, e la domanda a cui quasi nessuno risponde onestamente — di quanti esempi di test ho davvero bisogno?
L’aritmetica qui passa su 20.000 righe, quindi è tutta vettorializzata — NumPy fa il lavoro dal Capitolo 2, e da qui in poi non vale più la pena sottolinearlo.
Il nastro, con una domanda più rara
Link alla sezione: Il nastro, con una domanda più raraStessa fabbrica del Capitolo 1, domanda più difficile. Invece di accettare o scartare, la domanda è questo pezzo è difettoso — e i difetti sono rari, cosa che rende difficile la metà di misurazione di questo capitolo e ingannevolmente facile la metà di modellazione.
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 74Tre split, non due. Il motivo merita una sezione a parte e ne riceve una più sotto; per ora, addestra sul primo, regola sul secondo e non guardare il terzo.
Le feature sono standardizzate — sottratta la media, divise per la deviazione standard — usando solo le statistiche di training, per il motivo che il Capitolo 1 ha mostrato con il limite di convergenza del perceptron: dati non centrati rendono ostile la geometria. Da quali righe sia consentito calcolare quella media diventa una domanda viva più avanti in questo capitolo.
Da un verdetto a una probabilità
Link alla sezione: Da un verdetto a una probabilitàIl perceptron restituiva un segno. Un segno non può distinguere scarta da scarta, ma per un soffio, e quella differenza è esattamente ciò che serve a una fabbrica per decidere quali pezzi far ricontrollare prima da un essere umano.
Quindi segui alla lettera la ricetta del Capitolo 2. Scrivi cosa affermi su come viene prodotta un’etichetta, prendi la likelihood, prendi il log, cambiane il segno, e hai una loss. Per un risultato sì/no l’affermazione è una distribuzione di Bernoulli: c’è una probabilità che il pezzo sia difettoso, e
che è solo un modo compatto per scrivere « se , e se ». Prendi il log di questo e cambiane il segno, e la loss per un esempio è
Questa è binary cross-entropy. Non è stata scelta perché comoda; è la negative log-likelihood dell’unica distribuzione che un lancio di moneta può avere. Non c’era altro disponibile.
Manca ancora da dove venga . Il modello calcola una somma pesata , che è un numero reale e copre tutta la retta, mentre una probabilità deve vivere in . La funzione che passa dall’una all’altra è la logistic sigmoid:
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.9820Leggi la colonna di destra come un listino prezzi. Avere ragione con il 90% di fiducia costa 0,105. Rifiutarsi di prendere posizione costa 0,693 — cioè , il prezzo di un’alzata di spalle. Sbagliare con sicurezza costa 4,6, quarantaquattro volte di più, e il prezzo cresce senza limite man mano che il modello diventa più certo del proprio errore. La cross-entropy non conta semplicemente gli errori: fa pagare l’arroganza.
Il gradiente è predizione meno verità
Link alla sezione: Il gradiente è predizione meno veritàIl Capitolo 3 diceva: per addestrare qualunque cosa, ottieni la derivata della loss rispetto a ciascun parametro. Fallo per un esempio. Con e :
Mostra dettagli
Le due righe che fanno cancellare il caos. La sigmoid ha una derivata insolitamente piacevole, . E la loss si differenzia in
Moltiplica le due per la chain rule e compare una volta sopra e una volta sotto. Si cancella esattamente, e ciò che resta è . Quella cancellazione non è una coincidenza — è ciò che succede ogni volta che la loss è la negative log-likelihood di una distribuzione e la funzione di output è quella che quella distribuzione usa naturalmente. Questa coppia ha un nome — un modello lineare generalizzato — e il gradiente pulito è la sua impronta digitale.1
Quindi l’update è predizione meno verità, per l’input. Nient’altro. Ecco l’intero trainer, che è la discesa del Capitolo 3 con una riga cambiata:
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, bIl np.where in sigmoid non è cosmetico. Calcolare direttamente va in overflow per molto negativi; il ramo sceglie la forma algebricamente identica che mantiene negativo l’esponente. Questa è la scatola del floating-point del Capitolo 2 che incassa il suo primo debito, e ne incasserà uno più grande tra due sezioni.
Perché non l’errore quadratico, e perché la risposta riguarda il gradiente
Link alla sezione: Perché non l’errore quadratico, e perché la risposta riguarda il gradienteLa spiegazione standard per preferire la cross-entropy all’errore quadratico è l’argomento della likelihood sopra: l’errore quadratico è ciò che ottieni assumendo rumore gaussiano, le etichette non sono gaussiane, quindi non farlo. È corretto e non convince nessuno, perché puoi scrivere sopra una sigmoid e si addestrerà.
L’argomento che arriva a destinazione riguarda il gradiente. Metti l’errore quadratico sopra una sigmoid e la chain rule dà
Quel in più è quello che prima si cancellava. Ora non lo fa, e va a zero ogni volta che il modello è sicuro — anche quando il modello è sicuro e sbaglia. Valuta entrambi per alcuni punteggi, per un esempio la cui etichetta vera è 1:
| punteggio | cross-entropy | errore quadratico | rapporto | |
|---|---|---|---|---|
| 0,000335 | 1.491 | |||
| 0,017986 | 28,3 | |||
| 0,119203 | 4,8 | |||
| 0,500000 | 2,0 | |||
| 0,880797 | 4,8 |
A il modello è sbagliato quanto è possibile esserlo, e l’errore quadratico risponde con un gradiente 1.491 volte più piccolo di quello della cross-entropy. Più grave è l’errore, meno il modello impara da esso. Il gradiente della cross-entropy, invece, satura a : massimamente sbagliato produce un segnale massimamente grande, e non più grande.
Fai partire la gara. Duemila punti bilanciati, pesi iniziali identici scelti per essere sicuri e sbagliati (), learning rate identico, cambia solo la loss. Entrambe le esecuzioni sono valutate con cross-entropy così le colonne sono confrontabili.
| epoch | cross-entropy loss | accuracy | squared-error loss | 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 cross-entropy ha finito entro l’epoch 50. L’errore quadratico è ancora al 24% di accuracy all’epoch 100 — e non si era mosso dal 23% all’epoch 10 — peggio che tirare a indovinare, perché era partito sicuro e sbagliato e il gradiente che avrebbe dovuto salvarlo è stato moltiplicato per 0,0007. Scappa intorno all’epoch 500 e atterra nello stesso punto. Quindi il riassunto onesto è che l’errore quadratico sopra una sigmoid non è scorretto; è lento esattamente dove la velocità conta di più. Su un modello a due parametri perdi 450 epoch. Su una rete con cento layer, dove qualche unità da qualche parte è sempre sicura e sbagliata, perdi l’intera esecuzione di training.
Entropia, cross-entropy e KL, in una pagina
Link alla sezione: Entropia, cross-entropy e KL, in una paginaTre quantità, necessarie nel modo giusto nel Capitolo 8 per la perplexity e nel Capitolo 11 per la penalità che mantiene una policy fine-tuned vicina al suo riferimento. Sono più facili della loro reputazione.2
Entropia è il numero medio di bit che devi spendere per comunicare un’estrazione da una distribuzione, se usi il miglior codice possibile per essa:
Cross-entropy è ciò che spendi quando usi un codice costruito per su dati che in realtà arrivano da :
Divergenza KL è l’eccesso — lo spreco, in bit, causato dal credere a quando la verità è :
Controlla tutte e tre sul nastro:
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 bitsLì si vedono due cose. Primo, un modello che si limita a riportare il tasso base di training, 1,69%, ottiene una cross-entropy di 0,1330 bit, quasi esattamente l’entropia delle etichette di test — come deve essere, dato che ha la distribuzione giusta e nessun’altra informazione. L’entropia è il pavimento che ti compra l’ignoranza sul singolo caso. Secondo, un modello che alza le spalle e dice 0,5 paga esattamente 1 bit, e il divario tra i due, 0,8671 bit, è precisamente la divergenza KL. non è un’identità da memorizzare; è una fattura che puoi vedere sommarsi.
E il collegamento con il training: quando l’etichetta è una singola classe nota, la distribuzione “vera” è one-hot, la sua entropia è zero, e la cross-entropy è uguale alla divergenza KL. Minimizzare la cross-entropy e tirare la distribuzione del modello verso la verità sono lo stesso atto.
Più di due risposte: softmax, e lo shift che non costa nulla
Link alla sezione: Più di due risposte: softmax, e lo shift che non costa nullaDifettoso non è una sola cosa. Nello stampaggio, un pezzo può uscire come short shot (materiale insufficiente), flash (troppo materiale, schiacciato fuori dallo stampo), o burn. Quattro risultati, quindi quattro logits, e devono diventare quattro probabilità che sommano a uno. Questa è la softmax:
Ha una proprietà che sembra un incidente e in realtà è l’intera implementazione:
per qualunque costante , perché e si cancellano sopra e sotto. Solo le differenze tra logits significano qualcosa. Il livello assoluto non è informazione.
Per fortuna, perché il livello assoluto è ciò che rompe il computer:
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 manda in overflow un float a 64 bit, la somma diventa infinito, e infinito diviso infinito è nan — non un errore, non un crash, solo un buco silenzioso dove prima c’erano tre probabilità. Sottrarre il logit massimo non cambia nulla matematicamente e cambia tutto numericamente, perché l’esponente più grande diventa esattamente . Questo è il trucco logsumexp del Capitolo 2 con la tuta da lavoro, e ogni implementazione seria lo fa:
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, bIl gradiente è di nuovo predizione meno verità, ora con one-hot. Il caso binario era da sempre un caso speciale.
Addestrato su 3.000 pezzi e testato su 1.000, con tre misure ciascuno (larghezza, peso, temperatura di fusione), raggiunge 94,00% di accuracy. Ecco cosa quel numero sta nascondendo:
| verità ↓ / predetto → | ok | short shot | flash | burn | recall |
|---|---|---|---|---|---|
| ok | 850 | 5 | 9 | 0 | 0,984 |
| short shot | 22 | 21 | 0 | 0 | 0,488 |
| flash | 20 | 0 | 30 | 1 | 0,588 |
| burn | 3 | 0 | 0 | 39 | 0,929 |
| precision | 0,950 | 0,808 | 0,769 | 0,975 |
Il modello trova meno della metà degli short shot. L’accuracy non può vederlo, perché l’86% dei pezzi va bene e azzeccare quelli basta a sostenere la media. Macro F1 — la media degli F1 per classe, che pesa una classe rara quanto una comune — è 0,7983, contro un micro F1 di 0,9400 che per definizione è identico all’accuracy. Ogni volta che qualcuno riporta un solo numero F1, chiedi quale.
Questa è l’ultima parte della modellazione. Il resto del capitolo riguarda i numeri.
Tre modelli, una accuracy
Link alla sezione: Tre modelli, una accuracyPrendi il modello binario addestrato e crea due varianti moltiplicando ogni logit per una costante: 0,35 per una versione esitante, 4 per una troppo sicura. Moltiplicare per un numero positivo non può cambiare alcun segno, quindi tutti e tre i modelli predicono esattamente la stessa etichetta per tutti i 4.000 pezzi di test. L’accuracy non riesce a distinguerli. La cross-entropy non ha alcuna difficoltà:
| modello | accuracy | cross-entropy | loss media quando ha ragione | loss media quando sbaglia | peggiore singola loss |
|---|---|---|---|---|---|
| esitante (logits × 0,35) | 0,9830 | 0,1549 | 0,1369 | 1,1990 | 2,80 |
| come addestrato | 0,9830 | 0,0564 | 0,0147 | 2,4689 | 7,82 |
| troppo sicuro (logits × 4) | 0,9830 | 0,1563 | 0,0009 | 9,1427 | 27,63 |
Il modello esitante paga una piccola tassa su ogni pezzo, compresi le migliaia che azzecca. Quello troppo sicuro è quasi gratis quando ha ragione e catastrofico quando sbaglia — un pezzo in quel test set gli costa da solo 27,63 nats. I due arrivano a quasi lo stesso totale per strade opposte, e il modello addestrato, le cui probabilità sono calibrate sui dati, si colloca tre volte sotto entrambi.
Questo è il modo più netto per dichiarare la differenza tra una loss e una metric. La loss è ciò che ottimizzi: deve essere differenziabile, e vede tutto ciò che il modello ha detto, incluso quanto ne era sicuro. La metric è ciò su cui vieni giudicato: può essere una funzione a gradino, una regola di business, un conteggio dei difetti mancati. Non sono lo stesso oggetto e non sempre concordano — per questo definisci entrambe prima di iniziare, e non lasci mai che la loss sostituisca la metric solo perché capita di essere sullo schermo.
La baseline stupida viene prima
Link alla sezione: La baseline stupida viene primaPrima di qualsiasi modello, il requisito: che punteggio ottiene la risposta più pigra possibile? Su questo nastro, dire sempre che va bene:
always-say-fine baseline: accuracy = 0.9815
confusion (tn, fp, fn, tp) = (3926, 0, 74, 0)98,15%. Ora il modello logistico addestrato, alla soglia predefinita di 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 battuto la baseline di 0,15 punti percentuali, e qualunque report che si fermi all’accuracy lo chiamerà una vittoria. La matrice di confusione dice cosa è successo davvero:
| predetto buono | predetto difettoso | |
|---|---|---|
| davvero buono | 3.924 | 2 |
| davvero difettoso | 66 | 8 |
Ha trovato 8 pezzi difettosi su 74 e ne ha lasciati passare 66. Tre numeri danno un nome ai tre modi di leggere quella tabella:
- Precision . Dei pezzi che ha segnalato, quanti erano davvero difettosi. Questo è il costo delle ispezioni sprecate.
- Recall . Dei pezzi difettosi, quanti ne ha intercettati. Questo è il costo di spedire un pezzo difettoso a un cliente.
- F1 , la loro media armonica, che resta vicina al più piccolo dei due e quindi rifiuta di farsi lusingare da uno solo.
Quale conta dipende dalla fabbrica, non dalla matematica: un’ispezione costa pochi secondi e un difetto spedito costa un richiamo, quindi qui domina la recall e 0,108 è un fallimento.
Ma il problema non è il modello. È la soglia, e la soglia non fa parte del modello — è una decisione di business applicata dopo a una probabilità. Scorrila:
| soglia | 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 |
Leggi la colonna dell’accuracy verso il basso. Scende per tutto il percorso — dal 98,30% al 65,93% — mentre il modello passa dal catturare 8 difetti al catturarne 71 su 74. Ogni cosa utile che questo modello può fare peggiora la sua accuracy. Un team che ottimizza il numero da titolo spedirebbe la versione che non trova nulla.
Mostra dettagli
Il class weighting non crea segnale, sposta l’operating point. Il primo riflesso abituale con classi sbilanciate è pesare la classe rara nella loss. Facendolo, con pesi di 1, 10 e 60 sui positivi:
| peso sui positivi | 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 e recall si spostano molto. L’AUC — la probabilità che il modello classifichi un pezzo difettoso casuale sopra uno buono casuale, ignorando del tutto la soglia — si muove di 0,0002, cioè nulla. Il reweighting ha fatto scorrere lo stesso modello lungo la stessa curva di compromesso. Spesso è ciò che vuoi, e non è mai nuova informazione: se il ranking è cattivo, nessuno schema di pesatura lo salverà.
Tre split, e il leak che stai per trovare
Link alla sezione: Tre split, e il leak che stai per trovarePerché tre split e non due? Perché nel momento in cui usi un insieme di esempi per scegliere qualcosa — una soglia, un learning rate, quale dei sei modelli spedire — quell’insieme è stato usato per il fitting, e il suo punteggio smette di essere non distorto.3 Misurato su questo nastro: scorrere la soglia sul validation set sceglie 0,196, e il modello poi ottiene F1 = 0,4122 sul test set intatto. Se lo sweep fosse stato eseguito direttamente sul test set, il miglior risultato ottenibile lì sarebbe stato 0,4186 — un numero che nessuno ha il diritto di riportare.
Qui il divario è piccolo, 0,006, perché è un solo hyperparameter scansionato una volta su 4.000 esempi di validation. Cresce con ogni decisione in più e con ogni riduzione del validation set. Nota anche che la direzione non è garantita in una singola esecuzione: la soglia scelta ha ottenuto 0,3902 su validation e 0,4122 su test, quindi questa volta la validation l’ha sottostimata. Il bias è sistematico su molte decisioni, non visibile in una sola.4
Ora l’esercizio. Il log del nastro arriva con una terza colonna, station_seconds: quanto tempo ogni pezzo ha passato alla stazione di ispezione. Aggiungerla è una modifica di una riga al preprocessing. Ecco cosa fa:
| modello | accuracy | precision | recall | F1 | cross-entropy | AUC |
|---|---|---|---|---|---|---|
| larghezza + 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 |
La recall passa dal 10,8% al 77,0%. L’F1 più che quadruplica. E nota cosa ha fatto l’accuracy: 98,30% → 99,20%, un guadagno di nove decimi di punto, il tipo di numero che in una slide riassuntiva viene arrotondato a “circa 99% in entrambi i casi”. L’accuracy prima non è riuscita a vedere il fallimento e ora non riesce a vedere la frode.
Prima di continuare: il modello sta barando. Scopri come.
Come dare la caccia a un leak, nell’ordine che lo trova più in fretta.
-
Confronta train e test. L’overfitting si manifesta come un grande divario. Qui: modello onesto 0,9838 train / 0,9830 test; modello con leak 0,9936 train / 0,9920 test. Entrambi i divari sono sotto 0,2 punti. Un leak non assomiglia all’overfitting — la feature con leak è altrettanto disponibile al test time, quindi il modello generalizza magnificamente a un mondo che non esiste.
-
Addestra un modello per feature, da sola. Qualunque cosa porti la risposta si annuncerà:
feature da sola accuracy recall F1 AUC larghezza 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 colonna, da sola, ordina i difetti con AUC 0,9960. Due misure prese con calibro e bilancia arrivano a 0,87 e 0,79. Questa asimmetria è l’allarme.
-
Chiedi quando è stato scritto ogni numero. Tempo medio di permanenza: 2,23 secondi per i pezzi passati, 15,56 secondi per quelli falliti. Certo che sì. Un pezzo resta alla stazione perché un ispettore lo ha tolto dal nastro — cosa che succede dopo, e solo perché, qualcuno ha deciso che era difettoso. La colonna non è una misura del pezzo. È una misura del verdetto.
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 riga evidenziata è il leak: il tempo di permanenza di un pezzo difettoso viene estratto da una distribuzione diversa, perché un essere umano lo ha tolto dal nastro. Questo è il bug serio più comune nel machine learning applicato, e ha un nome: target leakage — informazione nelle feature di training che non sarebbe disponibile nel momento in cui la predizione deve essere fatta.5 Non lancia eccezioni. Produce un numero migliore. Ogni incentivo in un progetto spinge a tenerlo.
La difesa è una domanda, posta a ogni colonna: nell’istante in cui mi serve questa predizione, questo valore esiste già? Su un nastro live, station_seconds è sconosciuto fino a dopo che il pezzo è stato ispezionato — che è la cosa che il modello avrebbe dovuto sostituire.
Di quanti esempi di test ho bisogno?
Link alla sezione: Di quanti esempi di test ho bisogno?Supponi di valutare un modello su 20 esempi e che ne azzecchi 17. Riporti 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 lettura onesta di 17/20 è da qualche parte tra 64% e 95%. Un modello realmente al 65% produce questo risultato il 4,4% delle volte — una volta su ventitré — e se hai provato una manciata di prompts e riportato il migliore, quella run l’hai fabbricata tu. Diciassette su venti non distinguono un modello all’85% da uno al 65%.
Due modi per mettere un intervallo su un tasso, ed entrambi devono stare nel tuo toolkit:
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 per un semplice tasso di successo; si comporta bene a qualunque e non richiede casualità. Nota sopra che a l’estremo superiore del bootstrap è 1,0000 — ricampionare 20 punti può facilmente estrarne 20 corretti, quindi non può rappresentare un intervallo più stretto della propria granularità. Usa il bootstrap7 dove non esiste una formula, cioè nella maggior parte dei casi interessanti: F1, macro-medie, BLEU, pass@1, il punteggio di un giudice basato su rubric. Su questo nastro, l’F1 del modello regolato, 0,4122, porta un intervallo bootstrap di [0,3009, 0,5156] — ed è questo il numero che dovrebbe apparire nel report, perché la stima puntuale da sola invita a un confronto che non può sostenere.
Un’altra misura, perché cambia il modo in cui dovresti confrontare due modelli. Due modelli valutati sugli stessi 500 esempi:
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)I loro intervalli si sovrappongono, e la regola popolare — barre d’errore sovrapposte significano nessuna differenza significativa — dichiarerebbe il confronto inconcludente. Non lo è. I due modelli hanno girato sugli stessi esempi, quindi la quantità giusta è la differenza per esempio, il cui intervallo è [0.0260, 0.0680], comodamente sopra zero. Non concordano solo su 31 item su 500, e A vince 27 di quei disaccordi; gli esempi condivisi, facili e difficili allo stesso modo, si cancellano invece di aggiungere rumore. Confronta i modelli in modo paired, e arrivi alla stessa conclusione con una frazione dei dati.
Dove si va ora
Link alla sezione: Dove si va oraOra hai un modello che emette probabilità calibrate, una loss derivata da un’affermazione sui dati invece che scelta per comodità, un gradiente che è letteralmente predizione meno verità, e — ancora più importante — gli strumenti per scoprire se qualcosa di tutto questo funziona. L’intervallo Wilson di dieci righe qui sopra viene riutilizzato identico: accompagna le varianti di prompt nel Capitolo 15, le tabelle di retrieval nel Capitolo 19, e il golden set nel Capitolo 29. Il bootstrap è ciò a cui ricorri quando non esiste una formula.
Ma il modello è ancora un solo layer. Disegna una linea, e il Capitolo 1 ha dimostrato con quattro righe di XOR che una linea non basta. La correzione è impilare: un primo layer che piega lo spazio, un secondo che disegna la linea nello spazio piegato.
È qui che il gradiente ordinato di questo capitolo si esaurisce. Tutto sopra ha funzionato perché poteva essere scritto a mano, una volta, per un modello con un solo layer tra input e loss. Metti un secondo layer nel mezzo e la domanda cambia forma: qual è la derivata della loss rispetto a un peso che non tocca affatto l’output — uno la cui influenza arriva solo attraverso un altro layer, magari lungo diversi percorsi contemporaneamente?
Quella derivata esiste. Calcolarla a mano è senza speranza per qualunque cosa più grande di un giocattolo, e calcolarla un parametro alla volta è senza speranza su una scala diversa. Serve una procedura che ottenga ogni derivata nella rete da un singolo passaggio all’indietro sullo stesso grafo che il forward pass ha appena percorso.
Questo è il Capitolo 5, ed è il motore su cui gira il resto del corso.
Fonti e metodo
Link alla sezione: Fonti e metodoVale la pena leggere insieme a questo capitolo anche: Bishop, Pattern Recognition and Machine Learning §1.2, §1.5, §1.6 e §4.3, che copre probabilità, teoria delle decisioni, teoria dell’informazione e classificazione lineare nell’ordine seguito da questo capitolo; Murphy, Probabilistic Machine Learning: An Introduction, capitoli 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) — perché l’AUC citata sopra non dovrebbe essere l’unico numero indipendente dalla soglia che guardi quando l’1,7% dei pezzi è difettoso.
Riferimenti
Link alla sezione: Riferimenti-
Ma, T. e Ng, A. CS229 Lecture Notes, Stanford University, capitoli 2 e 3. Dove la cancellazione che produce smette di sembrare fortuna: scegli la distribuzione della famiglia esponenziale che corrisponde al tuo output, usa il suo canonical link, e il gradiente è sempre predizione meno verità. ↩
-
Olah, C. Visual Information Theory (2015),
colah.github.io/posts/2015-09-Visual-Information. La spiegazione più chiara disponibile di entropia, cross-entropy e divergenza KL come costi in bit invece che come formule. ↩ -
Abu-Mostafa, Y. S., Magdon-Ismail, M. e Lin, H.-T. Learning From Data (AMLBook, 2012), lezioni 13 e 17 del corso Caltech. La lezione 13 è sulla validation; la lezione 17, sui tre principi dell’apprendimento, è dove viene nominato il data snooping. Insieme sono la fonte della disciplina di questo capitolo: ogni sguardo a un data set è una decisione di fitting, che tu abbia eseguito un ottimizzatore o no. ↩
-
James, G., Witten, D., Hastie, T. e Tibshirani, R. An Introduction to Statistical Learning, 2ª edizione (Springer, 2021), capitoli 2 e 5, per la decomposizione bias–variance e per il resampling. Il volume companion è dove la trappola della selezione viene dichiarata esplicitamente: Hastie, Tibshirani e Friedman, The Elements of Statistical Learning, 2ª edizione, §7.10.2, The Wrong and Right Way to Do Cross-validation. ↩
-
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. Una trattazione formale del fallimento dimostrato sopra, con casi studio da competizioni vinte da un modello che aveva imparato un artefatto di come erano stati assemblati i dati. ↩
-
Wilson, E. B. Probable Inference, the Law of Succession, and Statistical Inference. Journal of the American Statistical Association 22(158), pp. 209–212 (1927). L’intervallo score usato in
wilson()sopra, ancora il default giusto per una proporzione. L’intervallo da manuale è quello da evitare: dà assurdità vicino a 0 e 1, e sotto-copre pesantemente a piccoli . ↩ -
Efron, B. Bootstrap Methods: Another Look at the Jackknife. The Annals of Statistics 7(1), pp. 1–26 (1979). L’idea che ti consente di mettere un intervallo su qualunque statistica tu possa calcolare, incluse quelle senza teoria del campionamento. ↩