Clasificare, cross-entropy și cum să nu te păcălești singur
Construiește un clasificator logistic și vezi de ce 98 % acuratețe poate însemna un model care nu găsește nimic.
Pe această pagină
Un model care răspunde piesa aceasta este în regulă pentru fiecare piesă care iese de pe bandă are dreptate în 98,15 % din cazuri. Este și inutil: dintre cele 74 de piese defecte din setul de test, nu prinde niciuna.
Ambele propoziții descriu același model. Distanța dintre ele este acest capitol.
Prima jumătate construiește clasificatorul. Nu are nevoie de aproape nimic nou: Capitolul 2 a dat rețeta pentru a transforma o presupunere despre cum sunt produse datele într-o funcție de pierdere, iar Capitolul 3 a dat mecanismul pentru a coborî pe orice pierdere îți oferă rețeta. Aplică-le pe amândouă unei întrebări cu da/nu și obții regresia logistică, plus o idee nouă — un logit — care va fi taxată din nou în Capitolul 17.
A doua jumătate este cea mai grea. Tot ce urmează după acest punct în curs este judecat după un număr pe care cineva l-a măsurat, iar dacă nu poți deosebi o îmbunătățire reală de un artefact de măsurare, fiecare capitol care urmează este decor. Așadar: matricea de confuzie, precizia și recall, cele trei împărțiri, leakage și întrebarea la care aproape nimeni nu răspunde sincer — de câte exemple de test am nevoie, de fapt?
Aritmetica de aici rulează peste 20.000 de rânduri, deci este vectorizată peste tot — NumPy face treaba încă din Capitolul 2, iar de aici încolo nu mai merită menționat.
Banda, cu o întrebare mai rară
Link către secțiunea: Banda, cu o întrebare mai rarăAceeași fabrică precum în Capitolul 1, o întrebare mai grea. În loc de acceptă sau respinge, întrebarea este este această piesă defectă — iar defectele sunt rare, ceea ce face jumătatea de măsurare a acestui capitol grea și jumătatea de modelare înșelător de ușoară.
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 74Trei împărțiri, nu două. Motivul merită o secțiune proprie și primește una mai jos; pentru moment, antrenează pe prima, ajustează pe a doua și nu te uita la a treia.
Feature-urile sunt standardizate — se scade media, se împarte la deviația standard — folosind doar statisticile de antrenare, din motivul pe care Capitolul 1 l-a demonstrat cu limita de convergență a perceptronului: datele necentrate fac geometria ostilă. Din ce rânduri ai voie să calculezi acea medie devine o întrebare vie mai târziu în acest capitol.
De la verdict la probabilitate
Link către secțiunea: De la verdict la probabilitatePerceptronul întorcea un semn. Un semn nu poate distinge respinge de respinge, dar la limită, iar diferența aceasta este exact ce îi trebuie unei fabrici ca să decidă ce piese ar trebui re-inspectate mai întâi de un om.
Așa că urmează literalmente rețeta din Capitolul 2. Scrie ce susții despre cum este produsă o etichetă, ia likelihood, ia logaritmul, neagă-l și ai o pierdere. Pentru un rezultat da/nu, afirmația este o distribuție Bernoulli: există o probabilitate ca piesa să fie defectă, iar
ceea ce este doar un mod compact de a scrie „ dacă și dacă ”. Ia logaritmul și neagă-l, iar pierderea pentru un exemplu este
Aceasta este binary cross-entropy. Nu a fost aleasă pentru că este convenabilă; este negative log-likelihood a singurei distribuții pe care o poate avea aruncarea unei monede. Nu era disponibil nimic altceva.
Încă lipsește de unde vine . Modelul calculează o sumă ponderată , care este un număr real și se întinde pe toată dreapta reală, iar o probabilitate trebuie să trăiască în . Funcția care mută între ele este 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.9820Citește coloana din dreapta ca pe o listă de prețuri. Să ai dreptate cu 90 % încredere costă 0,105. Refuzul de a te angaja costă 0,693 — adică , prețul unei ridicări din umeri. Să greșești cu încredere costă 4,6, de patruzeci și patru de ori mai mult, iar prețul crește fără limită pe măsură ce modelul devine mai sigur de o greșeală. Cross-entropy nu doar numără erorile: taxează aroganța.
Gradientul este predicția minus adevărul
Link către secțiunea: Gradientul este predicția minus adevărulCapitolul 3 spunea: ca să antrenezi orice, obține derivata pierderii în raport cu fiecare parametru. Fă asta pentru un exemplu. Cu și :
Afișează detaliile
Cele două linii care fac dezordinea să se anuleze. Sigmoidul are o derivată neobișnuit de plăcută, . Iar pierderea se diferențiază în
Înmulțește-le pe cele două prin regula lanțului și apare o dată sus și o dată jos. Se anulează exact, iar este ceea ce rămâne. Acea anulare nu este o coincidență — este ce se întâmplă ori de câte ori pierderea este negative log-likelihood a unei distribuții, iar funcția de ieșire este cea pe care acea distribuție o folosește în mod natural. Această pereche are un nume — un model liniar generalizat — iar gradientul curat este amprenta lui.1
Deci actualizarea este predicția minus adevărul, înmulțită cu inputul. Nimic altceva. Iată întregul antrenor, care este descent-ul din Capitolul 3 cu o singură linie schimbată:
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, bnp.where din sigmoid nu este cosmetic. Calcularea directă a lui face overflow pentru negativ mare; ramura alege forma algebric identică ce păstrează exponentul negativ. Aceasta este cutia de floating-point din Capitolul 2 care își încasează prima datorie, iar peste două secțiuni va încasa una mai mare.
De ce nu eroarea pătratică și de ce răspunsul este despre gradient
Link către secțiunea: De ce nu eroarea pătratică și de ce răspunsul este despre gradientExplicația standard pentru preferarea cross-entropy în locul erorii pătratice este argumentul de likelihood de mai sus: eroarea pătratică este ce obții presupunând zgomot gaussian, etichetele nu sunt gaussiene, deci nu face asta. Este corectă și nu convinge pe nimeni, pentru că poți scrie peste un sigmoid și se va antrena.
Argumentul care prinde este despre gradient. Pune eroare pătratică peste un sigmoid și regula lanțului dă
Acel suplimentar este cel care s-a anulat înainte. Acum nu se anulează și merge la zero ori de câte ori modelul este încrezător — inclusiv când modelul este încrezător și greșește. Evaluează-le pe amândouă la câteva scoruri, pentru un exemplu a cărui etichetă reală este 1:
| scor | cross-entropy | eroare pătratică | raport | |
|---|---|---|---|---|
| 0,000335 | 1.491 | |||
| 0,017986 | 28,3 | |||
| 0,119203 | 4,8 | |||
| 0,500000 | 2,0 | |||
| 0,880797 | 4,8 |
La modelul greșește cât de mult se poate, iar eroarea pătratică răspunde cu un gradient de 1.491 de ori mai mic decât cel al cross-entropy. Cu cât greșeala este mai rea, cu atât modelul învață mai puțin din ea. Gradientul cross-entropy, între timp, se saturează la : maxim greșit produce un semnal maxim de mare, și nu mai mare.
Rulează cursa. Două mii de puncte echilibrate, greutăți inițiale identice alese să fie încrezător greșite (), learning rate identic, doar pierderea diferă. Ambele rulări sunt evaluate cu cross-entropy, ca să fie comparabile coloanele.
| epocă | pierdere cross-entropy | acuratețe | pierdere eroare pătratică | acuratețe |
|---|---|---|---|---|
| 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 |
Cross-entropy termină până la epoca 50. Eroarea pătratică este încă la 24 % acuratețe la epoca 100 — și nu se mișcase de la 23 % la epoca 10 — mai rău decât ghicitul, pentru că a pornit încrezător greșit, iar gradientul care ar fi salvat-o a fost înmulțit cu 0,0007. Scapă pe la epoca 500 și ajunge în același loc. Deci rezumatul sincer este că eroarea pătratică peste un sigmoid nu este incorectă; este lentă exact acolo unde viteza contează cel mai mult. Pe un model cu doi parametri pierzi 450 de epoci. Pe o rețea cu o sută de straturi, unde o unitate de undeva este mereu încrezător greșită, pierzi rularea de antrenare.
Entropie, cross-entropy și KL, într-o pagină
Link către secțiunea: Entropie, cross-entropy și KL, într-o paginăTrei cantități, necesare cum trebuie în Capitolul 8 pentru perplexitate și în Capitolul 11 pentru penalizarea care ține o politică fine-tuned aproape de referința ei. Sunt mai ușoare decât reputația lor.2
Entropia este numărul mediu de biți pe care trebuie să-i cheltuiești ca să comunici o extragere dintr-o distribuție, dacă folosești cel mai bun cod posibil pentru ea:
Cross-entropy este ce cheltuiești când folosești un cod construit pentru pe date care vin de fapt din :
Divergența KL este excesul — risipa, în biți, cauzată de faptul că îl crezi pe când adevărul este :
Verifică-le pe toate trei pe bandă:
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 bitsDouă lucruri se văd acolo. Mai întâi, un model care raportează pur și simplu rata de bază din antrenare, 1,69 %, obține o cross-entropy de 0,1330 biți, aproape exact entropia etichetelor de test — așa cum trebuie, fiindcă are distribuția corectă și nicio altă informație. Entropia este podeaua pe care ți-o cumpără ignoranța despre individ. În al doilea rând, un model care ridică din umeri și spune 0,5 plătește exact 1 bit, iar diferența dintre cele două, 0,8671 biți, este exact divergența KL. nu este o identitate de memorat; este o factură pe care o poți urmări cum se adună.
Iar legătura înapoi la antrenare: când eticheta este o singură clasă cunoscută, distribuția „adevărată” este one-hot, entropia ei este zero, iar cross-entropy egalează divergența KL. A minimiza cross-entropy și a trage distribuția modelului spre adevăr sunt același act.
Mai mult de două răspunsuri: softmax și deplasarea care nu costă nimic
Link către secțiunea: Mai mult de două răspunsuri: softmax și deplasarea care nu costă nimicDefect nu înseamnă un singur lucru. În turnare, o piesă poate ieși ca short shot (material insuficient), flash (prea mult, împins afară din matriță) sau burn. Patru rezultate, deci patru logits, iar ele trebuie să devină patru probabilități care însumează unu. Acesta este softmax:
Are o proprietate care arată ca un accident și este, de fapt, întreaga implementare:
pentru orice constantă , deoarece și se anulează sus și jos. Doar diferențele dintre logits înseamnă ceva. Nivelul absolut nu este informație.
Din fericire, pentru că nivelul absolut este cel care strică computerul:
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 face overflow într-un float pe 64 de biți, suma devine infinit, iar infinit împărțit la infinit este nan — nu o eroare, nu un crash, doar o gaură tăcută acolo unde erau trei probabilități. Scăderea celui mai mare logit nu schimbă nimic matematic și schimbă totul numeric, pentru că cel mai mare exponent devine exact . Este trucul logsumexp din Capitolul 2 în haine de lucru, iar fiecare implementare serioasă îl folosește:
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, bGradientul este din nou predicția minus adevărul, acum cu one-hot. Cazul binar a fost tot timpul un caz special.
Antrenat pe 3.000 de piese și testat pe 1.000, cu trei măsurători fiecare (lățime, greutate, temperatură de topire), ajunge la 94,00 % acuratețe. Iată ce ascunde acel număr:
| adevăr ↓ / prezis → | 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 |
| precizie | 0,950 | 0,808 | 0,769 | 0,975 |
Modelul găsește mai puțin de jumătate dintre short shots. Acuratețea nu poate vedea asta, pentru că 86 % dintre piese sunt bune, iar faptul că le nimerește pe acelea este suficient ca să tragă media. Macro F1 — media scorurilor F1 pe clasă, care cântărește o clasă rară la fel ca una comună — este 0,7983, față de un micro F1 de 0,9400 care, prin definiție, este identic cu acuratețea. Ori de câte ori cineva raportează un singur număr F1, întreabă care.
Asta a fost ultima parte de modelare. Restul capitolului este despre numere.
Trei modele, o singură acuratețe
Link către secțiunea: Trei modele, o singură acuratețeIa modelul binar antrenat și fă două variante înmulțind fiecare logit cu o constantă: 0,35 pentru o versiune ezitantă, 4 pentru una prea încrezătoare. Înmulțirea cu un număr pozitiv nu poate schimba niciun semn, deci toate cele trei modele prezic exact aceeași etichetă pentru toate cele 4.000 de piese de test. Acuratețea nu le poate deosebi. Cross-entropy nu are nicio problemă:
| model | acuratețe | cross-entropy | pierdere medie când are dreptate | pierdere medie când greșește | cea mai mare pierdere individuală |
|---|---|---|---|---|---|
| ezitant (logits × 0,35) | 0,9830 | 0,1549 | 0,1369 | 1,1990 | 2,80 |
| așa cum a fost antrenat | 0,9830 | 0,0564 | 0,0147 | 2,4689 | 7,82 |
| prea încrezător (logits × 4) | 0,9830 | 0,1563 | 0,0009 | 9,1427 | 27,63 |
Modelul ezitant plătește o taxă mică pe fiecare piesă, inclusiv pe miile pe care le nimerește. Cel prea încrezător este aproape gratuit când are dreptate și catastrofal când greșește — o singură piesă din acel set de test îl costă 27,63 nats de una singură. Cele două ajung aproape la același total pe rute opuse, iar modelul antrenat, ale cărui probabilități sunt calibrate pe date, stă de trei ori mai jos decât ambele.
Acesta este cel mai tăios mod de a enunța diferența dintre o pierdere și o metrică. Pierderea este ce optimizezi: trebuie să fie diferențiabilă și vede tot ce a spus modelul, inclusiv cât de sigur era. Metrica este după ce ești judecat: poate fi o funcție treaptă, o regulă de business, un număr de defecte ratate. Nu sunt același obiect și nu sunt mereu de acord — de aceea le definești pe amândouă înainte să începi și nu lași niciodată pierderea să țină loc de metrică doar pentru că se întâmplă să fie pe ecran.
Baseline-ul prost merge primul
Link către secțiunea: Baseline-ul prost merge primulÎnainte de orice model, cerința: ce scor obține cel mai leneș răspuns posibil? Pe această bandă, spune mereu că e în regulă:
always-say-fine baseline: accuracy = 0.9815
confusion (tn, fp, fn, tp) = (3926, 0, 74, 0)98,15 %. Acum modelul logistic antrenat, la pragul implicit 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 %. A bătut baseline-ul cu 0,15 puncte procentuale, iar orice raport care se oprește la acuratețe va numi asta o victorie. Matricea de confuzie spune ce s-a întâmplat de fapt:
| prezis bună | prezis defectă | |
|---|---|---|
| de fapt bună | 3.924 | 2 |
| de fapt defectă | 66 | 8 |
A găsit 8 piese defecte din 74 și a lăsat 66 să treacă. Trei numere numesc cele trei moduri de a citi tabelul:
- Precizie . Dintre piesele pe care le-a marcat, câte erau cu adevărat defecte. Acesta este costul inspecțiilor irosite.
- Recall . Dintre piesele defecte, câte a prins. Acesta este costul de a trimite o piesă proastă unui client.
- F1 , media lor armonică, care rămâne aproape de cea mai mică dintre cele două și, prin urmare, refuză să fie flatată doar de una dintre ele.
Ce contează depinde de fabrică, nu de matematică: o inspecție costă câteva secunde, iar un defect livrat costă o notificare de rechemare, deci aici recall domină și 0,108 este un eșec.
Dar modelul nu este problema. Pragul este, iar pragul nu face parte din model — este o decizie de business aplicată după aceea unei probabilități. Parcurge-l:
| prag | TP | FP | FN | acuratețe | precizie | 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 |
Citește coloana de acuratețe în jos. Cade pe tot parcursul — de la 98,30 % la 65,93 % — în timp ce modelul trece de la a prinde 8 defecte la a prinde 71 din 74. Fiecare lucru util pe care îl poate face acest model îi înrăutățește acuratețea. O echipă care optimizează numărul din titlu ar livra versiunea care nu găsește nimic.
Afișează detaliile
Ponderarea claselor nu creează semnal, mută punctul de operare. Primul reflex obișnuit cu clase dezechilibrate este să ponderezi clasa rară în pierdere. Făcând asta, cu ponderi de 1, 10 și 60 pe pozitive:
| pondere pe pozitive | acuratețe | precizie | 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 |
Precizia și recall se mișcă mult. AUC — probabilitatea ca modelul să claseze o piesă defectă aleatoare deasupra uneia bune aleatoare, ignorând pragul complet — se mișcă cu 0,0002, adică nimic. Reponderarea a glisat același model de-a lungul aceleiași curbe de compromis. Asta este adesea ce vrei și nu este niciodată informație nouă: dacă ranking-ul este prost, nicio schemă de ponderare nu îl va salva.
Trei împărțiri și leakage-ul pe care ești pe cale să-l găsești
Link către secțiunea: Trei împărțiri și leakage-ul pe care ești pe cale să-l găseștiDe ce trei împărțiri și nu două? Pentru că în clipa în care folosești un set de exemple ca să alegi ceva — un prag, un learning rate, care dintre șase modele să fie livrat — acel set a fost folosit pentru fitting, iar scorul lui încetează să fie nepartinitor.3 Măsurat pe această bandă: parcurgerea pragului pe setul de validare alege 0,196, iar modelul obține apoi F1 = 0,4122 pe setul de test neatins. Dacă parcurgerea ar fi fost rulată direct pe setul de test, cel mai bun scor posibil acolo era 0,4186 — un număr pe care nimeni nu are dreptul să-l raporteze.
Diferența este mică aici, 0,006, pentru că este un singur hyperparameter parcurs o dată pe 4.000 de exemple de validare. Crește cu fiecare decizie suplimentară și cu fiecare micșorare a setului de validare. Observă și că direcția nu este garantată într-o singură rulare: pragul ales a obținut 0,3902 pe validare și 0,4122 pe test, deci validarea l-a subestimat de data aceasta. Bias-ul este sistematic peste multe decizii, nu vizibil într-una.4
Acum exercițiul. Jurnalul benzii vine cu o a treia coloană, station_seconds: cât timp a petrecut fiecare piesă la stația de inspecție. Adăugarea ei este o schimbare de o linie în preprocesare. Iată ce face:
| model | acuratețe | precizie | recall | F1 | cross-entropy | AUC |
|---|---|---|---|---|---|---|
| lățime + greutate | 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 |
Recall crește de la 10,8 % la 77,0 %. F1 se multiplică de peste patru ori. Și observă ce a făcut acuratețea: 98,30 % → 99,20 %, un câștig de nouă zecimi de punct, genul de număr rotunjit la „aproximativ 99 % oricum” într-un slide de sumar. Acuratețea nu a văzut eșecul mai devreme și acum nu vede frauda.
Înainte să citești mai departe: modelul trișează. Află cum.
Cum vânezi un leak, în ordinea care îl găsește cel mai repede.
-
Compară train și test. Overfitting-ul apare ca o diferență mare. Aici: model onest 0,9838 train / 0,9830 test; model cu leak 0,9936 train / 0,9920 test. Ambele diferențe sunt sub 0,2 puncte. Un leak nu arată ca overfitting — feature-ul cu leak este la fel de disponibil la test, deci modelul generalizează superb într-o lume care nu există.
-
Antrenează câte un model pe fiecare feature, singur. Orice poartă răspunsul se va anunța singur:
feature singur acuratețe recall F1 AUC lățime 0,9815 0,014 0,026 0,8691 greutate 0,9815 0,000 0,000 0,7914 station_seconds0,9850 0,405 0,500 0,9960 O coloană, de una singură, ordonează defectele la AUC 0,9960. Două măsurători luate cu un șubler și un cântar reușesc 0,87 și 0,79. Această asimetrie este alarma.
-
Întreabă când a fost notat fiecare număr. Timp mediu de staționare: 2,23 secunde pentru piesele care au trecut, 15,56 secunde pentru cele care au picat. Sigur că da. O piesă stă la stație pentru că un inspector a scos-o de pe bandă — ceea ce se întâmplă după, și doar pentru că, cineva a decis că era defectă. Coloana nu este o măsurare a piesei. Este o măsurare a verdictului.
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()) Linia evidențiată este leak-ul: timpul de staționare al unei piese defecte este extras dintr-o distribuție diferită, pentru că un om a scos-o de pe bandă. Acesta este cel mai comun bug serios în machine learning aplicat și are un nume: target leakage — informație în feature-urile de antrenare care nu ar fi disponibilă în momentul în care predicția trebuie făcută.5 Nu aruncă nicio excepție. Produce un număr mai bun. Fiecare stimulent dintr-un proiect împinge spre păstrarea lui.
Apărarea este o întrebare, pusă fiecărei coloane: în clipa în care am nevoie de această predicție, valoarea aceasta există deja? Pe o bandă live, station_seconds este necunoscut până după ce piesa a fost inspectată — adică exact lucrul pe care modelul trebuia să-l înlocuiască.
De câte exemple de test am nevoie?
Link către secțiunea: De câte exemple de test am nevoie?Să presupunem că evaluezi un model pe 20 de exemple și nimerește 17. Raportezi 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.6477Citirea onestă a lui 17/20 este undeva între 64 % și 95 %. Un model cu adevărat de 65 % produce acest rezultat în 4,4 % din cazuri — o rulare din douăzeci și trei — iar dacă ai încercat o mână de prompts și ai raportat cel mai bun, ai fabricat chiar tu acea rulare. Șaptesprezece din douăzeci nu poate distinge un model de 85 % de unul de 65 %.
Două moduri de a pune un interval pe o rată, iar ambele își au locul în trusa ta:
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)Folosește Wilson6 pentru o rată simplă de succes; se comportă bine la orice și nu are nevoie de randomness. Observă mai sus că la capătul superior al bootstrap este 1,0000 — reeșantionarea a 20 de puncte poate trage ușor 20 corecte, deci nu poate reprezenta un interval mai îngust decât propria granularitate. Folosește bootstrap7 acolo unde nu există formulă, adică în majoritatea cazurilor interesante: F1, macro-medii, BLEU, pass@1, scorul unui judecător bazat pe rubrică. Pe această bandă, F1-ul modelului ajustat, 0,4122, poartă un interval bootstrap de [0,3009, 0,5156] — acesta este numărul care ar trebui să apară în raport, pentru că estimarea punctuală singură invită o comparație pe care nu o poate susține.
Încă o măsurătoare, pentru că schimbă felul în care ar trebui să compari două modele. Două modele evaluate pe aceleași 500 de exemple:
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)Intervalele lor se suprapun, iar regula populară — bare de eroare suprapuse înseamnă nicio diferență semnificativă — ar numi comparația neconcludentă. Nu este. Cele două modele au rulat pe aceleași exemple, deci cantitatea corectă este diferența pe exemplu, al cărei interval este [0.0260, 0.0680], confortabil peste zero. Nu sunt de acord pe doar 31 din 500 de itemi, iar A câștigă 27 dintre aceste dezacorduri; exemplele comune, ușoare și grele deopotrivă, se anulează în loc să adauge zgomot. Compară modelele pereche și ajungi la aceeași concluzie cu o fracțiune din date.
Încotro mergem mai departe
Link către secțiunea: Încotro mergem mai departeAi acum un model care produce probabilități calibrate, o pierdere derivată dintr-o afirmație despre date, nu aleasă pentru comoditate, un gradient care este literalmente predicția minus adevărul și — mai important — mecanismul pentru a afla dacă ceva din toate acestea funcționează. Intervalul Wilson de zece linii de mai sus este refolosit verbatim: el susține variantele de prompt din Capitolul 15, tabelele de retrieval din Capitolul 19 și setul golden din Capitolul 29. Bootstrap este la ce apelezi când nu există formulă.
Dar modelul încă are un singur strat. Desenează o linie, iar Capitolul 1 a demonstrat cu patru rânduri de XOR că o linie nu este suficientă. Soluția este să stivuiești: un prim strat care îndoaie spațiul, un al doilea care desenează linia în spațiul îndoit.
Acolo se termină gradientul curat din acest capitol. Tot ce a fost mai sus a funcționat pentru că putea fi scris de mână, o singură dată, pentru un model cu un strat între input și pierdere. Pune un al doilea strat la mijloc și întrebarea își schimbă forma: care este derivata pierderii în raport cu o greutate care nu atinge deloc ieșirea — una a cărei influență ajunge doar printr-un alt strat, posibil de-a lungul mai multor căi simultan?
Acea derivată există. Calcularea ei de mână este fără speranță pentru orice lucru mai mare decât o jucărie, iar calcularea ei câte un parametru pe rând este fără speranță la o altă scară. Este nevoie de o procedură care obține fiecare derivată din rețea dintr-o singură trecere înapoi peste același graf pe care trecerea înainte tocmai l-a parcurs.
Acesta este Capitolul 5, iar el este motorul pe care rulează restul acestui curs.
Surse și metodă
Link către secțiunea: Surse și metodăMerită citite alături de acest capitol și: Bishop, Pattern Recognition and Machine Learning §1.2, §1.5, §1.6 și §4.3, care acoperă probabilitatea, teoria deciziei, teoria informației și clasificarea liniară în ordinea urmată de acest capitol; Murphy, Probabilistic Machine Learning: An Introduction, capitolele 6 și 10; Prince, Understanding Deep Learning §5.4–5.7; și Saito și Rehmsmeier, The Precision-Recall Plot Is More Informative than the ROC Plot When Evaluating Binary Classifiers on Imbalanced Datasets (PLOS ONE, 2015) — de ce AUC citat mai sus nu ar trebui să fie singurul număr fără prag la care te uiți când 1,7 % dintre piese sunt defecte.
Referințe
Link către secțiunea: Referințe-
Ma, T. și Ng, A. CS229 Lecture Notes, Stanford University, capitolele 2 și 3. Unde anularea care produce încetează să mai pară noroc: alege distribuția din familia exponențială care se potrivește ieșirii tale, folosește legătura ei canonică, iar gradientul este mereu predicția minus adevărul. ↩
-
Olah, C. Visual Information Theory (2015),
colah.github.io/posts/2015-09-Visual-Information. Cea mai clară explicație disponibilă a entropiei, cross-entropy și divergenței KL ca costuri în biți, nu ca formule. ↩ -
Abu-Mostafa, Y. S., Magdon-Ismail, M. și Lin, H.-T. Learning From Data (AMLBook, 2012), cursurile 13 și 17 din cursul Caltech. Cursul 13 este despre validare; cursul 17, despre cele trei principii ale învățării, este locul unde data snooping este numit. Împreună sunt sursa disciplinei din acest capitol: fiecare privire asupra unui set de date este o decizie de fitting, indiferent dacă ai rulat sau nu un optimizator. ↩
-
James, G., Witten, D., Hastie, T. și Tibshirani, R. An Introduction to Statistical Learning, ediția a 2-a (Springer, 2021), capitolele 2 și 5, pentru descompunerea bias–variance și pentru resampling. Volumul companion este locul unde capcana selecției este formulată direct: Hastie, Tibshirani și Friedman, The Elements of Statistical Learning, ediția a 2-a, §7.10.2, The Wrong and Right Way to Do Cross-validation. ↩
-
Kaufman, S., Rosset, S., Perlich, C. și Stitelman, O. Leakage in Data Mining: Formulation, Detection, and Avoidance. ACM Transactions on Knowledge Discovery from Data 6(4), 2012. Un tratament formal al eșecului demonstrat mai sus, cu studii de caz din competiții câștigate de un model care învățase un artefact al felului în care datele au fost asamblate. ↩
-
Wilson, E. B. Probable Inference, the Law of Succession, and Statistical Inference. Journal of the American Statistical Association 22(158), pp. 209–212 (1927). Intervalul score folosit în
wilson()de mai sus, încă default-ul corect pentru o proporție. Intervalul de manual este cel de evitat: dă absurdități aproape de 0 și 1 și subacoperă grav la mic. ↩ -
Efron, B. Bootstrap Methods: Another Look at the Jackknife. The Annals of Statistics 7(1), pp. 1–26 (1979). Ideea care îți permite să pui un interval pe orice statistică poți calcula, inclusiv pe cele fără teorie de eșantionare. ↩