Classificação, entropia cruzada e como não se enganar
Construa um classificador logístico e veja por que 98% de acurácia pode significar um modelo que não encontra nada.
Nesta página
Um modelo que responde esta peça está boa para todas as peças que saem da esteira está certo em 98,15% das vezes. Ele também é inútil: das 74 peças defeituosas no conjunto de teste, não pega nenhuma.
As duas frases descrevem o mesmo modelo. A distância entre elas é este capítulo.
A primeira metade constrói o classificador. Ela quase não precisa de nada novo: o Capítulo 2 deu a receita para transformar uma suposição sobre como os dados são produzidos em uma função de perda, e o Capítulo 3 deu a maquinaria para descer ladeira abaixo em qualquer perda que essa receita entregue. Aplique as duas a uma pergunta de sim/não e a regressão logística aparece, mais uma ideia nova — um logit — que será cobrada de novo no Capítulo 17.
A segunda metade é a mais difícil. Tudo depois deste ponto no curso é julgado por um número que alguém mediu, e se você não consegue diferenciar uma melhoria real de um artefato de medição, todo capítulo seguinte é enfeite. Então: a matriz de confusão, precision e recall, as três divisões, vazamento, e a pergunta que quase ninguém responde com honestidade — de quantos exemplos de teste eu realmente preciso?
A aritmética aqui roda sobre 20.000 linhas, então tudo é vetorizado do começo ao fim — NumPy vem fazendo o trabalho desde o Capítulo 2, e daqui em diante isso deixa de valer comentário.
A esteira, com uma pergunta mais rara
Link para a seção: A esteira, com uma pergunta mais raraA mesma fábrica do Capítulo 1, pergunta mais difícil. Em vez de aceitar ou rejeitar, a pergunta é esta peça está defeituosa — e defeitos são raros, o que torna a metade de medição deste capítulo difícil e a metade de modelagem enganosamente fácil.
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 74Três divisões, não duas. O motivo merece uma seção própria e ganha uma abaixo; por enquanto, treine na primeira, ajuste na segunda e não olhe para a terceira.
As features são padronizadas — subtrai-se a média, divide-se pelo desvio padrão — usando apenas as estatísticas de treino, pelo motivo que o Capítulo 1 demonstrou com o limite de convergência do perceptron: dados não centralizados tornam a geometria hostil. De quais linhas você tem permissão para calcular essa média vira uma questão viva mais adiante neste capítulo.
De um veredito a uma probabilidade
Link para a seção: De um veredito a uma probabilidadeO perceptron retornava um sinal. Um sinal não consegue distinguir rejeitar de rejeitar, mas por pouco, e essa diferença é exatamente o que uma fábrica precisa para decidir quais peças um humano deve reinspecionar primeiro.
Então siga literalmente a receita do Capítulo 2. Escreva o que você afirma sobre como um rótulo é produzido, pegue a verossimilhança, pegue o log, negue, e você tem uma perda. Para um resultado de sim/não, a afirmação é uma distribuição Bernoulli: existe uma probabilidade de que a peça esteja defeituosa, e
que é só uma forma compacta de escrever " se , e se ". Pegue o log disso e negue, e a perda para um exemplo é
Isso é entropia cruzada binária. Ela não foi escolhida porque é conveniente; é a log-verossimilhança negativa da única distribuição que um cara ou coroa pode ter. Não havia mais nada disponível.
O que ainda falta é de onde vem . O modelo calcula uma soma ponderada , que é um número real e percorre a reta inteira, e uma probabilidade precisa viver em . A função que faz a passagem entre os dois é a sigmoide logística:
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.9820Leia a coluna da direita como uma lista de preços. Estar certo com 90% de confiança custa 0,105. Recusar-se a se comprometer custa 0,693 — que é , o preço de dar de ombros. Estar confiantemente errado custa 4,6, quarenta e quatro vezes mais, e o preço cresce sem limite conforme o modelo fica mais certo sobre um erro. A entropia cruzada não apenas conta erros: ela cobra pela arrogância.
O gradient é previsão menos verdade
Link para a seção: O gradient é previsão menos verdadeO Capítulo 3 disse: para treinar qualquer coisa, obtenha a derivada da perda em relação a cada parâmetro. Faça isso para um exemplo. Com e :
Mostrar detalhes
As duas linhas que fazem a bagunça cancelar. A sigmoide tem uma derivada incomumente agradável, . E a perda deriva para
Multiplique as duas pela regra da cadeia e aparece uma vez em cima e uma vez embaixo. Cancela exatamente, e é o que sobra. Esse cancelamento não é coincidência — é o que acontece sempre que a perda é a log-verossimilhança negativa de uma distribuição e a função de saída é a que essa distribuição usa naturalmente. Esse pareamento tem um nome — um modelo linear generalizado — e o gradient arrumado é sua impressão digital.1
Então a atualização é previsão menos verdade, vezes a entrada. Nada mais. Aqui está o treinador inteiro, que é o gradient descent do Capítulo 3 com uma linha alterada:
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, bO np.where em sigmoid não é cosmético. Calcular diretamente estoura para negativos grandes; o desvio escolhe a forma algebricamente idêntica que mantém o expoente negativo. Esta é a caixa de ponto flutuante do Capítulo 2 cobrando sua primeira dívida, e ela cobrará uma maior daqui a duas seções.
Por que não erro quadrático, e por que a resposta é sobre o gradient
Link para a seção: Por que não erro quadrático, e por que a resposta é sobre o gradientA explicação padrão para preferir entropia cruzada a erro quadrático é o argumento da verossimilhança acima: erro quadrático é o que você obtém ao assumir ruído gaussiano, rótulos não são gaussianos, portanto não faça isso. Está correto e não convence ninguém, porque você pode escrever sobre uma sigmoide e isso vai treinar.
O argumento que funciona é sobre o gradient. Coloque erro quadrático em cima de uma sigmoide e a regra da cadeia dá
Esse extra é o que cancelou antes. Agora não cancela, e vai a zero sempre que o modelo está confiante — inclusive quando o modelo está confiantemente errado. Avalie os dois em algumas pontuações, para um exemplo cujo rótulo verdadeiro é 1:
| pontuação | entropia cruzada | erro quadrático | razão | |
|---|---|---|---|---|
| 0,000335 | 1.491 | |||
| 0,017986 | 28,3 | |||
| 0,119203 | 4,8 | |||
| 0,500000 | 2,0 | |||
| 0,880797 | 4,8 |
Em o modelo está tão errado quanto é possível estar, e o erro quadrático responde com um gradient 1.491 vezes menor que o da entropia cruzada. Quanto pior o erro, menos o modelo aprende com ele. O gradient da entropia cruzada, enquanto isso, satura em : maximamente errado produz um sinal maximamente grande, e não maior.
Rode a corrida. Dois mil pontos balanceados, pesos iniciais idênticos escolhidos para estarem confiantemente errados (), taxa de aprendizado idêntica, só a perda muda. As duas execuções são pontuadas com entropia cruzada para que as colunas sejam comparáveis.
| epoch | perda de entropia cruzada | acurácia | perda de erro quadrático | acurácia |
|---|---|---|---|---|
| 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 |
A entropia cruzada terminou no epoch 50. O erro quadrático ainda está em 24% de acurácia no epoch 100 — e não tinha saído de 23% no epoch 10 — pior que chutar, porque começou confiantemente errado e o gradient que o resgataria foi multiplicado por 0,0007. Ele escapa por volta do epoch 500 e chega ao mesmo lugar. Então o resumo honesto é que erro quadrático sobre uma sigmoide não é incorreto; ele é lento exatamente onde velocidade mais importa. Em um modelo de dois parâmetros, você perde 450 epochs. Em uma rede com cem camadas, onde alguma unidade em algum lugar está sempre confiantemente errada, você perde a execução de treino.
Entropia, entropia cruzada e KL, em uma página
Link para a seção: Entropia, entropia cruzada e KL, em uma páginaTrês quantidades, necessárias de verdade no Capítulo 8 para perplexidade e no Capítulo 11 para a penalidade que mantém uma política fine-tuned perto de sua referência. Elas são mais fáceis do que a fama sugere.2
Entropia é o número médio de bits que você precisa gastar para comunicar uma amostra de uma distribuição, se usar o melhor código possível para ela:
Entropia cruzada é o que você gasta quando usa um código construído para em dados que na verdade vêm de :
Divergência KL é o excesso — o desperdício, em bits, causado por acreditar em quando a verdade é :
Confira as três na esteira:
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 bitsDuas coisas ficam visíveis ali. Primeiro, um modelo que simplesmente reporta a taxa base de treino, 1,69%, alcança uma entropia cruzada de 0,1330 bits, quase exatamente a entropia dos rótulos de teste — como deve ser, já que ele tem a distribuição certa e nenhuma outra informação. Entropia é o piso que a ignorância sobre o indivíduo compra para você. Segundo, um modelo que dá de ombros e diz 0,5 paga exatamente 1 bit, e a lacuna entre os dois, 0,8671 bits, é precisamente a divergência KL. não é uma identidade para decorar; é uma conta que você pode ver sendo somada.
E a conexão de volta ao treino: quando o rótulo é uma única classe conhecida, a distribuição “verdadeira” é one-hot, sua entropia é zero, e a entropia cruzada é igual à divergência KL. Minimizar a entropia cruzada e puxar a distribuição do modelo em direção à verdade são o mesmo ato.
Mais de duas respostas: softmax, e o deslocamento que não custa nada
Link para a seção: Mais de duas respostas: softmax, e o deslocamento que não custa nadaDefeituoso não é uma coisa só. Na moldagem, uma peça pode sair com injeção incompleta (material insuficiente), rebarba (material demais, espremido para fora do molde) ou queima. Quatro resultados, portanto quatro logits, e eles precisam virar quatro probabilidades que somam um. Isso é softmax:
Ele tem uma propriedade que parece acidente e na verdade é toda a implementação:
para qualquer constante , porque e o se cancelam em cima e embaixo. Só diferenças entre logits significam alguma coisa. O nível absoluto não é informação.
Ainda bem, porque o nível absoluto é o que quebra o computador:
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 estoura um float de 64 bits, a soma vira infinito, e infinito dividido por infinito é nan — não um erro, não uma falha, apenas um buraco silencioso onde três probabilidades costumavam estar. Subtrair o maior logit não muda nada matematicamente e muda tudo numericamente, porque o maior expoente se torna exatamente . Este é o truque logsumexp do Capítulo 2 vestindo roupa de trabalho, e toda implementação séria faz isso:
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, bO gradient é novamente previsão menos verdade, agora com one-hot. O caso binário foi um caso especial o tempo todo.
Treinado em 3.000 peças e testado em 1.000, com três medições cada (largura, peso, temperatura de fusão), ele chega a 94,00% de acurácia. Aqui está o que esse número esconde:
| verdade ↓ / previsto → | ok | injeção incompleta | rebarba | queima | recall |
|---|---|---|---|---|---|
| ok | 850 | 5 | 9 | 0 | 0,984 |
| injeção incompleta | 22 | 21 | 0 | 0 | 0,488 |
| rebarba | 20 | 0 | 30 | 1 | 0,588 |
| queima | 3 | 0 | 0 | 39 | 0,929 |
| precision | 0,950 | 0,808 | 0,769 | 0,975 |
O modelo encontra menos da metade das peças com injeção incompleta. A acurácia não consegue ver isso, porque 86% das peças estão boas e acertar essas já basta para carregar a média. Macro F1 — a média dos F1 por classe, que dá a uma classe rara o mesmo peso de uma classe comum — é 0,7983, contra um micro F1 de 0,9400 que por definição é idêntico à acurácia. Sempre que alguém reportar um único número de F1, pergunte qual.
Esse é o fim da modelagem. O resto do capítulo é sobre os números.
Três modelos, uma acurácia
Link para a seção: Três modelos, uma acuráciaPegue o modelo binário treinado e crie duas variantes multiplicando todo logit por uma constante: 0,35 para uma versão hesitante, 4 para uma superconfiante. Multiplicar por um número positivo não pode mudar nenhum sinal, então os três modelos preveem exatamente o mesmo rótulo para todas as 4.000 peças de teste. A acurácia não consegue diferenciá-los. A entropia cruzada não tem dificuldade nenhuma:
| modelo | acurácia | entropia cruzada | perda média quando acerta | perda média quando erra | pior perda única |
|---|---|---|---|---|---|
| hesitante (logits × 0,35) | 0,9830 | 0,1549 | 0,1369 | 1,1990 | 2,80 |
| como treinado | 0,9830 | 0,0564 | 0,0147 | 2,4689 | 7,82 |
| superconfiante (logits × 4) | 0,9830 | 0,1563 | 0,0009 | 9,1427 | 27,63 |
O modelo hesitante paga uma pequena taxa em cada peça, inclusive nas milhares que acerta. O superconfiante é quase gratuito quando acerta e catastrófico quando erra — uma peça nesse conjunto de teste custa sozinha 27,63 nats. Os dois chegam quase ao mesmo total por caminhos opostos, e o modelo treinado, cujas probabilidades são calibradas aos dados, fica três vezes abaixo de ambos.
Esta é a forma mais afiada de declarar a diferença entre uma perda e uma métrica. A perda é o que você otimiza: precisa ser diferenciável, e vê tudo que o modelo disse, inclusive o quanto ele tinha certeza. A métrica é aquilo pelo qual você é julgado: pode ser uma função degrau, uma regra de negócio, uma contagem de defeitos perdidos. Elas não são o mesmo objeto e nem sempre concordam — por isso você define ambas antes de começar, e nunca deixa a perda substituir a métrica só porque ela apareceu na tela.
O baseline idiota vem primeiro
Link para a seção: O baseline idiota vem primeiroAntes de qualquer modelo, a exigência: quanto pontua a resposta mais preguiçosa possível? Nesta esteira, sempre diga que está boa:
always-say-fine baseline: accuracy = 0.9815
confusion (tn, fp, fn, tp) = (3926, 0, 74, 0)98,15%. Agora o modelo logístico treinado, no limiar padrão de 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%. Ele bateu o baseline por 0,15 ponto percentual, e qualquer relatório que pare na acurácia chamará isso de vitória. A matriz de confusão diz o que realmente aconteceu:
| previsto como boa | previsto como defeituosa | |
|---|---|---|
| realmente boa | 3.924 | 2 |
| realmente defeituosa | 66 | 8 |
Ele encontrou 8 peças defeituosas de 74 e deixou 66 passarem. Três números nomeiam as três formas de ler essa tabela:
- Precision . Das peças que ele sinalizou, quantas eram realmente defeituosas. Este é o custo de inspeções desperdiçadas.
- Recall . Das peças defeituosas, quantas ele pegou. Este é o custo de enviar uma peça ruim a um cliente.
- F1 , a média harmônica das duas, que fica perto da menor e portanto se recusa a ser bajulada por uma delas sozinha.
O que importa depende da fábrica, não da matemática: uma inspeção custa alguns segundos e um defeito enviado custa um aviso de recall, então aqui recall domina e 0,108 é uma falha.
Mas o modelo não é o problema. O limiar é, e o limiar não faz parte do modelo — é uma decisão de negócio aplicada depois a uma probabilidade. Faça a varredura:
| limiar | TP | FP | FN | acurácia | 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 |
Leia a coluna de acurácia de cima para baixo. Ela cai o tempo todo — de 98,30% para 65,93% — enquanto o modelo passa de pegar 8 defeitos para pegar 71 de 74. Tudo de útil que este modelo consegue fazer piora sua acurácia. Uma equipe otimizando o número do título enviaria a versão que não encontra nada.
Mostrar detalhes
Ponderação de classes não cria sinal, ela move o ponto de operação. O reflexo inicial comum com classes desbalanceadas é ponderar a classe rara na perda. Fazendo isso, com pesos 1, 10 e 60 nos positivos:
| peso nos positivos | acurácia | 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 e recall se movem bastante. A AUC — a probabilidade de que o modelo ranqueie uma peça defeituosa aleatória acima de uma peça boa aleatória, ignorando o limiar por completo — se move 0,0002, que é nada. Reponderar deslizou o mesmo modelo ao longo da mesma curva de trade-off. Muitas vezes é isso que você quer, e nunca é informação nova: se o ranking é ruim, nenhum esquema de ponderação vai salvá-lo.
Três divisões, e o vazamento que você está prestes a encontrar
Link para a seção: Três divisões, e o vazamento que você está prestes a encontrarPor que três divisões e não duas? Porque no momento em que você usa um conjunto de exemplos para escolher qualquer coisa — um limiar, uma taxa de aprendizado, qual de seis modelos colocar em produção — esse conjunto foi usado para ajuste, e sua pontuação deixa de ser não enviesada.3 Medido nesta esteira: varrer o limiar no conjunto de validação escolhe 0,196, e o modelo então pontua F1 = 0,4122 no conjunto de teste intocado. Se a varredura tivesse sido feita diretamente no conjunto de teste, o melhor alcançável ali seria 0,4186 — um número que ninguém tem direito de reportar.
A lacuna é pequena aqui, 0,006, porque é um hiperparâmetro varrido uma vez contra 4.000 exemplos de validação. Ela cresce a cada decisão extra e a cada encolhimento do conjunto de validação. Note também que a direção não é garantida em uma única execução: o limiar escolhido pontuou 0,3902 na validação e 0,4122 no teste, então a validação subestimou desta vez. O viés é sistemático ao longo de muitas decisões, não visível em uma só.4
Agora o exercício. O log da esteira chega com uma terceira coluna, station_seconds: quanto tempo cada peça passou na estação de inspeção. Adicioná-la é uma mudança de uma linha no pré-processamento. Eis o que ela faz:
| modelo | acurácia | precision | recall | F1 | entropia cruzada | AUC |
|---|---|---|---|---|---|---|
| largura + peso | 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 vai de 10,8% para 77,0%. F1 mais que quadruplica. E repare no que a acurácia fez: 98,30% → 99,20%, um ganho de nove décimos de ponto, que é o tipo de número arredondado para “cerca de 99% de qualquer jeito” em um slide de resumo. Acurácia não conseguiu ver a falha antes e agora não consegue ver a fraude.
Antes de continuar: o modelo está trapaceando. Descubra como.
Como caçar um vazamento, na ordem que o encontra mais rápido.
-
Compare treino e teste. Overfitting aparece como uma lacuna grande. Aqui: modelo honesto 0,9838 treino / 0,9830 teste; modelo com vazamento 0,9936 treino / 0,9920 teste. As duas lacunas ficam abaixo de 0,2 ponto. Um vazamento não parece overfitting — a feature vazada está igualmente disponível no teste, então o modelo generaliza lindamente para um mundo que não existe.
-
Treine um modelo por feature, isoladamente. Qualquer coisa que carregue a resposta se anuncia:
feature isolada acurácia recall F1 AUC largura 0,9815 0,014 0,026 0,8691 peso 0,9815 0,000 0,000 0,7914 station_seconds0,9850 0,405 0,500 0,9960 Uma coluna, sozinha, ranqueia defeitos com AUC 0,9960. Duas medições feitas por um paquímetro e uma balança conseguem 0,87 e 0,79. Essa assimetria é o alarme.
-
Pergunte quando cada número foi anotado. Tempo médio de permanência: 2,23 segundos para peças aprovadas, 15,56 segundos para peças reprovadas. Claro que sim. Uma peça permanece na estação porque um inspetor a tirou da esteira — o que acontece depois, e apenas porque, alguém decidiu que ela era defeituosa. A coluna não é uma medição da peça. É uma medição do veredito.
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()) A linha destacada é o vazamento: o tempo de permanência de uma peça defeituosa é sorteado de uma distribuição diferente, porque um humano a tirou da esteira. Este é o bug sério mais comum em machine learning aplicado, e tem nome: vazamento do alvo — informação nas features de treino que não estaria disponível no momento em que a previsão precisa ser feita.5 Ele não lança exceção. Produz um número melhor. Todo incentivo em um projeto aponta para mantê-lo.
A defesa é uma pergunta, feita a cada coluna: no instante em que preciso desta previsão, este valor já existe? Em uma esteira ao vivo, station_seconds é desconhecido até depois que a peça foi inspecionada — que é justamente o que o modelo deveria substituir.
De quantos exemplos de teste eu preciso?
Link para a seção: De quantos exemplos de teste eu preciso?Suponha que você pontue um modelo em 20 exemplos e ele acerte 17. Você reporta 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.6477A leitura honesta de 17/20 é algo entre 64% e 95%. Um modelo genuinamente de 65% produz esse resultado 4,4% das vezes — uma execução em vinte e três — e, se você tentou um punhado de prompts e reportou o melhor, você mesmo fabricou essa execução. Dezessete de vinte não conseguem distinguir um modelo de 85% de um de 65%.
Duas formas de colocar um intervalo em uma taxa, e ambas pertencem ao seu kit de ferramentas:
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)Use Wilson6 para uma taxa de sucesso simples; ele se comporta bem em qualquer e não precisa de aleatoriedade. Note acima que em a extremidade superior do bootstrap é 1,0000 — reamostrar 20 pontos pode facilmente sortear 20 corretos, então ele não consegue representar um intervalo mais estreito que sua própria granularidade. Use o bootstrap7 quando não existe fórmula, que é a maior parte dos casos interessantes: F1, médias macro, BLEU, pass@1, a pontuação de um juiz baseado em rubrica. Nesta esteira, o F1 de 0,4122 do modelo ajustado carrega um intervalo bootstrap de [0,3009, 0,5156] — que é o número que deveria aparecer no relatório, porque a estimativa pontual sozinha convida uma comparação que ela não pode sustentar.
Mais uma medição, porque ela muda como você deve comparar dois modelos. Dois modelos pontuados nos mesmos 500 exemplos:
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)Seus intervalos se sobrepõem, e a regra popular — barras de erro sobrepostas significam ausência de diferença significativa — chamaria a comparação de inconclusiva. Não é. Os dois modelos rodaram nos mesmos exemplos, então a quantidade certa é a diferença por exemplo, cujo intervalo é [0.0260, 0.0680], confortavelmente acima de zero. Eles discordam em apenas 31 dos 500 itens, e A vence 27 dessas discordâncias; os exemplos compartilhados, fáceis e difíceis, se cancelam em vez de adicionar ruído. Compare modelos de forma pareada, e você chega à mesma conclusão com uma fração dos dados.
Para onde isso vai agora
Link para a seção: Para onde isso vai agoraAgora você tem um modelo que gera probabilidades calibradas, uma perda derivada de uma afirmação sobre os dados em vez de escolhida por conveniência, um gradient que é literalmente previsão menos verdade, e — mais importante — a maquinaria para descobrir se qualquer uma dessas coisas funciona. O intervalo de Wilson de dez linhas acima é reutilizado literalmente: ele carrega as variantes de prompt no Capítulo 15, as tabelas de recuperação no Capítulo 19 e o conjunto dourado no Capítulo 29. O bootstrap é o que você procura quando não existe fórmula.
Mas o modelo ainda tem uma camada. Ele desenha uma linha, e o Capítulo 1 provou com quatro linhas de XOR que uma linha não basta. A correção é empilhar: uma primeira camada que dobra o espaço, uma segunda que desenha a linha no espaço dobrado.
É aí que o gradient arrumado deste capítulo acaba. Tudo acima funcionou porque podia ser escrito à mão, uma vez, para um modelo com uma camada entre a entrada e a perda. Coloque uma segunda camada no meio e a pergunta muda de forma: qual é a derivada da perda em relação a um peso que não toca a saída de forma alguma — um peso cuja influência chega apenas por meio de outra camada, possivelmente por vários caminhos ao mesmo tempo?
Essa derivada existe. Calculá-la à mão é inviável para qualquer coisa maior que um brinquedo, e calculá-la um parâmetro por vez é inviável em outra escala. O que é necessário é um procedimento que obtenha cada derivada na rede a partir de uma única passada para trás sobre o mesmo grafo que a passada para frente acabou de percorrer.
Esse é o Capítulo 5, e ele é o motor em que o resto deste curso roda.
Fontes e método
Link para a seção: Fontes e métodoTambém vale ler junto com este capítulo: Bishop, Pattern Recognition and Machine Learning §1.2, §1.5, §1.6 e §4.3, que cobre probabilidade, teoria da decisão, teoria da informação e classificação linear na ordem que este capítulo segue; Murphy, Probabilistic Machine Learning: An Introduction, capítulos 6 e 10; Prince, Understanding Deep Learning §5.4–5.7; e Saito e Rehmsmeier, The Precision-Recall Plot Is More Informative than the ROC Plot When Evaluating Binary Classifiers on Imbalanced Datasets (PLOS ONE, 2015) — por que a AUC citada acima não deve ser o único número independente de limiar que você olha quando 1,7% das peças são defeituosas.
Referências
Link para a seção: Referências-
Ma, T. e Ng, A. CS229 Lecture Notes, Stanford University, capítulos 2 e 3. Onde o cancelamento que produz deixa de parecer sorte: escolha a distribuição da família exponencial que corresponde à sua saída, use seu link canônico, e o gradient é sempre previsão menos verdade. ↩
-
Olah, C. Visual Information Theory (2015),
colah.github.io/posts/2015-09-Visual-Information. A explicação mais clara disponível de entropia, entropia cruzada e divergência KL como custos em bits, não como fórmulas. ↩ -
Abu-Mostafa, Y. S., Magdon-Ismail, M. e Lin, H.-T. Learning From Data (AMLBook, 2012), aulas 13 e 17 do curso da Caltech. A aula 13 é validação; a aula 17, sobre os três princípios de aprendizado, é onde data snooping é nomeado. Juntas, elas são a fonte da disciplina neste capítulo: cada olhada em um conjunto de dados é uma decisão de ajuste, tenha você rodado um otimizador ou não. ↩
-
James, G., Witten, D., Hastie, T. e Tibshirani, R. An Introduction to Statistical Learning, 2ª edição (Springer, 2021), capítulos 2 e 5, para a decomposição viés–variância e para reamostragem. O volume complementar é onde a armadilha de seleção é declarada diretamente: Hastie, Tibshirani e Friedman, The Elements of Statistical Learning, 2ª edição, §7.10.2, The Wrong and Right Way to Do Cross-validation. ↩
-
Kaufman, S., Rosset, S., Perlich, C. e Stitelman, O. Leakage in Data Mining: Formulation, Detection, and Avoidance. ACM Transactions on Knowledge Discovery from Data 6(4), 2012. Um tratamento formal da falha demonstrada acima, com estudos de caso de competições vencidas por um modelo que aprendeu um artefato de como os dados foram montados. ↩
-
Wilson, E. B. Probable Inference, the Law of Succession, and Statistical Inference. Journal of the American Statistical Association 22(158), pp. 209–212 (1927). O intervalo de score usado em
wilson()acima, ainda o padrão certo para uma proporção. O intervalo de livro-texto é o que deve ser evitado: ele dá absurdos perto de 0 e 1, e subcobre gravemente em pequeno. ↩ -
Efron, B. Bootstrap Methods: Another Look at the Jackknife. The Annals of Statistics 7(1), pp. 1–26 (1979). A ideia que permite colocar um intervalo em qualquer estatística que você consiga calcular, incluindo as que não têm teoria amostral. ↩