Classification, Cross-Entropy และวิธีไม่หลอกตัวเอง
สร้าง logistic classifier จาก loss และ descent แล้วดูว่าทำไม accuracy 98% อาจหมายถึง model ที่หาอะไรไม่เจอเลย
ในหน้านี้
model ที่ตอบว่า ชิ้นส่วนนี้ปกติ กับทุกชิ้นส่วนที่ออกจากสายพาน จะตอบถูก 98.15 % ของเวลา และมันก็ไร้ค่าเช่นกัน: จากชิ้นส่วนเสีย 74 ชิ้นใน test set มันจับได้ 0 ชิ้น
ทั้งสองประโยคอธิบาย model เดียวกัน ระยะห่างระหว่างสองประโยคนี้คือบทนี้
ครึ่งแรกสร้าง classifier แทบไม่ต้องมีอะไรใหม่: บทที่ 2 ให้สูตรในการแปลงสมมติฐานว่าข้อมูลถูกสร้างอย่างไรให้เป็น loss function และ บทที่ 3 ให้กลไกสำหรับเดินลงเขาบน loss ใดก็ตามที่สูตรนั้นส่งมาให้ เมื่อนำทั้งสองอย่างไปใช้กับคำถามแบบใช่/ไม่ใช่ ก็ได้ logistic regression ออกมา พร้อมแนวคิดใหม่หนึ่งอย่าง — logit — ซึ่งจะถูกคิดเงินซ้ำอีกครั้งใน บทที่ 17
ครึ่งหลังยากกว่า ทุกอย่างหลังจากจุดนี้ในคอร์สจะถูกตัดสินด้วยตัวเลขที่ใครบางคนวัดมา และถ้าคุณแยกไม่ออกว่าการปรับปรุงจริงกับ artefact จากการวัดต่างกันอย่างไร ทุกบทที่ตามมาก็เป็นแค่ของตกแต่ง ดังนั้น: confusion matrix, precision และ recall, การแบ่งข้อมูลสามชุด, leakage และคำถามที่แทบไม่มีใครตอบอย่างซื่อสัตย์ — จริง ๆ แล้วฉันต้องมี test examples กี่ตัว?
การคำนวณตรงนี้วิ่งบน 20,000 แถว ดังนั้นจึงทำแบบ vectorised ตลอด — NumPy ทำงานนี้มาตั้งแต่บทที่ 2 และจากนี้ไปมันไม่คุ้มที่จะพูดถึงทุกครั้งแล้ว
สายพาน กับคำถามที่พบได้น้อยกว่า
ลิงก์ไปยังส่วน: สายพาน กับคำถามที่พบได้น้อยกว่าโรงงานเดียวกับ บทที่ 1 แต่คำถามยากกว่า แทนที่จะถามว่า รับหรือปฏิเสธ คำถามคือ ชิ้นส่วนนี้มี defect หรือไม่ — และ defect พบได้น้อย ซึ่งทำให้ครึ่งของการวัดในบทนี้ยาก และครึ่งของการทำ model ดูง่ายอย่างหลอกตา
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 74สามชุด ไม่ใช่สองชุด เหตุผลควรมี section ของตัวเองและจะมีด้านล่าง ตอนนี้ให้ train บนชุดแรก tune บนชุดที่สอง และอย่ามองชุดที่สาม
features ถูก standardised — ลบค่าเฉลี่ย หารด้วยส่วนเบี่ยงเบนมาตรฐาน — โดยใช้ สถิติของ training เท่านั้น ด้วยเหตุผลที่บทที่ 1 แสดงไว้ผ่าน convergence bound ของ perceptron: ข้อมูลที่ไม่ถูกจัดศูนย์ทำให้ geometry เป็นศัตรู ส่วนคำถามว่าคุณได้รับอนุญาตให้คำนวณค่าเฉลี่ยจากแถวไหน จะกลายเป็นประเด็นจริงจังในบทนี้ภายหลัง
จากคำตัดสินเป็นความน่าจะเป็น
ลิงก์ไปยังส่วน: จากคำตัดสินเป็นความน่าจะเป็นperceptron คืนค่าเป็นเครื่องหมาย sign หนึ่งตัว sign แยกไม่ได้ระหว่าง ปฏิเสธ กับ ปฏิเสธ แต่เฉียดฉิว และความต่างนั้นคือสิ่งที่โรงงานต้องใช้เพื่อตัดสินว่าชิ้นส่วนไหนควรให้มนุษย์ตรวจซ้ำก่อน
ดังนั้นทำตามสูตรของบทที่ 2 แบบตรงตัว เขียนสิ่งที่คุณอ้างว่า label ถูกสร้างขึ้นอย่างไร หา likelihood เอา log แล้วใส่เครื่องหมายลบ คุณก็ได้ loss สำหรับผลลัพธ์แบบใช่/ไม่ใช่ ข้ออ้างคือการแจกแจง Bernoulli: มีความน่าจะเป็น ที่ชิ้นส่วนจะ defective และ
ซึ่งเป็นเพียงวิธีเขียนแบบกะทัดรัดว่า " ถ้า และ ถ้า " เอา log ของสิ่งนั้นแล้วใส่เครื่องหมายลบ loss สำหรับตัวอย่างหนึ่งตัวคือ
นี่คือ binary cross-entropy มันไม่ได้ถูกเลือกเพราะสะดวก แต่มันคือ negative log-likelihood ของการแจกแจงเดียวที่การโยนเหรียญจะมีได้ ไม่มีอย่างอื่นให้เลือก
สิ่งที่ยังขาดคือ มาจากไหน model คำนวณ weighted sum ซึ่งเป็นจำนวนจริงและมีค่าได้ทั้งเส้นจำนวน แต่ความน่าจะเป็นต้องอยู่ใน ฟังก์ชันที่ย้ายระหว่างสองโลกนี้คือ 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.9820อ่านคอลัมน์ขวาเหมือนรายการราคา ตอบถูกด้วยความมั่นใจ 90 % มีราคา 0.105 ปฏิเสธที่จะตัดสินใจมีราคา 0.693 — ซึ่งคือ ราคาของการยักไหล่ ตอบผิดอย่างมั่นใจมีราคา 4.6 มากกว่า 44 เท่า และราคาจะเพิ่มขึ้นอย่างไร้ขีดจำกัดเมื่อ model มั่นใจมากขึ้นกับความผิดพลาด cross-entropy ไม่ได้แค่นับ error: มันคิดเงินจากความโอหัง
gradient คือ prediction ลบ truth
ลิงก์ไปยังส่วน: gradient คือ prediction ลบ truthบทที่ 3 บอกว่า: เพื่อ train อะไรก็ตาม ให้หา derivative ของ loss เทียบกับ parameter แต่ละตัว ทำกับตัวอย่างหนึ่งตัว ด้วย และ :
แสดงรายละเอียด
สองบรรทัดที่ทำให้ความยุ่งเหยิงหักล้างกัน sigmoid มี derivative ที่ดีผิดปกติ คือ และ loss differentiate ได้เป็น
คูณทั้งสองด้วย chain rule แล้ว จะปรากฏครั้งหนึ่งด้านบนและอีกครั้งด้านล่าง มันหักล้างกันพอดี และสิ่งที่เหลือคือ การหักล้างนี้ไม่ใช่เรื่องบังเอิญ — มันคือสิ่งที่เกิดขึ้นเสมอเมื่อ loss เป็น negative log-likelihood ของการแจกแจงหนึ่ง และ output function เป็นฟังก์ชันที่การแจกแจงนั้นใช้โดยธรรมชาติ การจับคู่นี้มีชื่อว่า generalised linear model และ gradient ที่เรียบร้อยคือรอยนิ้วมือของมัน1
ดังนั้น update คือ prediction ลบ truth คูณ input แค่นั้น นี่คือ trainer ทั้งหมด ซึ่งเป็น descent ของบทที่ 3 ที่เปลี่ยนเพียงบรรทัดเดียว:
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 ใน sigmoid ไม่ใช่เพื่อความสวยงาม การคำนวณ ตรง ๆ จะ overflow สำหรับ ที่ติดลบมาก ๆ; branch จะเลือก form ที่เท่ากันทางพีชคณิตแต่ทำให้ exponent เป็นลบ นี่คือกล่อง floating-point ของบทที่ 2 เริ่มเก็บหนี้ก้อนแรก และมันจะเก็บก้อนใหญ่กว่าในอีกสอง section ข้างหน้า
ทำไมไม่ใช้ squared error และทำไมคำตอบจึงเกี่ยวกับ gradient
ลิงก์ไปยังส่วน: ทำไมไม่ใช้ squared error และทำไมคำตอบจึงเกี่ยวกับ gradientคำอธิบายมาตรฐานว่าทำไมควรใช้ cross-entropy แทน squared error คือเหตุผลเรื่อง likelihood ข้างต้น: squared error คือสิ่งที่ได้จากการสมมติว่า noise เป็น Gaussian, labels ไม่ใช่ Gaussian ดังนั้นอย่าทำ มันถูกต้องและไม่ทำให้ใครเชื่อ เพราะคุณสามารถเขียน ทับ sigmoid แล้วมันก็ train ได้
เหตุผลที่เข้าเป้าคือเรื่อง gradient ใส่ squared error ไว้บน sigmoid แล้ว chain rule ให้
ส่วนเกินนี้คือสิ่งที่ก่อนหน้านี้หักล้างไปแล้ว ตอนนี้มันไม่หักล้าง และมันเข้าใกล้ศูนย์เมื่อ model มั่นใจ — รวมถึงตอนที่ model ผิด อย่างมั่นใจด้วย ลองประเมินทั้งสองที่คะแนนต่าง ๆ สำหรับตัวอย่างที่ true label เป็น 1:
| 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 |
ที่ model ผิดสุดเท่าที่จะผิดได้ และ squared error ตอบสนองด้วย gradient ที่เล็กกว่า cross-entropy 1,491 เท่า ยิ่งความผิดพลาดแย่เท่าไร model ยิ่งเรียนรู้น้อยลงจากมัน ขณะเดียวกัน gradient ของ cross-entropy saturate ที่ : ผิดสูงสุดให้สัญญาณขนาดใหญ่สูงสุด และไม่ใหญ่ไปกว่านั้น
ให้มันแข่งกัน จุดข้อมูล balanced สองพันจุด น้ำหนักเริ่มต้นเหมือนกันและเลือกให้ผิดอย่างมั่นใจ (), learning rate เหมือนกัน ต่างกันแค่ loss ทั้งสอง run ถูกให้คะแนนด้วย cross-entropy เพื่อให้คอลัมน์เทียบกันได้
| 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 เสร็จตั้งแต่ epoch 50 squared error ยังอยู่ที่ accuracy 24 % ตอน epoch 100 — และยังไม่ขยับจาก 23 % ตอน epoch 10 — แย่กว่าการเดา เพราะมันเริ่มต้นแบบผิดอย่างมั่นใจ และ gradient ที่จะช่วยกู้มันถูกคูณด้วย 0.0007 มันหลุดออกมาได้ราว epoch 500 และลงเอยที่เดียวกัน ดังนั้นสรุปอย่างซื่อสัตย์คือ squared error บน sigmoid ไม่ได้ ผิด; แต่มัน ช้าพอดีในจุดที่ความเร็วสำคัญที่สุด บน model สอง parameter คุณเสีย 450 epochs บน network หนึ่งร้อย layer ที่ unit บางตัวที่ไหนสักแห่งผิดอย่างมั่นใจเสมอ คุณเสียทั้ง training run
Entropy, cross-entropy และ KL ในหน้าเดียว
ลิงก์ไปยังส่วน: Entropy, cross-entropy และ KL ในหน้าเดียวสามปริมาณที่ต้องใช้ให้ถูกใน บทที่ 8 สำหรับ perplexity และใน บทที่ 11 สำหรับ penalty ที่ทำให้ policy ที่ fine-tuned อยู่ใกล้ reference ของมัน มันง่ายกว่าชื่อเสียงของมัน2
Entropy คือจำนวน bits เฉลี่ยที่คุณต้องใช้เพื่อสื่อสารผลการสุ่มจากการแจกแจงหนึ่ง ถ้าคุณใช้ code ที่ดีที่สุดเท่าที่เป็นไปได้สำหรับมัน:
Cross-entropy คือสิ่งที่คุณใช้เมื่อคุณใช้ code ที่สร้างมาสำหรับ กับข้อมูลที่จริง ๆ แล้วมาจาก :
KL divergence คือส่วนเกิน — ความสูญเปล่าในหน่วย bits ที่เกิดจากการเชื่อ เมื่อความจริงคือ :
ตรวจทั้งสามบนสายพาน:
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 bitsมีสองอย่างที่เห็นได้ตรงนั้น อย่างแรก model ที่แค่รายงาน base rate ของ training คือ 1.69 % ได้ cross-entropy 0.1330 bits เกือบเท่ากับ entropy ของ test labels พอดี — ซึ่งต้องเป็นอย่างนั้น เพราะมันมีการแจกแจงที่ถูกต้องและไม่มีข้อมูลอื่น Entropy คือ floor ที่ความไม่รู้เกี่ยวกับตัวอย่างแต่ละตัวซื้อให้คุณ อย่างที่สอง model ที่ยักไหล่แล้วบอก 0.5 จ่ายพอดี 1 bit และช่องว่างระหว่างสองค่านั้น 0.8671 bits ก็คือ KL divergence อย่างแม่นยำ ไม่ใช่ identity ให้ท่องจำ; มันคือบิลที่คุณดูได้ว่าถูกบวกขึ้นอย่างไร
และความเชื่อมโยงกลับไปยัง training: เมื่อ label เป็น class ที่รู้แน่เพียง class เดียว การแจกแจง "true" เป็น one-hot, entropy ของมันเป็นศูนย์ และ cross-entropy เท่ากับ KL divergence การ minimise cross-entropy และดึงการแจกแจงของ model เข้าหาความจริงคือการกระทำเดียวกัน
มากกว่าสองคำตอบ: softmax และการ shift ที่ไม่เสียอะไร
ลิงก์ไปยังส่วน: มากกว่าสองคำตอบ: softmax และการ shift ที่ไม่เสียอะไรDefective ไม่ได้มีแบบเดียว ในการขึ้นรูป ชิ้นส่วนอาจออกมาเป็น short shot (วัสดุไม่พอ), flash (มากเกินไป ถูกบีบออกจากแม่พิมพ์) หรือ burn ผลลัพธ์สี่แบบ ดังนั้นมี logits สี่ตัว และมันต้องกลายเป็นความน่าจะเป็นสี่ค่าที่รวมกันได้หนึ่ง นั่นคือ softmax:
มันมีคุณสมบัติที่ดูเหมือนอุบัติเหตุ แต่จริง ๆ แล้วคือ implementation ทั้งหมด:
สำหรับค่าคงที่ใด ๆ เพราะ และ หักล้างกันทั้งบนและล่าง มีเพียง ความต่าง ระหว่าง logits เท่านั้นที่มีความหมาย ระดับ absolute ไม่ใช่ข้อมูล
โชคดีที่เป็นเช่นนั้น เพราะระดับ absolute คือสิ่งที่ทำให้คอมพิวเตอร์พัง:
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 float 64-bit, ผลรวมกลายเป็น infinity และ infinity หารด้วย infinity คือ nan — ไม่ใช่ error ไม่ใช่ crash แค่รูเงียบ ๆ ตรงที่เคยมีความน่าจะเป็นสามค่า การลบ logit สูงสุดไม่ได้เปลี่ยนอะไรทางคณิตศาสตร์ แต่เปลี่ยนทุกอย่างทางตัวเลข เพราะ exponent ที่ใหญ่ที่สุดกลายเป็น พอดี นี่คือ trick logsumexp ของบทที่ 2 ในชุดทำงาน และ implementation จริงจังทุกตัวทำแบบนี้:
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 ยังคงเป็น prediction ลบ truth ตอนนี้มี เป็น one-hot กรณี binary เป็นแค่กรณีพิเศษมาตลอด
เมื่อ train บนชิ้นส่วน 3,000 ชิ้น และ test บน 1,000 ชิ้น โดยมีการวัดสามอย่างต่อชิ้น (width, weight, melt temperature) มันได้ accuracy 94.00 % นี่คือสิ่งที่ตัวเลขนั้นซ่อนอยู่:
| 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 |
model หา short shot ได้ไม่ถึงครึ่ง Accuracy มองไม่เห็นสิ่งนี้ เพราะ 86 % ของชิ้นส่วนปกติ และการทายชิ้นส่วนเหล่านั้นให้ถูกก็พอแบกค่าเฉลี่ยแล้ว Macro F1 — ค่าเฉลี่ยของ F1 score ราย class ซึ่งให้น้ำหนัก class ที่หายากเท่ากับ class ที่พบบ่อย — คือ 0.7983 เทียบกับ micro F1 0.9400 ที่โดยนิยามแล้วเหมือนกับ accuracy ทุกครั้งที่มีคนรายงาน F1 ตัวเดียว ให้ถามว่าอันไหน
นี่คือส่วนสุดท้ายของการทำ model ส่วนที่เหลือของบทนี้เกี่ยวกับตัวเลข
สาม models หนึ่ง accuracy
ลิงก์ไปยังส่วน: สาม models หนึ่ง accuracyนำ binary model ที่ train แล้วมาสร้าง variant สองตัว โดยคูณทุก logit ด้วยค่าคงที่: 0.35 สำหรับเวอร์ชันลังเล และ 4 สำหรับเวอร์ชันมั่นใจเกินไป การคูณด้วยจำนวนบวกไม่สามารถเปลี่ยน sign ใด ๆ ได้ ดังนั้นทั้งสาม models predict label เดียวกันเป๊ะสำหรับ test parts ทั้ง 4,000 ชิ้น Accuracy แยกไม่ออก แต่ cross-entropy ไม่มีปัญหาเลย:
| model | accuracy | cross-entropy | mean loss when right | mean loss when wrong | worst single loss |
|---|---|---|---|---|---|
| hesitant (logits × 0.35) | 0.9830 | 0.1549 | 0.1369 | 1.1990 | 2.80 |
| as trained | 0.9830 | 0.0564 | 0.0147 | 2.4689 | 7.82 |
| overconfident (logits × 4) | 0.9830 | 0.1563 | 0.0009 | 9.1427 | 27.63 |
model ที่ลังเลจ่ายภาษีเล็ก ๆ กับทุกชิ้นส่วน รวมถึงหลายพันชิ้นที่มันทำถูก ตัวที่มั่นใจเกินไปแทบไม่เสียอะไรเมื่อถูก และพังยับเมื่อผิด — ชิ้นส่วนหนึ่งใน test set นั้นทำให้มันเสีย 27.63 nats ด้วยตัวเอง ทั้งสองจบที่ผลรวมเกือบเท่ากันด้วยเส้นทางตรงข้ามกัน และ model ที่ train แล้ว ซึ่ง probabilities ถูก calibrate ให้เข้ากับข้อมูล อยู่ต่ำกว่าทั้งคู่สามเท่า
นี่คือวิธีที่คมที่สุดในการบอกความต่างระหว่าง loss กับ metric loss คือสิ่งที่คุณ optimise: มันต้อง differentiable และมองเห็นทุกอย่างที่ model พูด รวมถึงมันมั่นใจแค่ไหน metric คือสิ่งที่คุณถูกตัดสินด้วย: มันอาจเป็น step function, business rule, หรือจำนวน defect ที่พลาดก็ได้ ทั้งสองไม่ใช่วัตถุเดียวกันและไม่ได้เห็นตรงกันเสมอ — นั่นคือเหตุผลที่คุณต้องนิยามทั้งคู่ก่อนเริ่ม และอย่าปล่อยให้ loss แทน metric เพียงเพราะมันบังเอิญอยู่บนหน้าจอ
baseline โง่ ๆ ต้องมาก่อน
ลิงก์ไปยังส่วน: baseline โง่ ๆ ต้องมาก่อนก่อนมี model ใด ๆ ต้องมีข้อกำหนด: คำตอบที่ขี้เกียจที่สุดเท่าที่เป็นไปได้ได้คะแนนเท่าไร? บนสายพานนี้คือ บอกว่าปกติเสมอ:
always-say-fine baseline: accuracy = 0.9815
confusion (tn, fp, fn, tp) = (3926, 0, 74, 0)98.15 % ตอนนี้มาดู logistic model ที่ train แล้ว ที่ 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 % มันชนะ baseline ด้วย 0.15 percentage point และรายงานใดก็ตามที่หยุดอยู่ที่ accuracy จะเรียกสิ่งนี้ว่าชัยชนะ confusion matrix บอกว่าเกิดอะไรขึ้นจริง:
| predicted fine | predicted defective | |
|---|---|---|
| actually fine | 3,924 | 2 |
| actually defective | 66 | 8 |
มันพบชิ้นส่วน defective 8 ชิ้นจาก 74 และปล่อยผ่านไป 66 ชิ้น ตัวเลขสามตัวตั้งชื่อวิธีอ่านตารางนี้สามแบบ:
- Precision จากชิ้นส่วนที่มัน flag มีสักกี่ชิ้นที่ defective จริง นี่คือต้นทุนของการตรวจซ้ำที่สูญเปล่า
- Recall จากชิ้นส่วน defective มีกี่ชิ้นที่มันจับได้ นี่คือต้นทุนของการส่งชิ้นส่วนเสียให้ลูกค้า
- F1 ค่า harmonic mean ของทั้งสอง ซึ่งจะอยู่ใกล้ค่าที่เล็กกว่า และจึงปฏิเสธที่จะถูกยกยอด้วยค่าใดค่าหนึ่งเพียงลำพัง
อะไรสำคัญขึ้นอยู่กับโรงงาน ไม่ใช่คณิตศาสตร์: การตรวจใช้เวลาไม่กี่วินาที แต่ defect ที่ถูกส่งออกไปมีต้นทุนเป็นประกาศเรียกคืนสินค้า ดังนั้นที่นี่ recall จึงสำคัญกว่า และ 0.108 คือความล้มเหลว
แต่ปัญหาไม่ใช่ model threshold ต่างหาก และ threshold ไม่ใช่ส่วนหนึ่งของ model — มันคือการตัดสินใจทางธุรกิจที่นำไปใช้ภายหลังกับความน่าจะเป็น ลอง sweep มัน:
| 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 |
อ่านคอลัมน์ accuracy ลงมา มันลดลงตลอดทาง — จาก 98.30 % เป็น 65.93 % — ขณะที่ model เปลี่ยนจากการจับ defect ได้ 8 ชิ้น ไปเป็น 71 จาก 74 ชิ้น ทุกสิ่งที่มีประโยชน์ที่ model นี้ทำได้ทำให้ accuracy แย่ลง ทีมที่ optimise ตัวเลขพาดหัวจะส่งเวอร์ชันที่หาอะไรไม่เจอออกไป
แสดงรายละเอียด
Class weighting ไม่ได้สร้าง signal มันย้าย operating point reflex แรกตามปกติเมื่อ classes ไม่สมดุลคือให้น้ำหนัก class ที่หายากใน loss ทำแบบนั้นด้วย weights 1, 10 และ 60 บน positives:
| weight on 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 และ recall ขยับไปไกลมาก AUC — ความน่าจะเป็นที่ model จัดอันดับชิ้นส่วน defective แบบสุ่มไว้เหนือชิ้นส่วนดีแบบสุ่ม ซึ่งไม่สน threshold เลย — ขยับ 0.0002 ซึ่งเท่ากับไม่ขยับ Reweighting เลื่อน model เดิมไปตาม trade-off curve เดิม นั่นมักเป็นสิ่งที่คุณต้องการ และมันไม่ใช่ข้อมูลใหม่เลย: ถ้า ranking แย่ ไม่มี weighting scheme ใดช่วยได้
สาม splits และ leak ที่คุณกำลังจะเจอ
ลิงก์ไปยังส่วน: สาม splits และ leak ที่คุณกำลังจะเจอทำไมต้องสาม splits ไม่ใช่สอง? เพราะทันทีที่คุณใช้ชุดตัวอย่างหนึ่งเพื่อ เลือก อะไรก็ตาม — threshold, learning rate, model ใดในหกตัวที่จะ ship — ชุดนั้นถูกใช้สำหรับ fitting แล้ว และคะแนนของมันก็หยุดเป็น unbiased3 วัดบนสายพานนี้: การ sweep threshold บน validation set เลือก 0.196 แล้ว model ได้คะแนน F1 = 0.4122 บน test set ที่ไม่ถูกแตะต้อง ถ้า sweep นั้นถูกรันบน test set โดยตรง ค่าดีที่สุดที่ได้ตรงนั้นคือ 0.4186 — ตัวเลขที่ไม่มีใครมีสิทธิ์รายงาน
ช่องว่างเล็กตรงนี้คือ 0.006 เพราะเป็น hyperparameter หนึ่งตัวที่ sweep ครั้งเดียวกับ validation examples 4,000 ตัว มันโตขึ้นกับทุกการตัดสินใจที่เพิ่มขึ้น และทุกครั้งที่ validation set เล็กลง สังเกตด้วยว่าทิศทางไม่ได้การันตีใน run เดียว: threshold ที่เลือกได้ 0.3902 บน validation และ 0.4122 บน test ดังนั้น validation ประเมินต่ำไป ครั้งนี้ bias เป็นระบบเมื่อมีการตัดสินใจจำนวนมาก ไม่ใช่สิ่งที่เห็นได้ในครั้งเดียว4
ตอนนี้เป็นแบบฝึกหัด log ของสายพานมาพร้อมคอลัมน์ที่สาม station_seconds: เวลาที่ชิ้นส่วนแต่ละชิ้นใช้ที่สถานีตรวจสอบ การเพิ่มมันเป็นการเปลี่ยน preprocessing เพียงบรรทัดเดียว นี่คือสิ่งที่มันทำ:
| 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 เพิ่มจาก 10.8 % เป็น 77.0 % F1 มากกว่าสี่เท่า และดูว่า accuracy ทำอะไร: 98.30 % → 99.20 % เพิ่มขึ้นเก้าส่วนสิบของจุด ซึ่งเป็นตัวเลขประเภทที่ถูกปัดเป็น "ประมาณ 99 % ทั้งคู่" ในสไลด์สรุป Accuracy ไม่เห็นความล้มเหลวก่อนหน้านี้ และตอนนี้ก็ไม่เห็นการโกง
ก่อนอ่านต่อ: model กำลังโกง หาว่าโกงอย่างไร
วิธีล่า leak ตามลำดับที่เจอเร็วที่สุด
-
เปรียบเทียบ train กับ test Overfitting จะแสดงเป็นช่องว่างขนาดใหญ่ ที่นี่: honest model 0.9838 train / 0.9830 test; leaky model 0.9936 train / 0.9920 test ช่องว่างทั้งคู่ต่ำกว่า 0.2 points leak ไม่ได้ดูเหมือน overfitting — leaky feature มีให้ใช้ที่ test time เช่นกัน ดังนั้น model จึง generalise ได้อย่างงดงามไปยังโลกที่ไม่มีอยู่จริง
-
Train model หนึ่งตัวต่อ feature โดยใช้ feature นั้นตัวเดียว อะไรก็ตามที่พกคำตอบมาจะประกาศตัวเอง:
feature alone 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 คอลัมน์เดียว เพียงลำพัง จัดอันดับ defects ได้ AUC 0.9960 การวัดสองอย่างที่มาจาก caliper และตาชั่งทำได้ 0.87 และ 0.79 ความไม่สมมาตรนี้คือสัญญาณเตือน
-
ถามว่าตัวเลขแต่ละตัวถูกบันทึกเมื่อไร Mean dwell time: 2.23 วินาทีสำหรับชิ้นส่วนที่ผ่าน, 15.56 วินาทีสำหรับชิ้นส่วนที่ตก แน่นอนอยู่แล้ว ชิ้นส่วนค้างอยู่ที่สถานี เพราะ inspector หยิบมันออกจากสายพาน — ซึ่งเกิดขึ้นหลังจาก และเกิดขึ้นเพราะ มีใครบางคนตัดสินแล้วว่ามัน defective คอลัมน์นี้ไม่ใช่การวัดชิ้นส่วน มันคือการวัดคำตัดสิน
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()) บรรทัดที่ highlight คือ leak: dwell time ของชิ้นส่วน defective ถูกสุ่มจากการแจกแจงที่ต่างออกไป เพราะมนุษย์หยิบมันออกจากสายพาน นี่คือ bug ร้ายแรงที่พบบ่อยที่สุดใน applied machine learning และมันมีชื่อว่า target leakage — ข้อมูลใน training features ที่จะไม่มีให้ใช้ ณ เวลาที่ต้องทำ prediction จริง5 มันไม่ throw exception มันสร้างตัวเลขที่ดีขึ้น ทุกแรงจูงใจในโปรเจกต์ผลักไปทางการเก็บมันไว้
การป้องกันคือคำถามเดียวที่ถามกับทุกคอลัมน์: ในวินาทีที่ฉันต้องการ prediction นี้ ค่านี้มีอยู่แล้วหรือยัง? บนสายพานสด station_seconds ยังไม่รู้จนกว่าชิ้นส่วนจะถูกตรวจ — ซึ่งเป็นสิ่งที่ model ควรมาแทนที่
ฉันต้องมี test examples กี่ตัว?
ลิงก์ไปยังส่วน: ฉันต้องมี test examples กี่ตัว?สมมติคุณให้คะแนน model บน 20 examples และมันตอบถูก 17 คุณรายงาน 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.6477การอ่าน 17/20 อย่างซื่อสัตย์คือ อยู่ที่ไหนสักแห่งระหว่าง 64 % ถึง 95 % model ที่จริง ๆ แล้ว 65 % ให้ผลลัพธ์นี้ 4.4 % ของเวลา — หนึ่ง run ในยี่สิบสาม — และถ้าคุณลอง prompts สักกำมือแล้วรายงานตัวที่ดีที่สุด คุณก็ผลิต run นั้นขึ้นมาเอง สิบเจ็ดจากยี่สิบแยก model 85 % ออกจาก model 65 % ไม่ได้
มีสองวิธีในการใส่ interval ให้ rate และทั้งสองควรอยู่ใน 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)ใช้ Wilson6 สำหรับ success rate ธรรมดา; มันประพฤติตัวดีที่ ใด ๆ และไม่ต้องใช้ randomness สังเกตด้านบนว่าที่ ปลายบนของ bootstrap เป็น 1.0000 — การ resample 20 points สามารถดึงตัวที่ถูกทั้ง 20 ได้ง่าย ๆ ดังนั้นมันไม่สามารถแทน interval ที่แคบกว่า granularity ของตัวเองได้ ใช้ bootstrap7 เมื่อไม่มีสูตร ซึ่งเป็นกรณีที่น่าสนใจส่วนใหญ่: F1, macro-averages, BLEU, pass@1, คะแนนของ judge ที่ใช้ rubric บนสายพานนี้ F1 ของ tuned model เท่ากับ 0.4122 มี bootstrap interval [0.3009, 0.5156] — ซึ่งคือตัวเลขที่ควรปรากฏในรายงาน เพราะ point estimate เพียงอย่างเดียวชวนให้เปรียบเทียบในแบบที่มันรองรับไม่ได้
อีกหนึ่งการวัด เพราะมันเปลี่ยนวิธีที่คุณควรเปรียบเทียบสอง models สอง models ถูกให้คะแนนบน 500 examples เดียวกัน:
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)intervals ของมัน overlap กัน และกฎชาวบ้าน — error bars ที่ overlap หมายถึงไม่มีความแตกต่างอย่างมีนัยสำคัญ — จะเรียกการเปรียบเทียบนี้ว่า inconclusive แต่มันไม่ใช่ สอง models run บน examples เดียวกัน ดังนั้นปริมาณที่ถูกต้องคือความต่างราย example ซึ่ง interval ของมันคือ [0.0260, 0.0680] สูงกว่าศูนย์อย่างสบาย ทั้งสองไม่เห็นด้วยกันเพียง 31 จาก 500 items และ A ชนะ 27 จากความไม่เห็นด้วยเหล่านั้น examples ที่ใช้ร่วมกัน ทั้งง่ายและยาก จะหักล้างกันแทนที่จะเพิ่ม noise เปรียบเทียบ models แบบ paired แล้วคุณจะได้ข้อสรุปเดียวกันจากข้อมูลเพียงเศษหนึ่ง
ต่อจากนี้ไปไหน
ลิงก์ไปยังส่วน: ต่อจากนี้ไปไหนตอนนี้คุณมี model ที่ output probabilities ที่ calibrate แล้ว, loss ที่ derive จากข้ออ้างเกี่ยวกับข้อมูลแทนที่จะเลือกเพราะสะดวก, gradient ที่เป็น prediction ลบ truth อย่างแท้จริง และ — สำคัญกว่านั้น — กลไกในการค้นหาว่าสิ่งเหล่านี้ใช้ได้จริงหรือไม่ Wilson interval สิบบรรทัดด้านบนถูกนำกลับมาใช้ซ้ำแบบคำต่อคำ: มันรองรับ prompt variants ใน บทที่ 15, ตาราง retrieval ใน บทที่ 19 และ golden set ใน บทที่ 29 bootstrap คือสิ่งที่คุณหยิบมาใช้เมื่อไม่มีสูตร
แต่ model ยังมีเพียง layer เดียว มันวาดเส้น และบทที่ 1 พิสูจน์ด้วย XOR สี่แถวแล้วว่าเส้นเดียวไม่พอ วิธีแก้คือ stack: layer แรกที่บิด space, layer ที่สองที่วาดเส้นใน space ที่ถูกบิดแล้ว
นั่นคือจุดที่ gradient เรียบร้อยของบทนี้หมดทาง ทุกอย่างข้างบนทำงานได้เพราะ สามารถเขียนด้วยมือได้ครั้งเดียว สำหรับ model ที่มี layer เดียวระหว่าง input กับ loss ใส่ layer ที่สองไว้ตรงกลาง แล้วคำถามเปลี่ยนรูปร่าง: derivative ของ loss เทียบกับ weight ที่ไม่ได้แตะ output เลยคืออะไร — weight ที่อิทธิพลของมันมาถึงเฉพาะผ่าน layer อื่น อาจตามหลายเส้นทางพร้อมกัน?
derivative นั้นมีอยู่ การคำนวณด้วยมือสิ้นหวังสำหรับอะไรก็ตามที่ใหญ่กว่า toy และการคำนวณทีละ parameter ก็สิ้นหวังในอีก scale หนึ่ง สิ่งที่ต้องการคือ procedure ที่ได้ derivative ทุกตัวใน network จาก backward pass เพียงครั้งเดียวบน graph เดียวกับที่ forward pass เพิ่งเดินผ่านมา
นั่นคือ บทที่ 5 และมันคือ engine ที่คอร์สที่เหลือทั้งหมดใช้ขับเคลื่อน
แหล่งที่มาและวิธีการ
ลิงก์ไปยังส่วน: แหล่งที่มาและวิธีการสิ่งที่ควรอ่านควบคู่กับบทนี้ด้วย: Bishop, Pattern Recognition and Machine Learning §1.2, §1.5, §1.6 and §4.3 ซึ่งครอบคลุม probability, decision theory, information theory และ linear classification ตามลำดับที่บทนี้เดินตาม; Murphy, Probabilistic Machine Learning: An Introduction, chapters 6 and 10; Prince, Understanding Deep Learning §5.4–5.7; และ Saito and Rehmsmeier, The Precision-Recall Plot Is More Informative than the ROC Plot When Evaluating Binary Classifiers on Imbalanced Datasets (PLOS ONE, 2015) — ทำไม AUC ที่ quote ข้างต้นไม่ควรเป็นตัวเลข threshold-free เพียงตัวเดียวที่คุณดู เมื่อ 1.7 % ของชิ้นส่วนเป็น defective
รายการอ้างอิง
ลิงก์ไปยังส่วน: รายการอ้างอิง-
Ma, T. and Ng, A. CS229 Lecture Notes, Stanford University, chapters 2 and 3. จุดที่การหักล้างซึ่งให้ เลิกดูเหมือนโชค: เลือกการแจกแจง exponential-family ที่ตรงกับ output ของคุณ ใช้ canonical link ของมัน แล้ว gradient จะเป็น prediction ลบ truth เสมอ ↩
-
Olah, C. Visual Information Theory (2015),
colah.github.io/posts/2015-09-Visual-Information. คำอธิบายที่ชัดที่สุดที่มีของ entropy, cross-entropy และ KL divergence ในฐานะต้นทุนเป็น bits แทนที่จะเป็นสูตร ↩ -
Abu-Mostafa, Y. S., Magdon-Ismail, M. and Lin, H.-T. Learning From Data (AMLBook, 2012), lectures 13 and 17 of the Caltech course. Lecture 13 คือ validation; lecture 17 ว่าด้วยหลักการเรียนรู้สามข้อ คือจุดที่ data snooping ถูกตั้งชื่อ ทั้งสองรวมกันเป็นที่มาของวินัยในบทนี้: ทุกครั้งที่มอง data set คือการตัดสินใจ fitting ไม่ว่าคุณจะรัน optimiser หรือไม่ ↩
-
James, G., Witten, D., Hastie, T. and Tibshirani, R. An Introduction to Statistical Learning, 2nd edition (Springer, 2021), chapters 2 and 5, สำหรับ bias–variance decomposition และ resampling. companion volume คือจุดที่กับดักการเลือกถูกพูดตรง ๆ: Hastie, Tibshirani and Friedman, The Elements of Statistical Learning, 2nd edition, §7.10.2, The Wrong and Right Way to Do Cross-validation. ↩
-
Kaufman, S., Rosset, S., Perlich, C. and Stitelman, O. Leakage in Data Mining: Formulation, Detection, and Avoidance. ACM Transactions on Knowledge Discovery from Data 6(4), 2012. การวิเคราะห์เชิง formal ของความล้มเหลวที่สาธิตข้างต้น พร้อม case studies จากการแข่งขันที่ชนะโดย model ซึ่งเรียนรู้ artefact ของวิธีประกอบข้อมูล ↩
-
Wilson, E. B. Probable Inference, the Law of Succession, and Statistical Inference. Journal of the American Statistical Association 22(158), pp. 209–212 (1927). score interval ที่ใช้ใน
wilson()ข้างต้น ซึ่งยังเป็น default ที่ถูกต้องสำหรับ proportion ส่วน textbook interval คือสิ่งที่ควรหลีกเลี่ยง: มันให้ผลไร้สาระใกล้ 0 และ 1 และ undercovers อย่างหนักเมื่อ เล็ก ↩ -
Efron, B. Bootstrap Methods: Another Look at the Jackknife. The Annals of Statistics 7(1), pp. 1–26 (1979). แนวคิดที่ให้คุณใส่ interval ให้ statistic ใดก็ได้ที่คุณคำนวณได้ รวมถึงสิ่งที่ไม่มี sampling theory ↩