Klasifikace, křížová entropie a jak neobelhat sami sebe
Sestavte logistický klasifikátor a zjistěte, proč 98 % přesnost může znamenat model, který nenajde vůbec nic.
Na této stránce
Model, který o každém dílu sjíždějícím z pásu odpoví tahle součástka je v pořádku, má pravdu v 98,15 % případů. Zároveň je k ničemu: ze 74 vadných dílů v testovací sadě nezachytí ani jeden.
Obě věty popisují tentýž model. Vzdálenost mezi nimi je tato kapitola.
První polovina klasifikátor sestaví. Nepotřebuje skoro nic nového: kapitola 2 dala recept, jak převést předpoklad o tom, jak data vznikají, na ztrátovou funkci, a kapitola 3 dala aparát pro sestup z kopce po jakékoli ztrátě, kterou vám tento recept předá. Použijte obojí na otázku ano/ne a vypadne z toho logistická regrese plus jedna nová myšlenka — logit — ke které se znovu vrátíme v kapitole 17.
Druhá polovina je ta těžší. Všechno od tohoto místa v kurzu se posuzuje podle čísla, které někdo změřil, a pokud nedokážete rozlišit skutečné zlepšení od artefaktu měření, každá další kapitola je jen ozdoba. Takže: matice záměn, precision a recall, tři splity, leakage a otázka, na kterou skoro nikdo neodpovídá poctivě — kolik testovacích příkladů vlastně potřebuji?
Výpočty zde běží přes 20 000 řádků, takže jsou celé vektorizované — NumPy dělá práci už od kapitoly 2 a odteď už nestojí za to na to upozorňovat.
Pás, s vzácnější otázkou
Odkaz na sekci: Pás, s vzácnější otázkouStejná továrna jako v kapitole 1, těžší otázka. Místo přijmout, nebo vyřadit zní otázka je tento díl vadný — a vady jsou vzácné, což dělá měřicí polovinu této kapitoly těžkou a modelovací polovinu klamně snadnou.
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 74Tři splity, ne dva. Důvod si zaslouží vlastní oddíl a jeden níže dostane; prozatím trénujte na prvním, laděte na druhém a na třetí se nedívejte.
Features jsou standardizované — odečtený průměr, vydělené směrodatnou odchylkou — za použití pouze statistik trénovací sady, z důvodu, který kapitola 1 ukázala na konvergenční mezi perceptronu: necentrovaná data dělají geometrii nepřátelskou. Z kterých řádků tento průměr smíte počítat, se později v této kapitole stane živá otázka.
Od verdiktu k pravděpodobnosti
Odkaz na sekci: Od verdiktu k pravděpodobnostiPerceptron vracel znaménko. Znaménko nedokáže rozlišit vyřadit od vyřadit, ale jen těsně, a přesně tento rozdíl továrna potřebuje, aby rozhodla, které díly má člověk znovu zkontrolovat jako první.
Držte se tedy doslova receptu z kapitoly 2. Zapište, co tvrdíte o tom, jak vzniká label, vezměte likelihood, vezměte logaritmus, změňte znaménko a máte ztrátu. Pro výsledek ano/ne je tím tvrzením Bernoulliho rozdělení: existuje pravděpodobnost , že díl je vadný, a
což je jen kompaktní způsob, jak napsat „ pokud , a pokud “. Vezměte z toho logaritmus a změňte znaménko; ztráta pro jeden příklad je
To je binární křížová entropie. Nebyla vybrána proto, že je pohodlná; je to záporný log-likelihood jediného rozdělení, které může mít hod mincí. Nic jiného k dispozici nebylo.
Stále chybí, odkud se bere . Model počítá vážený součet , což je reálné číslo v celém rozsahu přímky, zatímco pravděpodobnost musí ležet v . Funkce, která mezi nimi převádí, je logistický 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.9820Pravý sloupec čtěte jako ceník. Mít pravdu s 90% jistotou stojí 0,105. Odmítnout se rozhodnout stojí 0,693 — což je , cena pokrčení rameny. Mýlit se sebejistě stojí 4,6, čtyřiačtyřicetkrát víc, a cena roste bez omezení, jak si model je svou chybou jistější. Cross-entropy chyby jen nepočítá: účtuje si za aroganci.
Gradient je predikce minus pravda
Odkaz na sekci: Gradient je predikce minus pravdaKapitola 3 říkala: chcete-li cokoli trénovat, získejte derivaci ztráty podle každého parametru. Udělejte to pro jeden příklad. S a :
Zobrazit podrobnosti
Dva řádky, díky kterým se nepořádek vyruší. Sigmoid má neobvykle příjemnou derivaci, . A ztráta se derivuje na
Vynásobte obojí podle řetězového pravidla a se objeví jednou nahoře a jednou dole. Přesně se vyruší a přežije . Toto vyrušení není náhoda — děje se vždy, když je ztráta záporný log-likelihood rozdělení a výstupní funkce je ta, kterou toto rozdělení přirozeně používá. Tato dvojice má jméno — generalizovaný lineární model — a úhledný gradient je její otisk prstu.1
Aktualizace je tedy predikce minus pravda, krát vstup. Nic víc. Tady je celý trainer, tedy descent z kapitoly 3 s jedním změněným řádkem:
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 v sigmoid není kosmetika. Přímý výpočet přeteče pro velké záporné ; větev vybere tu algebraicky totožnou podobu, která udrží exponent záporný. To si floating-point krabice z kapitoly 2 vybírá svůj první dluh a o dva oddíly dál si vybere větší.
Proč ne čtvercovou chybu a proč odpověď souvisí s gradientem
Odkaz na sekci: Proč ne čtvercovou chybu a proč odpověď souvisí s gradientemStandardní vysvětlení, proč dát přednost cross-entropy před čtvercovou chybou, je likelihood argument výše: čtvercová chyba je to, co dostanete při předpokladu gaussovského šumu, labels gaussovské nejsou, tedy to nedělejte. Je to správně a nikoho to nepřesvědčí, protože můžete napsat přes sigmoid a bude se to trénovat.
Argument, který dopadne, je o gradientu. Položte čtvercovou chybu na sigmoid a řetězové pravidlo dá
To dodatečné je člen, který se předtím vyrušil. Teď se nevyruší a jde k nule vždy, když je model sebejistý — včetně situace, kdy se model sebejistě mýlí. Vyhodnoťte obojí pro několik skóre, pro příklad, jehož skutečný label je 1:
| skóre | cross-entropy | čtvercová chyba | poměr | |
|---|---|---|---|---|
| 0.000335 | 1 491 | |||
| 0.017986 | 28.3 | |||
| 0.119203 | 4.8 | |||
| 0.500000 | 2.0 | |||
| 0.880797 | 4.8 |
Při se model mýlí tak moc, jak je vůbec možné, a čtvercová chyba odpoví gradientem 1 491krát menším než cross-entropy. Čím horší chyba, tím méně se z ní model učí. Gradient cross-entropy se mezitím saturuje na : maximálně špatně znamená maximálně velký signál, a ne větší.
Spusťte závod. Dva tisíce vyvážených bodů, identické počáteční váhy zvolené tak, aby byly sebejistě špatně (), identická learning rate, liší se jen ztráta. Obě běhy jsou skórovány pomocí cross-entropy, aby byly sloupce srovnatelné.
| epocha | 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 je hotová v epoše 50. Čtvercová chyba je v epoše 100 pořád na 24 % accuracy — a od epochy 10 se nepohnula z 23 % — hůř než hádání, protože začala sebejistě špatně a gradient, který by ji zachránil, byl vynásoben 0,0007. Uteče z toho kolem epochy 500 a skončí na stejném místě. Poctivé shrnutí tedy zní: čtvercová chyba nad sigmoidem není nesprávná; je pomalá přesně tam, kde na rychlosti záleží nejvíc. Na modelu se dvěma parametry ztratíte 450 epoch. Na síti se sto vrstvami, kde je nějaká jednotka někde vždy sebejistě špatně, ztratíte celý trénovací běh.
Entropie, cross-entropy a KL na jedné stránce
Odkaz na sekci: Entropie, cross-entropy a KL na jedné stránceTři veličiny, které budete pořádně potřebovat v kapitole 8 pro perplexity a v kapitole 11 pro penalizaci, která drží fine-tuned politiku blízko její reference. Jsou snazší, než jakou mají pověst.2
Entropie je průměrný počet bitů, které musíte utratit, abyste sdělili losování z rozdělení, pokud pro něj použijete nejlepší možný kód:
Cross-entropy je to, co utratíte, když použijete kód postavený pro na data, která ve skutečnosti pocházejí z :
KL divergence je přebytek — plýtvání v bitech způsobené tím, že věříte , když pravda je :
Ověřte všechny tři na pásu:
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 bitsJsou tam vidět dvě věci. Za prvé, model, který prostě hlásí základní míru z tréninku, 1,69 %, dosáhne cross-entropy 0,1330 bitu, téměř přesně entropie testovacích labels — jak musí, protože má správné rozdělení a žádnou další informaci. Entropie je podlaha, kterou vám koupí neznalost jednotlivce. Za druhé, model, který pokrčí rameny a řekne 0,5, zaplatí přesně 1 bit, a rozdíl mezi nimi, 0,8671 bitu, je přesně KL divergence. není identita k zapamatování; je to účet, který můžete sledovat, jak se sčítá.
A spojení zpět k trénování: když je label jedna známá třída, „skutečné“ rozdělení je one-hot, jeho entropie je nula a cross-entropy se rovná KL divergenci. Minimalizace cross-entropy a přitahování rozdělení modelu k pravdě jsou tentýž čin.
Více než dvě odpovědi: softmax a posun, který nic nestojí
Odkaz na sekci: Více než dvě odpovědi: softmax a posun, který nic nestojíVadný neznamená jednu věc. Při vstřikování může díl vyjít jako nedostřik (málo materiálu), přetok (příliš mnoho, vytlačené z formy), nebo spálenina. Čtyři výsledky, tedy čtyři logits, a ty se musí stát čtyřmi pravděpodobnostmi se součtem jedna. To je softmax:
Má vlastnost, která vypadá jako náhoda a ve skutečnosti je celou implementací:
pro libovolnou konstantu , protože a se vyruší nahoře i dole. Význam mají jen rozdíly mezi logits. Absolutní úroveň není informace.
Naštěstí, protože absolutní úroveň je to, co rozbije počítač:
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 přeteče 64bitový float, součet se stane nekonečnem a nekonečno dělené nekonečnem je nan — žádná chyba, žádný pád, jen tichá díra tam, kde dřív byly tři pravděpodobnosti. Odečtení maximálního logit matematicky nezmění nic a numericky všechno, protože největší exponent se stane přesně . Je to trik logsumexp z kapitoly 2 v pracovním oděvu a dělá to každá seriózní implementace:
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, bGradient je znovu predikce minus pravda, nyní s one-hot. Binární případ byl celou dobu jen speciální případ.
Natrénováno na 3 000 dílech a otestováno na 1 000, se třemi měřeními pro každý (šířka, hmotnost, teplota taveniny), dosáhne 94,00 % accuracy. Tady je, co toto číslo skrývá:
| pravda ↓ / predikce → | ok | nedostřik | přetok | spálenina | recall |
|---|---|---|---|---|---|
| ok | 850 | 5 | 9 | 0 | 0.984 |
| nedostřik | 22 | 21 | 0 | 0 | 0.488 |
| přetok | 20 | 0 | 30 | 1 | 0.588 |
| spálenina | 3 | 0 | 0 | 39 | 0.929 |
| precision | 0.950 | 0.808 | 0.769 | 0.975 |
Model najde méně než polovinu nedostřiků. Accuracy to nevidí, protože 86 % dílů je v pořádku a správně trefit tyto díly stačí k utažení průměru. Macro F1 — průměr F1 skóre po třídách, který vzácné třídě dává stejnou váhu jako běžné — je 0,7983, oproti micro F1 0,9400, které je z definice totožné s accuracy. Kdykoli někdo nahlásí jedno číslo F1, zeptejte se které.
Tím končí modelování. Zbytek kapitoly je o číslech.
Tři modely, jedna accuracy
Odkaz na sekci: Tři modely, jedna accuracyVezměte natrénovaný binární model a vytvořte dvě varianty vynásobením každého logit konstantou: 0,35 pro váhavou verzi, 4 pro přehnaně sebejistou. Násobení kladným číslem nemůže změnit žádné znaménko, takže všechny tři modely predikují přesně stejný label pro všech 4 000 testovacích dílů. Accuracy je nerozliší. Cross-entropy s tím nemá žádný problém:
| model | accuracy | cross-entropy | průměrná ztráta při správné odpovědi | průměrná ztráta při chybě | nejhorší jednotlivá ztráta |
|---|---|---|---|---|---|
| váhavý (logits × 0.35) | 0.9830 | 0.1549 | 0.1369 | 1.1990 | 2.80 |
| jak byl natrénován | 0.9830 | 0.0564 | 0.0147 | 2.4689 | 7.82 |
| přehnaně sebejistý (logits × 4) | 0.9830 | 0.1563 | 0.0009 | 9.1427 | 27.63 |
Váhavý model platí malou daň za každý díl, včetně tisíců těch, které trefí správně. Přehnaně sebejistý je při správné odpovědi téměř zdarma a při chybě katastrofální — jeden díl v této testovací sadě ho sám o sobě stojí 27,63 nats. Ty dva skončí skoro na stejném součtu opačnými cestami a natrénovaný model, jehož pravděpodobnosti jsou kalibrované na data, sedí třikrát níž než oba.
To je nejostřejší způsob, jak vyjádřit rozdíl mezi ztrátou a metrikou. Ztráta je to, co optimalizujete: musí být diferencovatelná a vidí všechno, co model řekl, včetně toho, jak si byl jistý. Metrika je to, podle čeho jste souzeni: může to být schodová funkce, obchodní pravidlo, počet přehlédnutých vad. Nejsou to stejné objekty a ne vždy se shodnou — proto obojí definujete předem a nikdy nenecháte ztrátu zastupovat metriku jen proto, že se zrovna zobrazuje na obrazovce.
Hloupý baseline jde první
Odkaz na sekci: Hloupý baseline jde prvníPřed jakýmkoli modelem požadavek: jaké skóre má nejlínější možná odpověď? Na tomto pásu vždy řekněte v pořádku:
always-say-fine baseline: accuracy = 0.9815
confusion (tn, fp, fn, tp) = (3926, 0, 74, 0)98,15 %. Teď natrénovaný logistický model při výchozím threshold 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 %. Porazil baseline o 0,15 procentního bodu a každá zpráva, která skončí u accuracy, to označí za vítězství. Matice záměn říká, co se skutečně stalo:
| predikováno v pořádku | predikováno vadné | |
|---|---|---|
| skutečně v pořádku | 3,924 | 2 |
| skutečně vadné | 66 | 8 |
Našel 8 vadných dílů ze 74 a 66 pustil dál. Tři čísla pojmenovávají tři způsoby, jak tuto tabulku číst:
- Precision . Z dílů, které označil, kolik bylo skutečně vadných. To je cena zbytečných kontrol.
- Recall . Z vadných dílů, kolik jich zachytil. To je cena odeslání špatného dílu zákazníkovi.
- F1 , jejich harmonický průměr, který zůstává blízko menší z obou hodnot, a proto se nenechá lichotit jen jednou z nich.
Na čem záleží, určuje továrna, ne matematika: kontrola stojí pár sekund a odeslaná vada stojí svolávací akci, takže zde dominuje recall a 0,108 je selhání.
Problém ale není model. Problém je threshold, a threshold není součástí modelu — je to obchodní rozhodnutí aplikované následně na pravděpodobnost. Projeďte ho:
| 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 |
Čtěte sloupec accuracy směrem dolů. Celou dobu klesá — z 98,30 % na 65,93 % — zatímco model přechází od zachycení 8 vad k zachycení 71 ze 74. Každá užitečná věc, kterou tento model umí udělat, jeho accuracy zhoršuje. Tým optimalizující hlavní číslo by nasadil verzi, která nenajde nic.
Zobrazit podrobnosti
Vážení tříd nevytváří signál, posouvá provozní bod. Obvyklým prvním reflexem u nevyvážených tříd je navážit vzácnou třídu ve ztrátě. Když to uděláte s vahami 1, 10 a 60 na pozitivních příkladech:
| váha pozitivních | 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 a recall se posunou hodně. AUC — pravděpodobnost, že model seřadí náhodný vadný díl nad náhodný dobrý, která threshold úplně ignoruje — se posune o 0,0002, tedy o nic. Převážení posunulo stejný model po stejné křivce trade-offu. To je často přesně to, co chcete, a nikdy to není nová informace: pokud je řazení špatné, žádné váhovací schéma ho nezachrání.
Tři splity a leak, který právě najdete
Odkaz na sekci: Tři splity a leak, který právě najdeteProč tři splity a ne dva? Protože ve chvíli, kdy sadu příkladů použijete k tomu, abyste vybrali cokoli — threshold, learning rate, který ze šesti modelů nasadit — tato sada už byla použita k fittingu a její skóre přestane být nestranné.3 Změřeno na tomto pásu: sweep threshold na validační sadě vybere 0,196 a model pak na nedotčené testovací sadě dosáhne F1 = 0,4122. Kdyby sweep proběhl přímo na testovací sadě, nejlepší dosažitelná hodnota by tam byla 0,4186 — číslo, které nikdo nemá právo reportovat.
Rozdíl je zde malý, 0,006, protože jde o jeden hyperparametr jednou projetý proti 4 000 validačních příkladů. Roste s každým dalším rozhodnutím a každým zmenšením validační sady. Všimněte si také, že směr není v jednom běhu zaručený: zvolený threshold měl na validaci skóre 0,3902 a na testu 0,4122, takže validace ho tentokrát podhodnotila. Bias je systematický napříč mnoha rozhodnutími, ne viditelný v jednom.4
Teď cvičení. Log pásu dorazí s třetím sloupcem, station_seconds: jak dlouho každý díl strávil na kontrolním stanovišti. Jeho přidání je jednořádková změna v preprocessingu. Tady je, co udělá:
| model | accuracy | precision | recall | F1 | cross-entropy | AUC |
|---|---|---|---|---|---|---|
| šířka + hmotnost | 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 skočí z 10,8 % na 77,0 %. F1 se více než zečtyřnásobí. A všimněte si, co udělala accuracy: 98,30 % → 99,20 %, zisk devíti desetin bodu, tedy typ čísla, které se ve shrnujícím slidu zaokrouhlí na „tak jako tak asi 99 %“. Accuracy dříve neviděla selhání a teď nevidí podvod.
Než budete číst dál: model podvádí. Zjistěte jak.
Jak lovit leak, v pořadí, které ho najde nejrychleji.
-
Porovnejte train a test. Overfitting se projeví jako velký rozdíl. Zde: poctivý model 0,9838 train / 0,9830 test; model s leakem 0,9936 train / 0,9920 test. Oba rozdíly jsou pod 0,2 bodu. Leak nevypadá jako overfitting — uniklá feature je stejně dostupná v testu, takže model skvěle generalizuje do světa, který neexistuje.
-
Natrénujte jeden model na každou feature samostatně. Cokoli, co nese odpověď, se samo přihlásí:
feature samostatně accuracy recall F1 AUC šířka 0.9815 0.014 0.026 0.8691 hmotnost 0.9815 0.000 0.000 0.7914 station_seconds0.9850 0.405 0.500 0.9960 Jeden sloupec sám o sobě řadí vady s AUC 0,9960. Dvě měření pořízená posuvným měřidlem a vahou zvládnou 0,87 a 0,79. Tato asymetrie je alarm.
-
Zeptejte se, kdy bylo každé číslo zapsáno. Průměrná doba setrvání: 2,23 sekundy pro díly, které prošly, 15,56 sekundy pro díly, které neprošly. Samozřejmě. Díl setrvá na stanovišti protože ho inspektor stáhl z pásu — což se stane až poté, a jen proto, že se někdo rozhodl, že je vadný. Sloupec není měření dílu. Je to měření verdiktu.
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()) Zvýrazněný řádek je leak: doba setrvání vadného dílu je tažena z jiného rozdělení, protože ho člověk sundal z pásu. To je nejběžnější závažná chyba v aplikovaném machine learningu a má jméno: target leakage — informace v trénovacích features, která by nebyla dostupná ve chvíli, kdy je třeba predikci udělat.5 Nevyhodí žádnou výjimku. Vyrobí lepší číslo. Každá pobídka v projektu směřuje k tomu, abyste si ho nechali.
Obranou je jedna otázka položená u každého sloupce: existuje tato hodnota v okamžiku, kdy tuto predikci potřebuji? Na živém pásu je station_seconds neznámé až do chvíle po kontrole dílu — což je věc, kterou měl model nahradit.
Kolik testovacích příkladů potřebuji?
Odkaz na sekci: Kolik testovacích příkladů potřebuji?Představte si, že ohodnotíte model na 20 příkladech a 17 z nich trefí správně. Nahlásíte 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.6477Poctivé čtení 17/20 je někde mezi 64 % a 95 %. Skutečný 65% model vyprodukuje tento výsledek ve 4,4 % případů — jeden běh z třiadvaceti — a pokud jste vyzkoušeli několik prompts a nahlásili nejlepší, vyrobili jste si tento běh sami. Sedmnáct z dvaceti nedokáže odlišit 85% model od 65%.
Dva způsoby, jak dát intervalu míru, a oba patří do vaší výbavy:
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)Použijte Wilson6 pro obyčejnou míru úspěchu; chová se dobře při libovolném a nepotřebuje náhodnost. Všimněte si výše, že při je horní konec bootstrapu 1,0000 — resampling 20 bodů snadno vytáhne 20 správných, takže nedokáže reprezentovat interval užší než vlastní granularitu. Použijte bootstrap7 tam, kde neexistuje vzorec, což je většina zajímavých případů: F1, macro-průměry, BLEU, pass@1, skóre rubric-based judge. Na tomto pásu nese F1 vyladěného modelu 0,4122 bootstrap interval [0.3009, 0.5156] — a právě toto číslo má být ve zprávě, protože bodový odhad sám láká ke srovnání, které nemůže podpořit.
Ještě jedno měření, protože mění způsob, jakým byste měli porovnávat dva modely. Dva modely skórované na stejných 500 příkladech:
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)Jejich intervaly se překrývají a lidové pravidlo — překrývající se chybové úsečky znamenají žádný významný rozdíl — by označilo srovnání za neprůkazné. Není. Oba modely běžely na stejných příkladech, takže správnou veličinou je rozdíl po příkladech, jehož interval je [0.0260, 0.0680], pohodlně nad nulou. Neshodnou se jen na 31 z 500 položek a A vyhraje 27 z těchto neshod; sdílené příklady, snadné i těžké, se vyruší místo toho, aby přidávaly šum. Porovnávejte modely párově a dojdete ke stejnému závěru ze zlomku dat.
Kam to vede dál
Odkaz na sekci: Kam to vede dálTeď máte model, který vrací kalibrované pravděpodobnosti, ztrátu odvozenou z tvrzení o datech místo vybranou pro pohodlí, gradient, který je doslova predikce minus pravda, a — důležitěji — aparát, jak zjistit, jestli cokoli z toho funguje. Desetiřádkový Wilsonův interval výše se znovu používá doslova: nese varianty prompt v kapitole 15, retrieval tabulky v kapitole 19 a golden set v kapitole 29. Bootstrap je to, po čem sáhnete, když neexistuje vzorec.
Model má ale pořád jen jednu vrstvu. Kreslí čáru a kapitola 1 na čtyřech řádcích XOR dokázala, že čára nestačí. Oprava je stackování: první vrstva, která ohne prostor, druhá, která v ohnutém prostoru nakreslí čáru.
Tady úhledný gradient této kapitoly dochází. Všechno výše fungovalo, protože šlo zapsat ručně, jednou, pro model s jednou vrstvou mezi vstupem a ztrátou. Vložte doprostřed druhou vrstvu a otázka změní tvar: jaká je derivace ztráty podle váhy, která se výstupu vůbec nedotýká — takové, jejíž vliv dorazí jen přes jinou vrstvu, možná několika cestami najednou?
Tato derivace existuje. Počítat ji ručně je beznadějné pro cokoli většího než hračka a počítat ji po jednom parametru je beznadějné v jiném měřítku. Potřeba je postup, který získá každou derivaci v síti z jednoho zpětného průchodu přes tentýž graf, kterým právě prošel dopředný průchod.
To je kapitola 5 a je to motor, na kterém běží zbytek tohoto kurzu.
Zdroje a metoda
Odkaz na sekci: Zdroje a metodaVedle této kapitoly stojí za přečtení také: Bishop, Pattern Recognition and Machine Learning §1.2, §1.5, §1.6 a §4.3, který pokrývá pravděpodobnost, teorii rozhodování, teorii informace a lineární klasifikaci v pořadí, které tato kapitola sleduje; Murphy, Probabilistic Machine Learning: An Introduction, kapitoly 6 a 10; Prince, Understanding Deep Learning §5.4–5.7; a Saito a Rehmsmeier, The Precision-Recall Plot Is More Informative than the ROC Plot When Evaluating Binary Classifiers on Imbalanced Datasets (PLOS ONE, 2015) — proč by výše citované AUC nemělo být jediným threshold-free číslem, na které se díváte, když je vadných 1,7 % dílů.
Reference
Odkaz na sekci: Reference-
Ma, T. a Ng, A. CS229 Lecture Notes, Stanford University, kapitoly 2 a 3. Místo, kde vyrušení, které vytvoří , přestane vypadat jako štěstí: vyberte rozdělení z exponenciální rodiny, které odpovídá vašemu výstupu, použijte jeho kanonický link a gradient je vždy predikce minus pravda. ↩
-
Olah, C. Visual Information Theory (2015),
colah.github.io/posts/2015-09-Visual-Information. Nejjasnější dostupný výklad entropie, cross-entropy a KL divergence jako nákladů v bitech, ne jako vzorců. ↩ -
Abu-Mostafa, Y. S., Magdon-Ismail, M. a Lin, H.-T. Learning From Data (AMLBook, 2012), přednášky 13 a 17 kurzu Caltech. Přednáška 13 je validace; přednáška 17, o třech principech učení, je místo, kde je pojmenován data snooping. Dohromady jsou zdrojem disciplíny v této kapitole: každý pohled na datovou sadu je fitting rozhodnutí, ať jste spustili optimiser, nebo ne. ↩
-
James, G., Witten, D., Hastie, T. a Tibshirani, R. An Introduction to Statistical Learning, 2. vydání (Springer, 2021), kapitoly 2 a 5, pro rozklad bias–variance a resampling. Doprovodný svazek je místo, kde je selekční past řečena napřímo: Hastie, Tibshirani a Friedman, The Elements of Statistical Learning, 2. vydání, §7.10.2, The Wrong and Right Way to Do Cross-validation. ↩
-
Kaufman, S., Rosset, S., Perlich, C. a Stitelman, O. Leakage in Data Mining: Formulation, Detection, and Avoidance. ACM Transactions on Knowledge Discovery from Data 6(4), 2012. Formální zpracování selhání ukázaného výše, s případovými studiemi ze soutěží vyhraných modelem, který se naučil artefakt toho, jak byla data sestavena. ↩
-
Wilson, E. B. Probable Inference, the Law of Succession, and Statistical Inference. Journal of the American Statistical Association 22(158), s. 209–212 (1927). Score interval použitý výše v
wilson(), pořád správný výchozí interval pro podíl. Učebnicový interval je ten, kterému se vyhnout: dává nesmysly blízko 0 a 1 a při malém špatně pokrývá. ↩ -
Efron, B. Bootstrap Methods: Another Look at the Jackknife. The Annals of Statistics 7(1), s. 1–26 (1979). Myšlenka, která vám dovolí dát interval na jakoukoli statistiku, kterou umíte spočítat, včetně těch bez sampling theory. ↩