Classificatie, cross-entropy en hoe je jezelf niet voor de gek houdt
Bouw een logistic classifier en ontdek waarom 98% accuracy kan betekenen dat je model helemaal niets vindt.
Op deze pagina
Een model dat bij elk onderdeel van de band antwoordt dit onderdeel is in orde, heeft in 98,15% van de gevallen gelijk. Het is ook waardeloos: van de 74 defecte onderdelen in de testset vindt het er geen één.
Beide zinnen beschrijven hetzelfde model. De afstand ertussen is dit hoofdstuk.
De eerste helft bouwt de classifier. Daar is bijna niets nieuws voor nodig: Hoofdstuk 2 gaf het recept om een aanname over hoe data wordt geproduceerd om te zetten in een loss function, en Hoofdstuk 3 gaf de machinerie om bergafwaarts te lopen over welke loss dat recept ook oplevert. Pas beide toe op een ja/nee-vraag en logistic regression valt eruit, plus één nieuw idee — een logit — waarvoor opnieuw betaald wordt in Hoofdstuk 17.
De tweede helft is de moeilijkere. Alles na dit punt in de cursus wordt beoordeeld met een getal dat iemand heeft gemeten, en als je een echte verbetering niet kunt onderscheiden van een meetartefact, is elk volgend hoofdstuk decoratie. Dus: de confusion matrix, precision en recall, de drie splits, leakage, en de vraag die bijna niemand eerlijk beantwoordt — hoeveel testvoorbeelden heb ik eigenlijk nodig?
De rekenkunde hier loopt over 20.000 rijen, dus alles is gevectoriseerd — NumPy doet het werk al sinds Hoofdstuk 2, en vanaf nu is het niet meer de moeite waard om dat steeds te melden.
De band, met een zeldzamere vraag
Link naar de sectie: De band, met een zeldzamere vraagDezelfde fabriek als in Hoofdstuk 1, maar een moeilijkere vraag. In plaats van accepteren of afkeuren is de vraag is dit onderdeel defect — en defecten zijn zeldzaam, wat de meethelft van dit hoofdstuk moeilijk maakt en de modelleerhelft bedrieglijk makkelijk.
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 74Drie splits, niet twee. De reden verdient een eigen sectie en krijgt die hieronder; train voorlopig op de eerste, tune op de tweede, en kijk niet naar de derde.
De features zijn gestandaardiseerd — gemiddelde eraf, gedeeld door de standaarddeviatie — met alleen de trainingsstatistieken, om de reden die Hoofdstuk 1 liet zien met de convergentiegrens van de perceptron: ongecentreerde data maakt de geometrie vijandig. Uit welke rijen je dat gemiddelde mag berekenen wordt later in dit hoofdstuk een levende vraag.
Van oordeel naar probability
Link naar de sectie: Van oordeel naar probabilityDe perceptron gaf een teken terug. Een teken kan afkeuren niet onderscheiden van afkeuren, maar net aan, en precies dat verschil heeft een fabriek nodig om te bepalen welke onderdelen een mens als eerste opnieuw moet inspecteren.
Volg dus het recept uit Hoofdstuk 2 letterlijk. Schrijf op wat je beweert over hoe een label wordt geproduceerd, neem de likelihood, neem de log, maak die negatief, en je hebt een loss. Voor een ja/nee-uitkomst is de bewering een Bernoulli-verdeling: er is een probability dat het onderdeel defect is, en
wat gewoon een compacte manier is om te schrijven ‘ als , en als ’. Neem daarvan de log en maak die negatief, en de loss voor één voorbeeld is
Dit is binary cross-entropy. Die is niet gekozen omdat hij handig is; hij is de negative log-likelihood van de enige verdeling die een coin flip kan hebben. Iets anders was er niet.
Wat nog ontbreekt is waar vandaan komt. Het model berekent een gewogen som , een reëel getal dat over de hele getallenlijn loopt, en een probability moet in liggen. De functie die tussen die twee beweegt is de 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.9820Lees de rechterkolom als een prijslijst. Gelijk hebben met 90% confidence kost 0,105. Niet willen kiezen kost 0,693 — dat is , de prijs van een schouderophalen. Vol vertrouwen fout zitten kost 4,6, vierenveertig keer zoveel, en de prijs stijgt onbeperkt naarmate het model zekerder wordt van een vergissing. Cross-entropy telt fouten niet alleen: het rekent arrogantie af.
De gradiënt is prediction minus truth
Link naar de sectie: De gradiënt is prediction minus truthHoofdstuk 3 zei: om wat dan ook te trainen, neem de afgeleide van de loss naar elke parameter. Doe dat voor één voorbeeld. Met en :
Details tonen
De twee regels die de rommel laten wegvallen. De sigmoid heeft een ongewoon prettige afgeleide, . En de loss differentieert naar
Vermenigvuldig die twee met de kettingregel en verschijnt één keer boven en één keer onder. Het valt exact weg, en blijft over. Die wegstreping is geen toeval — dit gebeurt telkens wanneer de loss de negative log-likelihood van een verdeling is en de outputfunctie degene is die die verdeling van nature gebruikt. Dat koppel heeft een naam — een generalised linear model — en de nette gradiënt is zijn vingerafdruk.1
De update is dus prediction minus truth, maal de input. Niets anders. Hier is de hele trainer, de descent uit Hoofdstuk 3 met één aangepaste regel:
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, bDe np.where in sigmoid is niet cosmetisch. direct berekenen loopt over bij grote negatieve ; de branch kiest de algebraïsch identieke vorm die de exponent negatief houdt. Dit is de floating-point-box uit Hoofdstuk 2 die zijn eerste schuld int, en twee secties verder int hij een grotere.
Waarom geen squared error, en waarom het antwoord over de gradiënt gaat
Link naar de sectie: Waarom geen squared error, en waarom het antwoord over de gradiënt gaatDe standaarduitleg om cross-entropy boven squared error te verkiezen is het likelihood-argument hierboven: squared error krijg je als je Gaussian noise aanneemt, labels zijn niet Gaussian, dus doe het niet. Dat klopt en het overtuigt niemand, want je kunt over een sigmoid schrijven en het zal trainen.
Het argument dat landt gaat over de gradiënt. Zet squared error boven op een sigmoid en de kettingregel geeft
Die extra is degene die eerder wegviel. Nu niet, en hij gaat naar nul wanneer het model zeker is — ook wanneer het model vol vertrouwen fout zit. Evalueer beide bij een paar scores, voor een voorbeeld waarvan het echte label 1 is:
| score | cross-entropy | squared error | ratio | |
|---|---|---|---|---|
| 0,000335 | 1.491 | |||
| 0,017986 | 28,3 | |||
| 0,119203 | 4,8 | |||
| 0,500000 | 2,0 | |||
| 0,880797 | 4,8 |
Bij zit het model zo fout als maar kan, en squared error reageert met een gradiënt die 1.491 keer kleiner is dan die van cross-entropy. Hoe erger de fout, hoe minder het model ervan leert. De gradiënt van cross-entropy verzadigt ondertussen op : maximaal fout levert een maximaal groot signaal op, en niet groter.
Laat ze racen. Tweeduizend gebalanceerde punten, identieke beginweights gekozen om vol vertrouwen fout te zijn (), identieke learning rate, alleen de loss verschilt. Beide runs worden gescoord met cross-entropy zodat de kolommen vergelijkbaar zijn.
| 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 |
Cross-entropy is klaar bij epoch 50. Squared error zit bij epoch 100 nog steeds op 24% accuracy — en was sinds epoch 10 niet van 23% gekomen — erger dan gokken, omdat het vol vertrouwen fout begon en de gradiënt die het had moeten redden met 0,0007 is vermenigvuldigd. Het ontsnapt rond epoch 500 en komt op dezelfde plek uit. De eerlijke samenvatting is dus dat squared error over een sigmoid niet incorrect is; het is traag precies waar snelheid het belangrijkst is. Op een model met twee parameters verlies je 450 epochs. Op een netwerk met honderd lagen, waar ergens altijd wel een unit vol vertrouwen fout zit, verlies je de trainingsrun.
Entropy, cross-entropy en KL, op één pagina
Link naar de sectie: Entropy, cross-entropy en KL, op één paginaDrie grootheden, later goed nodig in Hoofdstuk 8 voor perplexity en in Hoofdstuk 11 voor de penalty die een fine-tuned policy dicht bij zijn reference houdt. Ze zijn makkelijker dan hun reputatie.2
Entropy is het gemiddelde aantal bits dat je moet besteden om een trekking uit een verdeling te communiceren, als je er de best mogelijke code voor gebruikt:
Cross-entropy is wat je besteedt wanneer je een code gebruikt die gebouwd is voor op data die eigenlijk uit komt:
KL divergence is het overschot — de verspilling, in bits, veroorzaakt doordat je gelooft terwijl de waarheid is:
Controleer alle drie op de 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 bitsTwee dingen zijn daar zichtbaar. Ten eerste haalt een model dat simpelweg de base rate van de training rapporteert, 1,69%, een cross-entropy van 0,1330 bits, bijna precies de entropy van de testlabels — zoals het moet, omdat het de juiste verdeling heeft en geen andere informatie. Entropy is de bodem die onwetendheid-over-het-individu je oplevert. Ten tweede betaalt een model dat zijn schouders ophaalt en 0,5 zegt exact 1 bit, en de kloof tussen die twee, 0,8671 bits, is precies de KL divergence. is geen identiteit om uit je hoofd te leren; het is een rekening die je kunt zien oplopen.
En de verbinding terug naar training: wanneer het label één bekende class is, is de ‘true’ verdeling one-hot, is de entropy nul, en is cross-entropy gelijk aan de KL divergence. Cross-entropy minimaliseren en de verdeling van het model naar de waarheid trekken zijn dezelfde handeling.
Meer dan twee antwoorden: softmax, en de verschuiving die niets kost
Link naar de sectie: Meer dan twee antwoorden: softmax, en de verschuiving die niets kostDefect is niet één ding. Bij spuitgieten kan een onderdeel eruit komen als een short shot (niet genoeg materiaal), flash (te veel, uit de mal geperst), of burn. Vier uitkomsten, dus vier logits, en die moeten vier probabilities worden die optellen tot één. Dat is softmax:
Hij heeft een eigenschap die op een ongeluk lijkt en in feite de hele implementatie is:
voor elke constante , omdat en boven en onder wegvallen. Alleen verschillen tussen logits betekenen iets. Het absolute niveau is geen informatie.
Gelukkig maar, want het absolute niveau is wat de computer breekt:
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 laat een 64-bit float overlopen, de som wordt infinity, en infinity gedeeld door infinity is nan — geen error, geen crash, alleen een stil gat waar eerst drie probabilities stonden. De maximale logit aftrekken verandert wiskundig niets en numeriek alles, omdat de grootste exponent precies wordt. Dit is de logsumexp-truc uit Hoofdstuk 2 in werkkleding, en elke serieuze implementatie doet het:
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, bDe gradiënt is opnieuw prediction minus truth, nu met one-hot. De binary case was al die tijd een speciaal geval.
Getraind op 3.000 onderdelen en getest op 1.000, met drie metingen per stuk (breedte, gewicht, smelttemperatuur), haalt het 94,00% accuracy. Dit is wat dat getal verbergt:
| truth ↓ / predicted → | 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 |
Het model vindt minder dan de helft van de short shots. Accuracy ziet dat niet, omdat 86% van de onderdelen in orde is en die goed hebben genoeg is om het gemiddelde te dragen. Macro F1 — het gemiddelde van de F1-scores per class, waarbij een zeldzame class even zwaar weegt als een veelvoorkomende — is 0,7983, tegenover een micro F1 van 0,9400 die per definitie identiek is aan accuracy. Wanneer iemand één F1-getal rapporteert, vraag dan welke.
Dat was het laatste van het modelleren. De rest van het hoofdstuk gaat over de getallen.
Drie modellen, één accuracy
Link naar de sectie: Drie modellen, één accuracyNeem het getrainde binary model en maak twee varianten door elke logit met een constante te vermenigvuldigen: 0,35 voor een aarzelende versie, 4 voor een overconfident versie. Vermenigvuldigen met een positief getal kan geen enkel teken veranderen, dus alle drie modellen voorspellen exact hetzelfde label voor alle 4.000 testonderdelen. Accuracy kan ze niet uit elkaar houden. Cross-entropy heeft daar geen enkele moeite mee:
| model | accuracy | cross-entropy | gemiddelde loss als goed | gemiddelde loss als fout | ergste enkele loss |
|---|---|---|---|---|---|
| aarzelend (logits × 0,35) | 0,9830 | 0,1549 | 0,1369 | 1,1990 | 2,80 |
| zoals getraind | 0,9830 | 0,0564 | 0,0147 | 2,4689 | 7,82 |
| overconfident (logits × 4) | 0,9830 | 0,1563 | 0,0009 | 9,1427 | 27,63 |
Het aarzelende model betaalt een kleine belasting op elk onderdeel, ook op de duizenden die het goed heeft. Het overconfident model is bijna gratis wanneer het gelijk heeft en catastrofaal wanneer het fout zit — één onderdeel in die testset kost het in zijn eentje 27,63 nats. De twee komen bijna op hetzelfde totaal uit via tegengestelde routes, en het getrainde model, waarvan de probabilities op de data zijn gekalibreerd, zit drie keer lager dan allebei.
Dit is de scherpste manier om het verschil tussen een loss en een metric te formuleren. De loss is wat je optimaliseert: hij moet differentieerbaar zijn, en hij ziet alles wat het model zei, inclusief hoe zeker het was. De metric is waarop je wordt beoordeeld: hij kan een step function zijn, een business rule, een telling van gemiste defecten. Ze zijn niet hetzelfde object en ze zijn het niet altijd eens — daarom definieer je beide voordat je begint, en laat je de loss nooit invallen voor de metric alleen omdat hij toevallig op het scherm staat.
De domme baseline gaat eerst
Link naar de sectie: De domme baseline gaat eerstVóór elk model komt de eis: wat scoort het luiest mogelijke antwoord? Op deze band: zeg altijd in orde:
always-say-fine baseline: accuracy = 0.9815
confusion (tn, fp, fn, tp) = (3926, 0, 74, 0)98,15%. Nu het getrainde logistic model, bij de standaardthreshold van 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%. Het versloeg de baseline met 0,15 procentpunt, en elk rapport dat bij accuracy stopt zal dat een winst noemen. De confusion matrix zegt wat er echt gebeurde:
| predicted fine | predicted defective | |
|---|---|---|
| actually fine | 3.924 | 2 |
| actually defective | 66 | 8 |
Het vond 8 defecte onderdelen van de 74 en liet er 66 door. Drie getallen benoemen de drie manieren om die tabel te lezen:
- Precision . Van de onderdelen die het markeerde, hoeveel waren er echt defect. Dit is de kost van verspilde inspecties.
- Recall . Van de defecte onderdelen, hoeveel ving het er. Dit is de kost van een slecht onderdeel naar een klant sturen.
- F1 , hun harmonisch gemiddelde, dat dicht bij de kleinste van de twee blijft en zich daarom niet door één van beide alleen laat vleien.
Wat ertoe doet hangt af van de fabriek, niet van de wiskunde: een inspectie kost een paar seconden en een verzonden defect kost een terugroepbericht, dus hier domineert recall en is 0,108 een mislukking.
Maar het model is niet het probleem. De threshold is dat wel, en de threshold is geen onderdeel van het model — het is een business decision die achteraf op een probability wordt toegepast. Sweep hem:
| threshold | 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 |
Lees de accuracy-kolom omlaag. Hij daalt de hele weg — van 98,30% naar 65,93% — terwijl het model van 8 defecten vinden naar 71 van de 74 gaat. Alles nuttigs wat dit model kan doen maakt zijn accuracy slechter. Een team dat het headline-getal optimaliseert zou de versie shippen die niets vindt.
Details tonen
Class weighting creëert geen signaal, het verplaatst het operating point. De gebruikelijke eerste reflex bij imbalanced classes is de zeldzame class zwaarder wegen in de loss. Doe dat, met weights van 1, 10 en 60 op de positives:
| weight op positives | 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 en recall bewegen enorm. De AUC — de probability dat het model een willekeurig defect onderdeel boven een willekeurig goed onderdeel rankt, en de threshold volledig negeert — beweegt met 0,0002, dus eigenlijk niet. Reweighting schoof hetzelfde model langs dezelfde trade-off curve. Dat is vaak wat je wilt, en het is nooit nieuwe informatie: als de ranking slecht is, redt geen enkel weighting scheme hem.
Drie splits, en de leak die je zo gaat vinden
Link naar de sectie: Drie splits, en de leak die je zo gaat vindenWaarom drie splits en geen twee? Omdat zodra je een set voorbeelden gebruikt om iets te kiezen — een threshold, een learning rate, welk van zes modellen je shipt — die set is gebruikt voor fitting, en zijn score niet langer unbiased is.3 Gemeten op deze band: de threshold sweep op de validation set kiest 0,196, en het model scoort daarna F1 = 0,4122 op de onaangeraakte testset. Was de sweep direct op de testset uitgevoerd, dan was het best haalbare daar 0,4186 geweest — een getal dat niemand mag rapporteren.
De kloof is hier klein, 0,006, omdat dit één hyperparameter is, één keer gesweept tegen 4.000 validation examples. Hij groeit met elke extra beslissing en elke krimp van de validation set. Let ook op dat de richting in één run niet gegarandeerd is: de gekozen threshold scoorde 0,3902 op validation en 0,4122 op test, dus validation onderschatte hem deze keer. De bias is systematisch over veel beslissingen, niet zichtbaar in één.4
Nu de oefening. Het bandlog komt binnen met een derde kolom, station_seconds: hoelang elk onderdeel bij het inspectiestation stond. Die toevoegen is een éénregelige wijziging in de preprocessing. Dit is wat hij doet:
| model | accuracy | precision | recall | F1 | cross-entropy | AUC |
|---|---|---|---|---|---|---|
| width + weight | 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 gaat van 10,8% naar 77,0%. F1 wordt meer dan vier keer zo groot. En let op wat accuracy deed: 98,30% → 99,20%, een winst van negen tiende punt, het soort getal dat in een samenvattingsslide wordt afgerond tot ‘ongeveer 99%, hoe dan ook’. Accuracy zag eerder de mislukking niet en ziet nu de fraude niet.
Lees nog niet verder: het model speelt vals. Zoek uit hoe.
Hoe je een leak opspoort, in de volgorde die hem het snelst vindt.
-
Vergelijk train en test. Overfitting verschijnt als een grote kloof. Hier: eerlijk model 0,9838 train / 0,9830 test; leaky model 0,9936 train / 0,9920 test. Beide kloven zijn kleiner dan 0,2 punten. Een leak lijkt niet op overfitting — de leaky feature is op testtijd net zo beschikbaar, dus het model generaliseert prachtig naar een wereld die niet bestaat.
-
Train één model per feature, alleen. Alles wat het antwoord draagt, zal zichzelf aankondigen:
feature alleen accuracy recall F1 AUC width 0,9815 0,014 0,026 0,8691 weight 0,9815 0,000 0,000 0,7914 station_seconds0,9850 0,405 0,500 0,9960 Eén kolom, op zichzelf, rankt defecten met AUC 0,9960. Twee metingen door een schuifmaat en een weegschaal halen 0,87 en 0,79. Die asymmetrie is het alarm.
-
Vraag wanneer elk getal is opgeschreven. Gemiddelde verblijftijd: 2,23 seconden voor onderdelen die slaagden, 15,56 seconden voor onderdelen die faalden. Natuurlijk. Een onderdeel blijft bij het station omdat een inspecteur het van de band haalde — wat gebeurt nadat, en alleen omdat, iemand besloot dat het defect was. De kolom is geen meting van het onderdeel. Het is een meting van het oordeel.
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()) De gemarkeerde regel is de leak: de verblijftijd van een defect onderdeel wordt uit een andere verdeling getrokken, omdat een mens het van de band haalde. Dit is de meest voorkomende ernstige bug in toegepaste machine learning, en hij heeft een naam: target leakage — informatie in de training features die niet beschikbaar zou zijn op het moment dat de voorspelling gemaakt moet worden.5 Hij gooit geen exception. Hij produceert een beter getal. Elke prikkel in een project wijst naar hem behouden.
De verdediging is één vraag, gesteld aan elke kolom: op het moment dat ik deze voorspelling nodig heb, bestaat deze waarde dan al? Op een live band is station_seconds onbekend tot nadat het onderdeel is geïnspecteerd — precies datgene wat het model moest vervangen.
Hoeveel testvoorbeelden heb ik nodig?
Link naar de sectie: Hoeveel testvoorbeelden heb ik nodig?Stel dat je een model op 20 voorbeelden scoort en het heeft er 17 goed. Je rapporteert 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.6477De eerlijke lezing van 17/20 is ergens tussen 64% en 95%. Een echt 65%-model produceert dit resultaat 4,4% van de tijd — één run op drieëntwintig — en als je een handvol prompts probeerde en de beste rapporteerde, heb je die run zelf gefabriceerd. Zeventien van de twintig kan een 85%-model niet onderscheiden van een 65%-model.
Twee manieren om een interval op een percentage te zetten, en beide horen in je 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)Gebruik Wilson6 voor een eenvoudig success rate; hij blijft zich goed gedragen bij elke en heeft geen randomness nodig. Let hierboven op dat bij de bovenkant van de bootstrap 1,0000 is — resamplen van 20 punten kan makkelijk 20 correcte trekken, dus hij kan geen interval weergeven dat smaller is dan zijn eigen granulariteit. Gebruik de bootstrap7 waar geen formule bestaat, en dat zijn de meeste interessante gevallen: F1, macro-gemiddelden, BLEU, pass@1, de score van een rubric-based judge. Op deze band heeft de F1 van 0,4122 van het getunede model een bootstrap-interval van [0,3009, 0,5156] — en dat is het getal dat in het rapport moet staan, omdat de puntschatting alleen een vergelijking uitnodigt die hij niet kan dragen.
Nog één meting, omdat die verandert hoe je twee modellen moet vergelijken. Twee modellen gescoord op dezelfde 500 voorbeelden:
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)Hun intervallen overlappen, en de vuistregel — overlappende error bars betekent geen significant verschil — zou de vergelijking inconclusief noemen. Dat is ze niet. De twee modellen draaiden op dezelfde voorbeelden, dus de juiste grootheid is het verschil per voorbeeld, waarvan het interval [0.0260, 0.0680] is, comfortabel boven nul. Ze zijn het oneens over slechts 31 van de 500 items, en A wint 27 van die meningsverschillen; de gedeelde voorbeelden, makkelijk en moeilijk, vallen tegen elkaar weg in plaats van noise toe te voegen. Vergelijk modellen paired, en je bereikt dezelfde conclusie met een fractie van de data.
Waar dit heen gaat
Link naar de sectie: Waar dit heen gaatJe hebt nu een model dat gekalibreerde probabilities output, een loss afgeleid uit een claim over de data in plaats van gekozen voor gemak, een gradiënt die letterlijk prediction minus truth is, en — belangrijker — de machinerie om uit te zoeken of iets ervan werkt. Het tienregelige Wilson-interval hierboven wordt letterlijk hergebruikt: het draagt de prompt-varianten in Hoofdstuk 15, de retrieval-tabellen in Hoofdstuk 19, en de golden set in Hoofdstuk 29. De bootstrap is waar je naar grijpt wanneer er geen formule bestaat.
Maar het model heeft nog steeds één laag. Het tekent een lijn, en Hoofdstuk 1 bewees met vier rijen XOR dat een lijn niet genoeg is. De fix is stapelen: een eerste laag die de ruimte buigt, een tweede die de lijn tekent in de gebogen ruimte.
Daar loopt de nette gradiënt van dit hoofdstuk op zijn einde. Alles hierboven werkte omdat één keer met de hand kon worden opgeschreven, voor een model met één laag tussen de input en de loss. Zet een tweede laag in het midden en de vraag verandert van vorm: wat is de afgeleide van de loss naar een weight die de output helemaal niet aanraakt — één waarvan de invloed alleen via een andere laag aankomt, mogelijk langs meerdere paden tegelijk?
Die afgeleide bestaat. Hem met de hand berekenen is hopeloos voor alles groter dan speelgoed, en hem één parameter tegelijk berekenen is hopeloos op een andere schaal. Wat nodig is, is een procedure die elke afgeleide in het netwerk uit één backward pass over dezelfde graph haalt waar de forward pass net doorheen liep.
Dat is Hoofdstuk 5, en het is de engine waarop de rest van deze cursus draait.
Bronnen en methode
Link naar de sectie: Bronnen en methodeOok de moeite waard om naast dit hoofdstuk te lezen: Bishop, Pattern Recognition and Machine Learning §1.2, §1.5, §1.6 en §4.3, dat probability, decision theory, information theory en linear classification behandelt in de volgorde die dit hoofdstuk volgt; Murphy, Probabilistic Machine Learning: An Introduction, hoofdstukken 6 en 10; Prince, Understanding Deep Learning §5.4–5.7; en Saito en Rehmsmeier, The Precision-Recall Plot Is More Informative than the ROC Plot When Evaluating Binary Classifiers on Imbalanced Datasets (PLOS ONE, 2015) — waarom de hierboven geciteerde AUC niet het enige threshold-free getal zou moeten zijn waar je naar kijkt wanneer 1,7% van de onderdelen defect is.
Referenties
Link naar de sectie: Referenties-
Ma, T. en Ng, A. CS229 Lecture Notes, Stanford University, hoofdstukken 2 en 3. Waar de wegstreping die oplevert ophoudt op geluk te lijken: kies de exponential-family distribution die bij je output past, gebruik zijn canonical link, en de gradiënt is altijd prediction minus truth. ↩
-
Olah, C. Visual Information Theory (2015),
colah.github.io/posts/2015-09-Visual-Information. De helderste beschikbare uitleg van entropy, cross-entropy en KL divergence als kosten in bits in plaats van als formules. ↩ -
Abu-Mostafa, Y. S., Magdon-Ismail, M. en Lin, H.-T. Learning From Data (AMLBook, 2012), colleges 13 en 17 van de Caltech-cursus. College 13 is validation; college 17, over de drie learning principles, is waar data snooping een naam krijgt. Samen zijn ze de bron van de discipline in dit hoofdstuk: elke blik op een dataset is een fitting decision, of je nu een optimiser hebt gedraaid of niet. ↩
-
James, G., Witten, D., Hastie, T. en Tibshirani, R. An Introduction to Statistical Learning, 2e editie (Springer, 2021), hoofdstukken 2 en 5, voor de bias–variance decomposition en voor resampling. Het companion volume is waar de selection trap ronduit wordt benoemd: Hastie, Tibshirani en Friedman, The Elements of Statistical Learning, 2e editie, §7.10.2, The Wrong and Right Way to Do Cross-validation. ↩
-
Kaufman, S., Rosset, S., Perlich, C. en Stitelman, O. Leakage in Data Mining: Formulation, Detection, and Avoidance. ACM Transactions on Knowledge Discovery from Data 6(4), 2012. Een formele behandeling van de fout die hierboven is gedemonstreerd, met case studies uit competities gewonnen door een model dat een artefact had geleerd van hoe de data was samengesteld. ↩
-
Wilson, E. B. Probable Inference, the Law of Succession, and Statistical Inference. Journal of the American Statistical Association 22(158), pp. 209–212 (1927). Het score-interval dat hierboven in
wilson()wordt gebruikt, nog steeds de juiste default voor een proportie. Het textbook-interval is degene die je moet vermijden: het geeft nonsens dicht bij 0 en 1, en undercovers sterk bij kleine . ↩ -
Efron, B. Bootstrap Methods: Another Look at the Jackknife. The Annals of Statistics 7(1), pp. 1–26 (1979). Het idee waarmee je een interval kunt zetten op elke statistic die je kunt berekenen, inclusief degene zonder sampling theory. ↩