Klasifikasi, Cross-Entropy, dan Cara Tidak Menipu Diri Sendiri
Bangun classifier logistik, lalu lihat mengapa akurasi 98% bisa berarti model tidak menemukan apa pun.
Di halaman ini
Model yang menjawab bagian ini baik-baik saja untuk setiap komponen yang keluar dari ban berjalan benar 98,15 % dari waktu. Model itu juga tidak berguna: dari 74 komponen cacat di test set, tidak satu pun tertangkap.
Kedua kalimat itu menggambarkan model yang sama. Jarak di antara keduanya adalah bab ini.
Paruh pertama membangun classifier. Hampir tidak ada hal baru yang dibutuhkan: Bab 2 memberi resep untuk mengubah asumsi tentang bagaimana data dihasilkan menjadi loss function, dan Bab 3 memberi mesin untuk berjalan menurun pada loss apa pun yang diberikan resep itu. Terapkan keduanya ke pertanyaan ya/tidak dan logistic regression muncul, ditambah satu gagasan baru — sebuah logit — yang akan ditagih lagi di Bab 17.
Paruh kedua lebih sulit. Semua setelah titik ini dalam kursus dinilai oleh angka yang diukur seseorang, dan jika kamu tidak bisa membedakan peningkatan nyata dari artefak pengukuran, setiap bab berikutnya hanyalah hiasan. Jadi: confusion matrix, precision dan recall, tiga split, leakage, dan pertanyaan yang hampir tidak pernah dijawab jujur — berapa banyak contoh test yang sebenarnya aku butuhkan?
Aritmetika di sini berjalan di atas 20.000 baris, jadi semuanya divektorisasi — NumPy sudah melakukan pekerjaan sejak Bab 2, dan mulai sekarang hal itu tidak lagi layak disebutkan.
Ban berjalan, dengan pertanyaan yang lebih langka
Tautan ke bagian: Ban berjalan, dengan pertanyaan yang lebih langkaPabrik yang sama seperti Bab 1, pertanyaan lebih sulit. Alih-alih terima atau tolak, pertanyaannya adalah apakah komponen ini cacat — dan cacat itu jarang, yang membuat paruh pengukuran bab ini sulit dan paruh pemodelannya tampak mudah secara menipu.
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 74Tiga split, bukan dua. Alasannya layak mendapat bagian sendiri dan akan dibahas di bawah; untuk sekarang, latih pada yang pertama, tune pada yang kedua, dan jangan lihat yang ketiga.
Fitur distandardisasi — mean dikurangkan, dibagi dengan standard deviation — menggunakan statistik training saja, karena alasan yang Bab 1 tunjukkan dengan batas konvergensi perceptron: data yang tidak dicenter membuat geometri menjadi tidak ramah. Baris mana yang boleh kamu pakai untuk menghitung mean itu menjadi pertanyaan nyata nanti dalam bab ini.
Dari vonis ke probabilitas
Tautan ke bagian: Dari vonis ke probabilitasPerceptron mengembalikan tanda. Tanda tidak bisa membedakan tolak dari tolak, tapi nyaris saja, dan perbedaan itu persis yang dibutuhkan pabrik untuk memutuskan komponen mana yang harus diperiksa ulang manusia terlebih dahulu.
Jadi ikuti resep Bab 2 secara harfiah. Tulis klaimmu tentang bagaimana label dihasilkan, ambil likelihood, ambil log, negasikan, dan kamu punya loss. Untuk hasil ya/tidak, klaimnya adalah distribusi Bernoulli: ada probabilitas bahwa komponen itu cacat, dan
yang hanyalah cara ringkas untuk menulis “ jika , dan jika ”. Ambil log dari itu dan negasikan, dan loss untuk satu contoh adalah
Ini adalah binary cross-entropy. Ia tidak dipilih karena nyaman; ia adalah negative log-likelihood dari satu-satunya distribusi yang bisa dimiliki lemparan koin. Tidak ada pilihan lain.
Yang masih hilang adalah dari mana berasal. Model menghitung jumlah berbobot , yang merupakan bilangan real dan menjangkau seluruh garis bilangan, sementara probabilitas harus hidup di . Fungsi yang memindahkan di antara keduanya adalah 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.9820Baca kolom kanan sebagai daftar harga. Benar dengan confidence 90 % berbiaya 0,105. Menolak berkomitmen berbiaya 0,693 — yaitu , harga sebuah angkat bahu. Salah dengan confidence tinggi berbiaya 4,6, empat puluh empat kali lebih mahal, dan harganya naik tanpa batas saat model makin yakin pada kesalahan. Cross-entropy tidak sekadar menghitung error: ia menagih arogansi.
Gradient adalah prediksi dikurangi kebenaran
Tautan ke bagian: Gradient adalah prediksi dikurangi kebenaranBab 3 berkata: untuk melatih apa pun, dapatkan turunan loss terhadap setiap parameter. Lakukan untuk satu contoh. Dengan dan :
Tampilkan detail
Dua baris yang membuat kekacauan saling menghapus. Sigmoid punya turunan yang luar biasa nyaman, . Dan loss terdiferensiasi menjadi
Kalikan keduanya dengan chain rule dan muncul sekali di atas dan sekali di bawah. Keduanya tepat saling hapus, dan yang tersisa. Pembatalan itu bukan kebetulan — itulah yang terjadi setiap kali loss adalah negative log-likelihood dari sebuah distribusi dan fungsi output adalah fungsi yang secara alami dipakai distribusi itu. Pasangan itu punya nama — generalised linear model — dan gradient yang rapi adalah sidik jarinya.1
Jadi update-nya adalah prediksi dikurangi kebenaran, dikali input. Tidak ada yang lain. Berikut seluruh trainer-nya, yaitu descent Bab 3 dengan satu baris diubah:
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 di sigmoid bukan kosmetik. Menghitung secara langsung overflow untuk negatif besar; cabang itu memilih bentuk aljabar identik mana pun yang menjaga eksponen tetap negatif. Ini adalah kotak floating-point Bab 2 menagih utang pertamanya, dan ia akan menagih yang lebih besar dua bagian lagi dari sekarang.
Mengapa bukan squared error, dan mengapa jawabannya tentang gradient
Tautan ke bagian: Mengapa bukan squared error, dan mengapa jawabannya tentang gradientPenjelasan standar untuk lebih memilih cross-entropy daripada squared error adalah argumen likelihood di atas: squared error adalah yang kamu dapat dari asumsi Gaussian noise, label bukan Gaussian, jadi jangan. Itu benar dan tidak meyakinkan siapa pun, karena kamu bisa menulis di atas sigmoid dan ia akan berlatih.
Argumen yang benar-benar mengena adalah tentang gradient. Letakkan squared error di atas sigmoid dan chain rule memberi
Tambahan itulah yang tadi terhapus. Sekarang tidak, dan ia menuju nol kapan pun model confident — termasuk ketika model confident salah. Evaluasi keduanya pada beberapa skor, untuk contoh yang label benarnya 1:
| score | cross-entropy | squared error | rasio | |
|---|---|---|---|---|
| 0.000335 | 1,491 | |||
| 0.017986 | 28.3 | |||
| 0.119203 | 4.8 | |||
| 0.500000 | 2.0 | |||
| 0.880797 | 4.8 |
Pada model salah sebesar mungkin, dan squared error merespons dengan gradient 1.491 kali lebih kecil daripada cross-entropy. Makin buruk kesalahannya, makin sedikit model belajar darinya. Sementara itu, gradient cross-entropy jenuh di : salah maksimal menghasilkan sinyal sebesar maksimal, dan tidak lebih besar.
Jalankan perlombaannya. Dua ribu titik seimbang, bobot awal identik yang dipilih agar confident salah (), learning rate identik, hanya loss yang berbeda. Kedua run dinilai dengan cross-entropy agar kolomnya sebanding.
| epoch | cross-entropy loss | akurasi | squared-error loss | akurasi |
|---|---|---|---|---|
| 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 selesai pada epoch 50. Squared error masih di akurasi 24 % pada epoch 100 — dan belum bergerak dari 23 % pada epoch 10 — lebih buruk daripada menebak, karena ia mulai dalam keadaan confident salah dan gradient yang seharusnya menyelamatkannya telah dikalikan 0,0007. Ia lolos sekitar epoch 500 dan mendarat di tempat yang sama. Jadi ringkasan jujurnya adalah squared error di atas sigmoid tidak salah; ia lambat persis saat kecepatan paling penting. Pada model dua parameter kamu kehilangan 450 epoch. Pada jaringan dengan seratus layer, ketika selalu ada suatu unit di suatu tempat yang confident salah, kamu kehilangan seluruh training run.
Entropy, cross-entropy, dan KL, dalam satu halaman
Tautan ke bagian: Entropy, cross-entropy, dan KL, dalam satu halamanTiga kuantitas, dibutuhkan dengan benar di Bab 8 untuk perplexity dan di Bab 11 untuk penalti yang menjaga policy hasil fine-tuning tetap dekat dengan referensinya. Mereka lebih mudah daripada reputasinya.2
Entropy adalah rata-rata jumlah bit yang harus kamu keluarkan untuk mengomunikasikan satu sampel dari sebuah distribusi, jika kamu menggunakan kode terbaik yang mungkin untuknya:
Cross-entropy adalah yang kamu keluarkan ketika kamu memakai kode yang dibangun untuk pada data yang sebenarnya berasal dari :
KL divergence adalah kelebihannya — pemborosan, dalam bit, yang disebabkan oleh mempercayai ketika kebenarannya adalah :
Periksa ketiganya pada ban berjalan:
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 bitsDua hal terlihat di sana. Pertama, model yang sekadar melaporkan base rate training, 1,69 %, mencapai cross-entropy 0,1330 bit, hampir persis entropy label test — sebagaimana mestinya, karena ia punya distribusi yang benar dan tidak punya informasi lain. Entropy adalah lantai yang dibeli oleh ketidaktahuan tentang individu. Kedua, model yang mengangkat bahu dan berkata 0,5 membayar tepat 1 bit, dan selisih di antara keduanya, 0,8671 bit, persis KL divergence. bukan identitas untuk dihafal; ia adalah tagihan yang bisa kamu lihat dijumlahkan.
Dan koneksinya kembali ke training: ketika label adalah satu kelas yang diketahui, distribusi “benar” adalah one-hot, entropy-nya nol, dan cross-entropy sama dengan KL divergence. Meminimalkan cross-entropy dan menarik distribusi model ke arah kebenaran adalah tindakan yang sama.
Lebih dari dua jawaban: softmax, dan pergeseran yang tidak berbiaya
Tautan ke bagian: Lebih dari dua jawaban: softmax, dan pergeseran yang tidak berbiayaCacat bukan satu hal. Dalam moulding, sebuah komponen bisa keluar sebagai short shot (material tidak cukup), flash (terlalu banyak, terdorong keluar dari mould), atau burn. Empat hasil, jadi empat logits, dan semuanya harus menjadi empat probabilitas yang jumlahnya satu. Itulah softmax:
Ia punya sifat yang tampak seperti kebetulan dan sebenarnya merupakan seluruh implementasinya:
untuk konstanta apa pun , karena dan saling hapus di atas dan bawah. Hanya perbedaan antara logits yang bermakna. Level absolut bukan informasi.
Untungnya begitu, karena level absolutlah yang merusak 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 overflow pada float 64-bit, jumlahnya menjadi infinity, dan infinity dibagi infinity adalah nan — bukan error, bukan crash, hanya lubang senyap di tempat tiga probabilitas dulu berada. Mengurangkan logit maksimum tidak mengubah apa pun secara matematis dan mengubah segalanya secara numerik, karena eksponen terbesar menjadi tepat . Ini adalah trik logsumexp dari Bab 2 mengenakan pakaian kerja, dan setiap implementasi serius melakukannya:
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 lagi-lagi prediksi dikurangi kebenaran, sekarang dengan one-hot. Kasus biner ternyata memang kasus khusus sejak awal.
Dilatih pada 3.000 komponen dan diuji pada 1.000, dengan tiga pengukuran masing-masing (lebar, berat, temperatur leleh), ia mencapai akurasi 94,00 %. Inilah yang disembunyikan angka itu:
| kebenaran ↓ / prediksi → | 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 |
Model menemukan kurang dari separuh short shot. Akurasi tidak bisa melihat ini, karena 86 % komponen baik-baik saja dan membenarkan yang itu cukup untuk mengangkat rata-rata. Macro F1 — mean dari skor F1 per kelas, yang memberi bobot kelas langka sama dengan kelas umum — adalah 0,7983, dibandingkan micro F1 0,9400 yang menurut definisi identik dengan akurasi. Setiap kali seseorang melaporkan satu angka F1, tanyakan yang mana.
Itulah akhir pemodelan. Sisa bab ini tentang angka.
Tiga model, satu akurasi
Tautan ke bagian: Tiga model, satu akurasiAmbil model biner yang sudah dilatih dan buat dua varian dengan mengalikan setiap logit dengan konstanta: 0,35 untuk versi ragu-ragu, 4 untuk yang terlalu confident. Mengalikan dengan angka positif tidak bisa mengubah tanda apa pun, jadi ketiga model memprediksi label yang persis sama untuk semua 4.000 komponen test. Akurasi tidak bisa membedakannya. Cross-entropy tidak mengalami kesulitan sama sekali:
| model | akurasi | cross-entropy | mean loss saat benar | mean loss saat salah | loss tunggal terburuk |
|---|---|---|---|---|---|
| ragu-ragu (logits × 0.35) | 0.9830 | 0.1549 | 0.1369 | 1.1990 | 2.80 |
| sebagaimana dilatih | 0.9830 | 0.0564 | 0.0147 | 2.4689 | 7.82 |
| terlalu confident (logits × 4) | 0.9830 | 0.1563 | 0.0009 | 9.1427 | 27.63 |
Model ragu-ragu membayar pajak kecil pada setiap komponen, termasuk ribuan yang ia benarkan. Yang terlalu confident hampir gratis saat benar dan katastrofik saat salah — satu komponen dalam test set itu menelan biaya 27,63 nats sendirian. Keduanya mendarat pada total yang hampir sama lewat rute berlawanan, dan model terlatih, yang probabilitasnya terkalibrasi ke data, berada tiga kali lebih rendah dari keduanya.
Ini cara paling tajam untuk menyatakan perbedaan antara loss dan metric. Loss adalah yang kamu optimalkan: ia harus differentiable, dan ia melihat semua yang dikatakan model, termasuk seberapa yakin model itu. Metric adalah yang dipakai untuk menilai kamu: ia bisa berupa step function, aturan bisnis, hitungan cacat yang terlewat. Keduanya bukan objek yang sama dan tidak selalu sepakat — karena itu kamu mendefinisikan keduanya sebelum mulai, dan tidak pernah membiarkan loss menggantikan metric hanya karena ia kebetulan ada di layar.
Baseline bodoh berjalan lebih dulu
Tautan ke bagian: Baseline bodoh berjalan lebih duluSebelum model apa pun, syaratnya: jawaban paling malas mungkin mendapat skor berapa? Pada ban berjalan ini, selalu bilang baik-baik saja:
always-say-fine baseline: accuracy = 0.9815
confusion (tn, fp, fn, tp) = (3926, 0, 74, 0)98,15 %. Sekarang model logistik terlatih, pada threshold default 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 %. Ia mengalahkan baseline sebesar 0,15 poin persentase, dan laporan apa pun yang berhenti di akurasi akan menyebutnya kemenangan. Confusion matrix mengatakan apa yang sebenarnya terjadi:
| diprediksi baik | diprediksi cacat | |
|---|---|---|
| sebenarnya baik | 3,924 | 2 |
| sebenarnya cacat | 66 | 8 |
Ia menemukan 8 komponen cacat dari 74 dan membiarkan 66 lolos. Tiga angka menamai tiga cara membaca tabel itu:
- Precision . Dari komponen yang ditandai, berapa banyak yang benar-benar cacat. Ini adalah biaya inspeksi yang terbuang.
- Recall . Dari komponen yang cacat, berapa banyak yang tertangkap. Ini adalah biaya mengirim komponen buruk ke pelanggan.
- F1 , harmonic mean keduanya, yang tetap dekat dengan yang lebih kecil sehingga menolak dirayu oleh salah satunya saja.
Mana yang penting bergantung pada pabrik, bukan pada matematika: inspeksi memakan beberapa detik dan cacat yang terkirim menelan pemberitahuan recall, jadi di sini recall mendominasi dan 0,108 adalah kegagalan.
Tapi model bukan masalahnya. Threshold-lah masalahnya, dan threshold bukan bagian dari model — ia adalah keputusan bisnis yang diterapkan setelahnya pada probabilitas. Sweep threshold-nya:
| threshold | TP | FP | FN | akurasi | 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 |
Baca kolom akurasi ke bawah. Ia terus turun — dari 98,30 % ke 65,93 % — sementara model berubah dari menangkap 8 cacat menjadi menangkap 71 dari 74. Setiap hal berguna yang bisa dilakukan model ini membuat akurasinya memburuk. Tim yang mengoptimalkan angka headline akan mengirim versi yang tidak menemukan apa pun.
Tampilkan detail
Class weighting tidak menciptakan sinyal, ia memindahkan operating point. Refleks pertama yang umum pada kelas tidak seimbang adalah memberi bobot kelas langka dalam loss. Melakukannya, dengan bobot 1, 10, dan 60 pada positif:
| bobot pada positif | akurasi | 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 dan recall bergerak jauh. AUC — probabilitas bahwa model memberi ranking komponen cacat acak di atas komponen baik acak, yang sepenuhnya mengabaikan threshold — bergerak 0,0002, yaitu tidak ada artinya. Reweighting menggeser model yang sama di sepanjang kurva trade-off yang sama. Itu sering kali yang kamu inginkan, dan tidak pernah merupakan informasi baru: jika ranking-nya buruk, tidak ada skema weighting yang akan menyelamatkannya.
Tiga split, dan leak yang akan kamu temukan
Tautan ke bagian: Tiga split, dan leak yang akan kamu temukanMengapa tiga split dan bukan dua? Karena begitu kamu memakai sekumpulan contoh untuk memilih apa pun — threshold, learning rate, model mana dari enam yang akan dikirim — set itu sudah dipakai untuk fitting, dan skornya berhenti unbiased.3 Diukur pada ban berjalan ini: sweeping threshold pada validation set memilih 0,196, dan model kemudian mencetak F1 = 0,4122 pada test set yang belum disentuh. Jika sweep dijalankan langsung pada test set, skor terbaik yang bisa dicapai di sana adalah 0,4186 — angka yang tidak berhak dilaporkan siapa pun.
Gap-nya kecil di sini, 0,006, karena itu satu hyperparameter yang di-sweep sekali terhadap 4.000 contoh validation. Ia tumbuh dengan setiap keputusan tambahan dan setiap penyusutan validation set. Perhatikan juga bahwa arahnya tidak dijamin pada satu run: threshold yang dipilih mencetak 0,3902 pada validation dan 0,4122 pada test, jadi validation mengecilkannya kali ini. Bias itu sistematis di banyak keputusan, bukan terlihat dalam satu keputusan.4
Sekarang latihannya. Log ban berjalan datang dengan kolom ketiga, station_seconds: berapa lama setiap komponen berada di stasiun inspeksi. Menambahkannya adalah perubahan satu baris pada preprocessing. Inilah hasilnya:
| model | akurasi | precision | recall | F1 | cross-entropy | AUC |
|---|---|---|---|---|---|---|
| lebar + berat | 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 naik dari 10,8 % ke 77,0 %. F1 lebih dari empat kali lipat. Dan perhatikan apa yang dilakukan akurasi: 98,30 % → 99,20 %, kenaikan sembilan persepuluh poin, jenis angka yang dalam slide ringkasan akan dibulatkan menjadi “sekitar 99 % dengan cara apa pun”. Akurasi gagal melihat kegagalan sebelumnya dan sekarang gagal melihat kecurangan.
Sebelum lanjut membaca: model sedang curang. Cari tahu caranya.
Cara memburu leak, dalam urutan yang menemukannya paling cepat.
-
Bandingkan train dan test. Overfitting muncul sebagai gap besar. Di sini: model jujur 0,9838 train / 0,9830 test; model bocor 0,9936 train / 0,9920 test. Kedua gap di bawah 0,2 poin. Leak tidak terlihat seperti overfitting — fitur bocor sama-sama tersedia pada test time, jadi model generalises dengan indah ke dunia yang tidak ada.
-
Latih satu model per fitur, sendirian. Apa pun yang membawa jawabannya akan mengumumkan diri:
fitur saja akurasi recall F1 AUC lebar 0.9815 0.014 0.026 0.8691 berat 0.9815 0.000 0.000 0.7914 station_seconds0.9850 0.405 0.500 0.9960 Satu kolom, sendirian, memberi ranking cacat pada AUC 0,9960. Dua pengukuran yang diambil oleh kaliper dan timbangan hanya mencapai 0,87 dan 0,79. Asimetri itulah alarmnya.
-
Tanyakan kapan setiap angka ditulis. Mean dwell time: 2,23 detik untuk komponen yang lolos, 15,56 detik untuk komponen yang gagal. Tentu saja. Komponen berhenti lama di stasiun karena inspector menariknya dari ban berjalan — yang terjadi setelah, dan hanya karena, seseorang memutuskan komponen itu cacat. Kolom itu bukan pengukuran komponen. Ia adalah pengukuran vonis.
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()) Baris yang disorot adalah leak: dwell time komponen cacat diambil dari distribusi berbeda, karena manusia mengambilnya dari ban berjalan. Ini bug serius paling umum dalam machine learning terapan, dan ia punya nama: target leakage — informasi dalam fitur training yang tidak akan tersedia pada saat prediksi harus dibuat.5 Ia tidak melempar exception. Ia menghasilkan angka yang lebih baik. Setiap insentif dalam proyek mengarah untuk mempertahankannya.
Pertahanannya adalah satu pertanyaan, diajukan pada setiap kolom: pada saat persis aku membutuhkan prediksi ini, apakah nilai ini sudah ada? Pada ban berjalan live, station_seconds tidak diketahui sampai setelah komponen diperiksa — yaitu hal yang seharusnya digantikan model.
Berapa banyak contoh test yang aku butuhkan?
Tautan ke bagian: Berapa banyak contoh test yang aku butuhkan?Misalkan kamu menilai model pada 20 contoh dan ia benar 17 kali. Kamu melaporkan 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.6477Pembacaan jujur dari 17/20 adalah di suatu tempat antara 64 % dan 95 %. Model yang benar-benar 65 % menghasilkan hasil ini 4,4 % dari waktu — satu run dalam dua puluh tiga — dan jika kamu mencoba beberapa prompt dan melaporkan yang terbaik, kamu menciptakan run itu sendiri. Tujuh belas dari dua puluh tidak bisa membedakan model 85 % dari model 65 %.
Dua cara memberi interval pada sebuah rate, dan keduanya harus ada di toolkit-mu:
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)Gunakan Wilson6 untuk success rate biasa; ia tetap berperilaku baik pada apa pun dan tidak membutuhkan randomness. Perhatikan di atas bahwa pada ujung atas bootstrap adalah 1,0000 — resampling 20 titik bisa dengan mudah menarik 20 yang benar, jadi ia tidak bisa merepresentasikan interval yang lebih sempit daripada granularitasnya sendiri. Gunakan bootstrap7 ketika tidak ada formula, yaitu sebagian besar kasus menarik: F1, macro-average, BLEU, pass@1, skor dari judge berbasis rubrik. Pada ban berjalan ini, F1 model yang di-tune sebesar 0,4122 membawa interval bootstrap [0.3009, 0.5156] — itulah angka yang seharusnya muncul di laporan, karena point estimate saja mengundang perbandingan yang tidak bisa ia dukung.
Satu pengukuran lagi, karena ia mengubah cara kamu seharusnya membandingkan dua model. Dua model dinilai pada 500 contoh yang sama:
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)Interval mereka overlap, dan aturan rakyat — error bar yang overlap berarti tidak ada perbedaan signifikan — akan menyebut perbandingan ini inkonklusif. Tidak begitu. Kedua model berjalan pada contoh yang sama, jadi kuantitas yang benar adalah perbedaan per contoh, yang intervalnya [0.0260, 0.0680], nyaman di atas nol. Mereka hanya berbeda pendapat pada 31 dari 500 item, dan A menang pada 27 dari perbedaan itu; contoh bersama, yang mudah maupun sulit, saling membatalkan alih-alih menambah noise. Bandingkan model secara berpasangan, dan kamu mencapai kesimpulan yang sama dari sebagian kecil data.
Ke mana ini berlanjut
Tautan ke bagian: Ke mana ini berlanjutSekarang kamu punya model yang mengeluarkan probabilitas terkalibrasi, loss yang diturunkan dari klaim tentang data alih-alih dipilih karena nyaman, gradient yang secara harfiah prediksi dikurangi kebenaran, dan — yang lebih penting — mesin untuk mencari tahu apakah semua itu bekerja. Interval Wilson sepuluh baris di atas dipakai ulang verbatim: ia membawa varian prompt di Bab 15, tabel retrieval di Bab 19, dan golden set di Bab 29. Bootstrap adalah yang kamu ambil ketika tidak ada formula.
Tapi modelnya masih satu layer. Ia menggambar garis, dan Bab 1 membuktikan dengan empat baris XOR bahwa garis tidak cukup. Perbaikannya adalah menumpuk: layer pertama yang membengkokkan ruang, layer kedua yang menggambar garis di ruang yang sudah dibengkokkan.
Di sanalah gradient rapi bab ini habis. Semua di atas bekerja karena bisa ditulis dengan tangan, sekali, untuk model dengan satu layer di antara input dan loss. Letakkan layer kedua di tengah dan pertanyaannya berubah bentuk: apa turunan loss terhadap bobot yang sama sekali tidak menyentuh output — bobot yang pengaruhnya datang hanya melalui layer lain, mungkin sepanjang beberapa jalur sekaligus?
Turunan itu ada. Menghitungnya dengan tangan mustahil untuk apa pun yang lebih besar dari mainan, dan menghitungnya satu parameter setiap kali mustahil pada skala yang berbeda. Yang dibutuhkan adalah prosedur yang mendapatkan setiap turunan dalam jaringan dari satu backward pass di atas graph yang sama yang baru saja dilalui forward pass.
Itulah Bab 5, dan itulah mesin yang menjalankan sisa kursus ini.
Sumber dan metode
Tautan ke bagian: Sumber dan metodeJuga layak dibaca berdampingan dengan bab ini: Bishop, Pattern Recognition and Machine Learning §1.2, §1.5, §1.6 dan §4.3, yang membahas probability, decision theory, information theory, dan linear classification dalam urutan yang diikuti bab ini; Murphy, Probabilistic Machine Learning: An Introduction, bab 6 dan 10; Prince, Understanding Deep Learning §5.4–5.7; serta Saito dan Rehmsmeier, The Precision-Recall Plot Is More Informative than the ROC Plot When Evaluating Binary Classifiers on Imbalanced Datasets (PLOS ONE, 2015) — mengapa AUC yang dikutip di atas tidak boleh menjadi satu-satunya angka bebas-threshold yang kamu lihat ketika 1,7 % komponen cacat.
Referensi
Tautan ke bagian: Referensi-
Ma, T. dan Ng, A. CS229 Lecture Notes, Stanford University, bab 2 dan 3. Di situlah pembatalan yang menghasilkan berhenti terlihat seperti keberuntungan: pilih distribusi exponential-family yang cocok dengan output-mu, gunakan canonical link-nya, dan gradient selalu prediksi dikurangi kebenaran. ↩
-
Olah, C. Visual Information Theory (2015),
colah.github.io/posts/2015-09-Visual-Information. Penjelasan paling jernih yang tersedia tentang entropy, cross-entropy, dan KL divergence sebagai biaya dalam bit, bukan sebagai formula. ↩ -
Abu-Mostafa, Y. S., Magdon-Ismail, M. dan Lin, H.-T. Learning From Data (AMLBook, 2012), kuliah 13 dan 17 dari kursus Caltech. Kuliah 13 adalah validation; kuliah 17, tentang tiga prinsip learning, adalah tempat data snooping diberi nama. Keduanya bersama-sama menjadi sumber disiplin dalam bab ini: setiap pandangan ke sebuah data set adalah keputusan fitting, apakah kamu menjalankan optimiser atau tidak. ↩
-
James, G., Witten, D., Hastie, T. dan Tibshirani, R. An Introduction to Statistical Learning, edisi ke-2 (Springer, 2021), bab 2 dan 5, untuk bias–variance decomposition dan resampling. Volume pendampingnya adalah tempat jebakan seleksi dinyatakan langsung: Hastie, Tibshirani dan Friedman, The Elements of Statistical Learning, edisi ke-2, §7.10.2, The Wrong and Right Way to Do Cross-validation. ↩
-
Kaufman, S., Rosset, S., Perlich, C. dan Stitelman, O. Leakage in Data Mining: Formulation, Detection, and Avoidance. ACM Transactions on Knowledge Discovery from Data 6(4), 2012. Perlakuan formal atas kegagalan yang didemonstrasikan di atas, dengan studi kasus dari kompetisi yang dimenangkan oleh model yang telah mempelajari artefak dari cara data disusun. ↩
-
Wilson, E. B. Probable Inference, the Law of Succession, and Statistical Inference. Journal of the American Statistical Association 22(158), hlm. 209–212 (1927). Score interval yang digunakan dalam
wilson()di atas, masih default yang tepat untuk proporsi. Interval textbook adalah yang harus dihindari: ia memberi hasil tidak masuk akal dekat 0 dan 1, dan undercovers parah pada kecil. ↩ -
Efron, B. Bootstrap Methods: Another Look at the Jackknife. The Annals of Statistics 7(1), hlm. 1–26 (1979). Gagasan yang memungkinkan kamu memberi interval pada statistik apa pun yang bisa kamu hitung, termasuk yang tidak punya sampling theory. ↩