Klasyfikacja, entropia krzyżowa i jak nie oszukiwać samego siebie
Zbuduj klasyfikator logistyczny, a potem zobacz, dlaczego 98% trafności może oznaczać model, który nie znajduje nic.
Na tej stronie
Model, który o każdej części zjeżdżającej z taśmy odpowiada ta część jest w porządku, ma rację w 98,15 % przypadków. Jest też bezwartościowy: spośród 74 wadliwych części w zbiorze testowym nie wyłapuje żadnej.
Oba zdania opisują ten sam model. Odległość między nimi to ten rozdział.
Pierwsza połowa buduje klasyfikator. Potrzebuje prawie niczego nowego: rozdział 2 dał przepis na zamianę założenia o tym, jak powstają dane, w funkcję straty, a rozdział 3 dał mechanikę schodzenia w dół po dowolnej stracie, którą ten przepis poda. Zastosuj oba do pytania tak/nie, a otrzymasz regresję logistyczną plus jeden nowy pomysł — logit — za który zapłacimy ponownie w rozdziale 17.
Druga połowa jest trudniejsza. Wszystko od tego miejsca w kursie będzie oceniane liczbą, którą ktoś zmierzył, a jeśli nie umiesz odróżnić prawdziwej poprawy od artefaktu pomiaru, każdy kolejny rozdział jest dekoracją. A więc: macierz pomyłek, precision i recall, trzy podziały, leakage oraz pytanie, na które prawie nikt nie odpowiada uczciwie — ilu przykładów testowych naprawdę potrzebuję?
Arytmetyka działa tu na 20 000 wierszy, więc wszędzie jest zwektoryzowana — NumPy wykonuje pracę od rozdziału 2, a od teraz nie warto już za każdym razem tego zaznaczać.
Taśma, z rzadszym pytaniem
Link do sekcji: Taśma, z rzadszym pytaniemTa sama fabryka co w rozdziale 1, trudniejsze pytanie. Zamiast zaakceptować czy odrzucić, pytanie brzmi czy ta część jest wadliwa — a wady są rzadkie, co sprawia, że pomiarowa połowa tego rozdziału jest trudna, a modelująca połowa zwodniczo łatwa.
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 74Trzy podziały, nie dwa. Powód zasługuje na własną sekcję i dostanie ją niżej; na razie trenuj na pierwszym, dostrajaj na drugim i nie zaglądaj do trzeciego.
Cechy są standaryzowane — odejmujemy średnią i dzielimy przez odchylenie standardowe — używając wyłącznie statystyk treningowych, z powodu, który rozdział 1 pokazał na ograniczeniu zbieżności perceptronu: niewycentrowane dane robią geometrię wrogą. To, z których wierszy wolno ci policzyć tę średnią, stanie się żywym pytaniem później w tym rozdziale.
Od werdyktu do prawdopodobieństwa
Link do sekcji: Od werdyktu do prawdopodobieństwaPerceptron zwracał znak. Znak nie umie odróżnić odrzuć od odrzuć, ale ledwo, a właśnie ta różnica jest fabryce potrzebna, żeby zdecydować, które części człowiek powinien sprawdzić ponownie jako pierwsze.
Podążaj więc dosłownie za przepisem z rozdziału 2. Zapisz, co twierdzisz o sposobie powstawania etykiety, weź likelihood, weź log, zaneguj go i masz stratę. Dla wyniku tak/nie twierdzeniem jest rozkład Bernoulliego: istnieje prawdopodobieństwo , że część jest wadliwa, oraz
co jest po prostu zwartym zapisem „ jeśli , oraz jeśli ”. Weź logarytm z tego i go zaneguj, a strata dla jednego przykładu wynosi
To jest binarna entropia krzyżowa. Nie wybrano jej dlatego, że jest wygodna; to ujemny log-likelihood jedynego rozkładu, jaki może mieć rzut monetą. Niczego innego nie było do wyboru.
Wciąż brakuje tego, skąd bierze się . Model liczy ważoną sumę , która jest liczbą rzeczywistą i obejmuje całą prostą, a prawdopodobieństwo musi mieścić się w . Funkcją, która przenosi jedno w drugie, jest logistyczna 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.9820Czytaj prawą kolumnę jak cennik. Mieć rację z pewnością 90 % kosztuje 0,105. Odmowa zajęcia stanowiska kosztuje 0,693 — czyli , cena wzruszenia ramionami. Bycie pewnym siebie i w błędzie kosztuje 4,6, czterdzieści cztery razy więcej, a cena rośnie bez ograniczeń, gdy model staje się coraz bardziej pewny pomyłki. Entropia krzyżowa nie tylko liczy błędy: pobiera opłatę za arogancję.
Gradient to predykcja minus prawda
Link do sekcji: Gradient to predykcja minus prawdaRozdział 3 mówił: żeby cokolwiek trenować, uzyskaj pochodną straty względem każdego parametru. Zrób to dla jednego przykładu. Przy i :
Pokaż szczegóły
Dwie linie, dzięki którym bałagan się znosi. Sigmoid ma wyjątkowo przyjemną pochodną, . A strata różniczkuje się do
Pomnóż jedno przez drugie regułą łańcuchową, a pojawia się raz na górze i raz na dole. Znosi się dokładnie i zostaje . To zniesienie nie jest przypadkiem — dzieje się zawsze wtedy, gdy strata jest ujemnym log-likelihood rozkładu, a funkcja wyjściowa jest tą, której ten rozkład naturalnie używa. Ta para ma nazwę — uogólniony model liniowy — a schludny gradient jest jej odciskiem palca.1
Aktualizacja to więc predykcja minus prawda, razy wejście. Nic więcej. Oto cały trener, czyli zejście z rozdziału 3 z jedną zmienioną linią:
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 w sigmoid nie jest kosmetyką. Bezpośrednie liczenie powoduje overflow dla dużych ujemnych ; gałąź wybiera tę algebraicznie identyczną postać, która utrzymuje wykładnik ujemny. To pudełko zmiennoprzecinkowe z rozdziału 2 ściąga swój pierwszy dług, a większy ściągnie za dwie sekcje.
Dlaczego nie błąd kwadratowy i dlaczego odpowiedź dotyczy gradientu
Link do sekcji: Dlaczego nie błąd kwadratowy i dlaczego odpowiedź dotyczy gradientuStandardowe wyjaśnienie, czemu przedkłada się entropię krzyżową nad błąd kwadratowy, to argument z likelihood powyżej: błąd kwadratowy dostajesz, zakładając szum Gaussowski, etykiety nie są Gaussowskie, więc tego nie rób. To poprawne i nikogo nie przekonuje, bo możesz zapisać na sigmoid i model będzie się trenował.
Argument, który trafia, dotyczy gradientu. Połóż błąd kwadratowy na sigmoid, a reguła łańcuchowa daje
To dodatkowe jest tym, co wcześniej się zniosło. Teraz się nie znosi i dąży do zera zawsze, gdy model jest pewny — także wtedy, gdy model jest pewny siebie i błędny. Oceń oba dla kilku wyników, dla przykładu, którego prawdziwa etykieta to 1:
| wynik | entropia krzyżowa | błąd kwadratowy | stosunek | |
|---|---|---|---|---|
| 0.000335 | 1,491 | |||
| 0.017986 | 28.3 | |||
| 0.119203 | 4.8 | |||
| 0.500000 | 2.0 | |||
| 0.880797 | 4.8 |
Przy model myli się tak bardzo, jak tylko może, a błąd kwadratowy odpowiada gradientem 1 491 razy mniejszym niż gradient entropii krzyżowej. Im gorsza pomyłka, tym mniej model się z niej uczy. Gradient entropii krzyżowej tymczasem nasyca się przy : maksymalnie błędna predykcja daje maksymalnie duży sygnał i nie większy.
Uruchom wyścig. Dwa tysiące zbalansowanych punktów, identyczne wagi początkowe wybrane tak, by były pewne siebie i błędne (), identyczny learning rate, różni się tylko strata. Oba przebiegi oceniane są entropią krzyżową, żeby kolumny były porównywalne.
| epoka | strata entropii krzyżowej | trafność | strata błędu kwadratowego | trafność |
|---|---|---|---|---|
| 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 |
Entropia krzyżowa kończy pracę przed epoką 50. Błąd kwadratowy nadal ma 24 % trafności w epoce 100 — i nie ruszył się z 23 % w epoce 10 — gorzej niż zgadywanie, bo zaczął pewny siebie i błędny, a gradient, który miał go uratować, został pomnożony przez 0.0007. Ucieka dopiero około epoki 500 i ląduje w tym samym miejscu. Uczciwe podsumowanie jest więc takie: błąd kwadratowy na sigmoid nie jest niepoprawny; jest wolny dokładnie tam, gdzie szybkość ma największe znaczenie. W modelu z dwoma parametrami tracisz 450 epok. W sieci ze stu warstwami, gdzie jakaś jednostka gdzieś zawsze jest pewna siebie i błędna, tracisz cały trening.
Entropia, entropia krzyżowa i KL na jednej stronie
Link do sekcji: Entropia, entropia krzyżowa i KL na jednej stronieTrzy wielkości, potrzebne porządnie w rozdziale 8 dla perplexity i w rozdziale 11 dla kary, która trzyma fine-tuned politykę blisko jej referencji. Są łatwiejsze, niż wynika z ich reputacji.2
Entropia to średnia liczba bitów, które musisz wydać, żeby zakomunikować losowanie z rozkładu, jeśli używasz najlepszego możliwego kodu dla tego rozkładu:
Entropia krzyżowa to koszt, który ponosisz, gdy używasz kodu zbudowanego dla na danych, które w rzeczywistości pochodzą z :
Dywergencja KL to nadwyżka — marnotrawstwo w bitach — spowodowana wiarą w , gdy prawdą jest :
Sprawdź wszystkie trzy na taśmie:
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 bitsWidać tam dwie rzeczy. Po pierwsze, model, który po prostu raportuje bazową częstość z treningu, 1,69 %, osiąga entropię krzyżową 0.1330 bitu, prawie dokładnie entropię etykiet testowych — tak jak musi, skoro ma właściwy rozkład i żadnych innych informacji. Entropia jest podłogą, którą kupuje ci niewiedza o pojedynczym przykładzie. Po drugie, model, który wzrusza ramionami i mówi 0.5, płaci dokładnie 1 bit, a różnica między nimi, 0.8671 bitu, jest dokładnie dywergencją KL. nie jest tożsamością do zapamiętania; to rachunek, którego narastanie możesz obserwować.
I połączenie z treningiem: gdy etykieta jest jedną znaną klasą, „prawdziwy” rozkład jest one-hot, jego entropia wynosi zero, a entropia krzyżowa równa się dywergencji KL. Minimalizowanie entropii krzyżowej i przyciąganie rozkładu modelu ku prawdzie to ten sam akt.
Więcej niż dwie odpowiedzi: softmax i przesunięcie, które nic nie kosztuje
Link do sekcji: Więcej niż dwie odpowiedzi: softmax i przesunięcie, które nic nie kosztujeWadliwość nie jest jedną rzeczą. Przy formowaniu wtryskowym część może wyjść jako niedolew (za mało materiału), wypływka (za dużo, wyciśnięte z formy) albo przypalenie. Cztery wyniki, więc cztery logits, i muszą stać się czterema prawdopodobieństwami sumującymi się do jednego. To jest softmax:
Ma własność, która wygląda jak przypadek, a w istocie jest całą implementacją:
dla dowolnej stałej , bo oraz znoszą się na górze i na dole. Znaczenie mają tylko różnice między logits. Poziom bezwzględny nie jest informacją.
Na szczęście, bo poziom bezwzględny jest tym, co psuje komputer:
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 przepełnia 64-bitowy float, suma staje się nieskończonością, a nieskończoność podzielona przez nieskończoność to nan — nie błąd, nie awaria, tylko cicha dziura w miejscu, gdzie przed chwilą były trzy prawdopodobieństwa. Odjęcie maksymalnego logit nie zmienia nic matematycznie i wszystko numerycznie, bo największy wykładnik staje się dokładnie . To trik logsumexp z rozdziału 2 w roboczym ubraniu, i każda poważna implementacja go stosuje:
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 znów jest predykcją minus prawda, teraz z one-hot. Przypadek binarny od początku był przypadkiem szczególnym.
Wytrenowany na 3000 częściach i przetestowany na 1000, z trzema pomiarami każdej (szerokość, waga, temperatura stopu), osiąga 94,00 % trafności. Oto co ta liczba ukrywa:
| prawda ↓ / przewidziane → | ok | niedolew | wypływka | przypalenie | recall |
|---|---|---|---|---|---|
| ok | 850 | 5 | 9 | 0 | 0.984 |
| niedolew | 22 | 21 | 0 | 0 | 0.488 |
| wypływka | 20 | 0 | 30 | 1 | 0.588 |
| przypalenie | 3 | 0 | 0 | 39 | 0.929 |
| precision | 0.950 | 0.808 | 0.769 | 0.975 |
Model znajduje mniej niż połowę niedolewów. Trafność tego nie widzi, bo 86 % części jest w porządku, a poprawne rozpoznanie ich wystarcza, by podnieść średnią. Macro F1 — średnia wyników F1 dla poszczególnych klas, ważąca klasę rzadką tak samo jak częstą — wynosi 0.7983, wobec micro F1 równego 0.9400, które z definicji jest identyczne z trafnością. Gdy ktoś raportuje jedną liczbę F1, zapytaj którą.
To koniec modelowania. Reszta rozdziału dotyczy liczb.
Trzy modele, jedna trafność
Link do sekcji: Trzy modele, jedna trafnośćWeź wytrenowany model binarny i zrób dwa warianty, mnożąc każdy logit przez stałą: 0,35 dla wersji wahającej się, 4 dla wersji nadmiernie pewnej siebie. Mnożenie przez liczbę dodatnią nie może zmienić żadnego znaku, więc wszystkie trzy modele przewidują dokładnie tę samą etykietę dla wszystkich 4000 części testowych. Trafność nie umie ich odróżnić. Entropia krzyżowa nie ma z tym żadnego problemu:
| model | trafność | entropia krzyżowa | średnia strata przy poprawnej predykcji | średnia strata przy błędnej predykcji | najgorsza pojedyncza strata |
|---|---|---|---|---|---|
| wahający się (logits × 0.35) | 0.9830 | 0.1549 | 0.1369 | 1.1990 | 2.80 |
| wytrenowany | 0.9830 | 0.0564 | 0.0147 | 2.4689 | 7.82 |
| nadmiernie pewny siebie (logits × 4) | 0.9830 | 0.1563 | 0.0009 | 9.1427 | 27.63 |
Model wahający się płaci mały podatek od każdej części, także od tysięcy tych, które rozpoznaje poprawnie. Ten nadmiernie pewny siebie jest prawie darmowy, gdy ma rację, i katastrofalny, gdy jej nie ma — jedna część w tym zbiorze testowym kosztuje go sama 27.63 nata. Oba docierają do prawie tego samego wyniku całkowitego przeciwnymi drogami, a model wytrenowany, którego prawdopodobieństwa są skalibrowane do danych, siedzi trzy razy niżej od obu.
To najostrzejszy sposób sformułowania różnicy między stratą a metryką. Strata jest tym, co optymalizujesz: musi być różniczkowalna i widzi wszystko, co model powiedział, także to, jak bardzo był pewny. Metryka jest tym, według czego jesteś oceniany: może być funkcją skokową, regułą biznesową, liczbą przeoczonych wad. To nie ten sam obiekt i nie zawsze się zgadzają — dlatego definiujesz oba, zanim zaczniesz, i nigdy nie pozwalasz stracie zastąpić metryki tylko dlatego, że akurat jest na ekranie.
Głupi baseline idzie pierwszy
Link do sekcji: Głupi baseline idzie pierwszyPrzed jakimkolwiek modelem wymaganie: jaki wynik osiąga najbardziej leniwa możliwa odpowiedź? Na tej taśmie: zawsze mów, że jest dobrze:
always-say-fine baseline: accuracy = 0.9815
confusion (tn, fp, fn, tp) = (3926, 0, 74, 0)98,15 %. Teraz wytrenowany model logistyczny, przy domyślnym progu 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 %. Pokonał baseline o 0,15 punktu procentowego, a każdy raport, który zatrzyma się na trafności, nazwie to zwycięstwem. Macierz pomyłek mówi, co naprawdę się stało:
| przewidziane: dobra | przewidziane: wadliwa | |
|---|---|---|
| faktycznie dobra | 3,924 | 2 |
| faktycznie wadliwa | 66 | 8 |
Znalazł 8 wadliwych części z 74 i przepuścił 66. Trzy liczby nazywają trzy sposoby czytania tej tabeli:
- Precision . Spośród części, które oznaczył, ile naprawdę było wadliwych. To koszt zmarnowanych kontroli.
- Recall . Spośród wadliwych części, ile złapał. To koszt wysłania złej części do klienta.
- F1 , ich średnia harmoniczna, która pozostaje blisko mniejszej z dwóch wartości i dlatego nie daje się pochlebić tylko jednej z nich.
Co ma znaczenie, zależy od fabryki, nie od matematyki: kontrola kosztuje kilka sekund, a wysłana wada kosztuje akcję serwisową, więc tutaj dominuje recall i 0.108 jest porażką.
Ale problemem nie jest model. Problemem jest próg, a próg nie jest częścią modelu — to decyzja biznesowa zastosowana później do prawdopodobieństwa. Przeskanuj go:
| próg | TP | FP | FN | trafność | 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 |
Czytaj kolumnę trafności w dół. Spada przez cały czas — z 98,30 % do 65,93 % — podczas gdy model przechodzi od złapania 8 wad do złapania 71 z 74. Każda użyteczna rzecz, jaką ten model może zrobić, pogarsza jego trafność. Zespół optymalizujący główną liczbę wysłałby wersję, która nie znajduje niczego.
Pokaż szczegóły
Ważenie klas nie tworzy sygnału, tylko przesuwa punkt pracy. Typowym pierwszym odruchem przy niezbalansowanych klasach jest ważenie rzadkiej klasy w stracie. Po zrobieniu tego, z wagami 1, 10 i 60 dla pozytywów:
| waga pozytywów | trafność | 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 i recall przesuwają się daleko. AUC — prawdopodobieństwo, że model uszereguje losową wadliwą część wyżej niż losową dobrą, całkowicie ignorujące próg — zmienia się o 0.0002, czyli wcale. Ponowne ważenie przesunęło ten sam model wzdłuż tej samej krzywej kompromisu. Często właśnie tego chcesz, ale to nigdy nie jest nowa informacja: jeśli ranking jest zły, żaden schemat ważenia go nie uratuje.
Trzy podziały i leak, który zaraz znajdziesz
Link do sekcji: Trzy podziały i leak, który zaraz znajdzieszDlaczego trzy podziały, a nie dwa? Bo w chwili, gdy używasz zbioru przykładów do wybrania czegokolwiek — progu, learning rate, tego, który z sześciu modeli wysłać — ten zbiór został użyty do dopasowania, a jego wynik przestaje być nieobciążony.3 Zmierzone na tej taśmie: przemiatanie progu na zbiorze walidacyjnym wybiera 0.196, a model osiąga potem F1 = 0.4122 na nietkniętym zbiorze testowym. Gdyby przemiatanie uruchomiono bezpośrednio na zbiorze testowym, najlepszy osiągalny tam wynik wynosił 0.4186 — liczba, której nikt nie ma prawa raportować.
Różnica jest tu mała, 0.006, bo to jeden hyperparameter przemiotany raz wobec 4000 przykładów walidacyjnych. Rośnie z każdą dodatkową decyzją i każdym zmniejszeniem zbioru walidacyjnego. Zauważ też, że kierunek nie jest gwarantowany w pojedynczym przebiegu: wybrany próg uzyskał 0.3902 na walidacji i 0.4122 na teście, więc walidacja tym razem go zaniżyła. Obciążenie jest systematyczne na wielu decyzjach, nie widoczne w jednej.4
Teraz ćwiczenie. Log z taśmy przychodzi z trzecią kolumną, station_seconds: ile czasu każda część spędziła na stanowisku kontroli. Dodanie jej to jednolinijkowa zmiana w preprocessingu. Oto co robi:
| model | trafność | precision | recall | F1 | entropia krzyżowa | AUC |
|---|---|---|---|---|---|---|
| szerokość + waga | 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 rośnie z 10,8 % do 77,0 %. F1 rośnie ponad czterokrotnie. I zauważ, co zrobiła trafność: 98,30 % → 99,20 %, zysk dziewięciu dziesiątych punktu, czyli liczba, którą w slajdzie podsumowującym zaokrągla się do „około 99 % w obie strony”. Trafność wcześniej nie zobaczyła porażki, a teraz nie widzi oszustwa.
Zanim czytasz dalej: model oszukuje. Znajdź jak.
Jak polować na leak, w kolejności, która znajduje go najszybciej.
-
Porównaj trening i test. Overfitting objawia się dużą luką. Tutaj: uczciwy model 0.9838 trening / 0.9830 test; model z leak 0.9936 trening / 0.9920 test. Obie luki są poniżej 0,2 punktu. Leak nie wygląda jak overfitting — nieszczelna cecha jest równie dostępna w czasie testu, więc model pięknie generalizuje do świata, który nie istnieje.
-
Trenuj po jednym modelu na każdą cechę osobno. Wszystko, co niesie odpowiedź, samo się ogłosi:
sama cecha trafność recall F1 AUC szerokość 0.9815 0.014 0.026 0.8691 waga 0.9815 0.000 0.000 0.7914 station_seconds0.9850 0.405 0.500 0.9960 Jedna kolumna sama w sobie szereguje wady z AUC 0.9960. Dwa pomiary wykonane suwmiarką i wagą osiągają 0.87 i 0.79. Ta asymetria jest alarmem.
-
Zapytaj, kiedy każda liczba została zapisana. Średni czas przebywania: 2,23 sekundy dla części, które przeszły, 15,56 sekundy dla części, które nie przeszły. Oczywiście, że tak. Część przebywa na stanowisku bo kontroler zdjął ją z taśmy — co dzieje się po tym, i tylko dlatego, że ktoś uznał ją za wadliwą. Kolumna nie jest pomiarem części. Jest pomiarem werdyktu.
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()) Podświetlona linia jest leak: czas przebywania wadliwej części jest losowany z innego rozkładu, bo człowiek zdjął ją z taśmy. To najczęstszy poważny błąd w stosowanym uczeniu maszynowym i ma nazwę: target leakage — informacja w cechach treningowych, która nie byłaby dostępna w momencie, gdy trzeba wykonać predykcję.5 Nie rzuca wyjątku. Produkuje lepszą liczbę. Każdy bodziec w projekcie pcha w stronę jej zostawienia.
Obrona to jedno pytanie zadawane każdej kolumnie: czy w chwili, gdy potrzebuję tej predykcji, ta wartość już istnieje? Na żywej taśmie station_seconds jest nieznane aż do momentu po kontroli części — czyli po tym, co model miał zastąpić.
Ilu przykładów testowych potrzebuję?
Link do sekcji: Ilu przykładów testowych potrzebuję?Załóżmy, że oceniasz model na 20 przykładach i ma 17 poprawnych. Raportujesz 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.6477Uczciwe odczytanie 17/20 to gdzieś między 64 % a 95 %. Prawdziwie 65-procentowy model daje taki wynik w 4,4 % przypadków — jeden przebieg na dwadzieścia trzy — a jeśli wypróbowałeś garść prompts i zaraportowałeś najlepszy, sam wyprodukowałeś sobie taki przebieg. Siedemnaście z dwudziestu nie odróżnia modelu 85 % od modelu 65 %.
Dwa sposoby na przedział dla odsetka, i oba należą do twojego zestawu narzędzi:
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)Używaj Wilsona6 dla zwykłego odsetka sukcesów; zachowuje się dobrze przy dowolnym i nie potrzebuje losowości. Zauważ wyżej, że przy górny koniec bootstrap wynosi 1.0000 — ponowne próbkowanie 20 punktów może łatwo wylosować 20 poprawnych, więc nie potrafi reprezentować przedziału węższego niż własna ziarnistość. Używaj bootstrap7 tam, gdzie nie istnieje wzór, czyli w większości interesujących przypadków: F1, średnie macro, BLEU, pass@1, wynik sędziego opartego na rubryce. Na tej taśmie F1 dostrojonego modelu równe 0.4122 ma przedział bootstrap [0.3009, 0.5156] — i to liczba, która powinna pojawić się w raporcie, bo sam estymator punktowy zaprasza do porównania, którego nie potrafi wesprzeć.
Jeszcze jeden pomiar, bo zmienia to, jak należy porównywać dwa modele. Dwa modele ocenione na tych samych 500 przykładach:
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)Ich przedziały się nakładają, a ludowa reguła — nakładające się paski błędu oznaczają brak istotnej różnicy — uznałaby porównanie za nierozstrzygające. Nie jest. Oba modele działały na tych samych przykładach, więc właściwą wielkością jest różnica per przykład, której przedział to [0.0260, 0.0680], wygodnie powyżej zera. Nie zgadzają się tylko w 31 z 500 elementów, a A wygrywa 27 z tych niezgodności; wspólne przykłady, łatwe i trudne, znoszą się zamiast dodawać szum. Porównuj modele parami, a dojdziesz do tego samego wniosku na ułamku danych.
Dokąd to prowadzi dalej
Link do sekcji: Dokąd to prowadzi dalejMasz teraz model, który wypisuje skalibrowane prawdopodobieństwa, stratę wyprowadzoną z twierdzenia o danych zamiast wybraną dla wygody, gradient będący dosłownie predykcją minus prawda oraz — co ważniejsze — mechanikę sprawdzania, czy cokolwiek z tego działa. Dziesięciolinijkowy przedział Wilsona powyżej jest używany dosłownie ponownie: niesie warianty prompt w rozdziale 15, tabele retrieval w rozdziale 19 i golden set w rozdziale 29. Bootstrap to narzędzie, po które sięgasz, gdy nie istnieje wzór.
Ale model wciąż ma jedną warstwę. Rysuje linię, a rozdział 1 udowodnił na czterech wierszach XOR, że linia nie wystarcza. Naprawą jest układanie w stos: pierwsza warstwa, która zakrzywia przestrzeń, druga, która rysuje linię w zakrzywionej przestrzeni.
To tam schludny gradient z tego rozdziału się kończy. Wszystko powyżej działało, bo dało się zapisać ręcznie, raz, dla modelu z jedną warstwą między wejściem a stratą. Wstaw drugą warstwę pośrodku, a pytanie zmienia kształt: jaka jest pochodna straty względem wagi, która w ogóle nie dotyka wyjścia — takiej, której wpływ dociera tylko przez inną warstwę, być może kilkoma ścieżkami naraz?
Ta pochodna istnieje. Liczenie jej ręcznie jest beznadziejne dla czegokolwiek większego niż zabawka, a liczenie jej po jednym parametrze naraz jest beznadziejne w innej skali. Potrzebna jest procedura, która uzyskuje każdą pochodną w sieci z jednego backward pass po tym samym grafie, którym przed chwilą przeszedł forward pass.
To jest rozdział 5, i to silnik, na którym jedzie reszta tego kursu.
Źródła i metoda
Link do sekcji: Źródła i metodaWarto też czytać równolegle z tym rozdziałem: Bishop, Pattern Recognition and Machine Learning §1.2, §1.5, §1.6 i §4.3, który omawia prawdopodobieństwo, teorię decyzji, teorię informacji i klasyfikację liniową w kolejności, za którą idzie ten rozdział; Murphy, Probabilistic Machine Learning: An Introduction, rozdziały 6 i 10; Prince, Understanding Deep Learning §5.4–5.7; oraz Saito i Rehmsmeier, The Precision-Recall Plot Is More Informative than the ROC Plot When Evaluating Binary Classifiers on Imbalanced Datasets (PLOS ONE, 2015) — dlaczego AUC cytowane wyżej nie powinno być jedyną liczbą bezprogową, na którą patrzysz, gdy 1,7 % części jest wadliwych.
Przypisy
Link do sekcji: Przypisy-
Ma, T. i Ng, A. CS229 Lecture Notes, Stanford University, rozdziały 2 i 3. Tam zniesienie prowadzące do przestaje wyglądać jak szczęście: wybierz rozkład z rodziny wykładniczej, który pasuje do twojego wyjścia, użyj jego kanonicznego łącza, a gradient zawsze jest predykcją minus prawda. ↩
-
Olah, C. Visual Information Theory (2015),
colah.github.io/posts/2015-09-Visual-Information. Najjaśniejsze dostępne wyjaśnienie entropii, entropii krzyżowej i dywergencji KL jako kosztów w bitach, a nie jako wzorów. ↩ -
Abu-Mostafa, Y. S., Magdon-Ismail, M. i Lin, H.-T. Learning From Data (AMLBook, 2012), wykłady 13 i 17 kursu Caltech. Wykład 13 to walidacja; wykład 17, o trzech zasadach uczenia, to miejsce, gdzie nazwane zostaje data snooping. Razem są źródłem dyscypliny w tym rozdziale: każde spojrzenie na zbiór danych jest decyzją dopasowującą, niezależnie od tego, czy uruchomiłeś optimiser. ↩
-
James, G., Witten, D., Hastie, T. i Tibshirani, R. An Introduction to Statistical Learning, wyd. 2 (Springer, 2021), rozdziały 2 i 5, dla rozkładu bias–variance oraz resampling. Tom towarzyszący to miejsce, gdzie pułapka selekcji jest powiedziana wprost: Hastie, Tibshirani i Friedman, The Elements of Statistical Learning, wyd. 2, §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. Formalne ujęcie porażki zademonstrowanej powyżej, ze studiami przypadków z konkursów wygranych przez model, który nauczył się artefaktu sposobu złożenia danych. ↩
-
Wilson, E. B. Probable Inference, the Law of Succession, and Statistical Inference. Journal of the American Statistical Association 22(158), s. 209–212 (1927). Przedział score użyty w
wilson()powyżej, nadal właściwy domyślny wybór dla proporcji. Podręcznikowy przedział jest tym, którego należy unikać: daje nonsens blisko 0 i 1 oraz mocno zaniża pokrycie przy małych . ↩ -
Efron, B. Bootstrap Methods: Another Look at the Jackknife. The Annals of Statistics 7(1), s. 1–26 (1979). Pomysł, który pozwala nałożyć przedział na dowolną statystykę, jaką umiesz policzyć, także te bez teorii próbkowania. ↩