Classificação, entropia cruzada e como não se enganar
Crie um classificador logístico com a perda do cap. 2 e a descida do cap. 3; veja porque 98 % de exatidão pode não encontrar nada.
Nesta página
Um modelo que responde esta peça está boa sobre todas as peças que saem da linha está certo 98,15 % das vezes. Também não vale nada: das 74 peças defeituosas no conjunto de teste, não apanha nenhuma.
As duas frases descrevem o mesmo modelo. A distância entre elas é este capítulo.
A primeira metade constrói o classificador. Quase não precisa de nada novo: o Capítulo 2 deu a receita para transformar uma hipótese sobre como os dados são produzidos numa função de perda, e o Capítulo 3 deu a maquinaria para descer qualquer perda que essa receita lhe entregue. Aplique ambos a uma pergunta de sim/não e obtém regressão logística, mais uma ideia nova — um logit — que voltará a ser cobrada no Capítulo 17.
A segunda metade é a mais difícil. Tudo a partir daqui no curso é julgado por um número que alguém mediu, e se não conseguir distinguir uma melhoria real de um artefacto de medição, todos os capítulos seguintes são decoração. Portanto: a matriz de confusão, precisão e revocação, as três divisões, leakage, e a pergunta a que quase ninguém responde honestamente — de quantos exemplos de teste preciso realmente?
A aritmética aqui percorre 20.000 linhas, por isso é sempre vectorizada — o NumPy faz o trabalho desde o Capítulo 2 e, a partir daqui, isso deixa de merecer comentário.
A linha, com uma pergunta mais rara
Ligação para a secção: A linha, 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 os defeitos são raros, o que torna a metade de medição deste capítulo difícil e a metade de modelação enganadoramente 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. A razão merece uma secção própria e terá uma abaixo; por agora, treine na primeira, ajuste na segunda e não olhe para a terceira.
As features são normalizadas — subtrai-se a média, divide-se pelo desvio-padrão — usando apenas as estatísticas de treino, pela razão que o Capítulo 1 demonstrou com o limite de convergência do perceptron: dados não centrados tornam a geometria hostil. As linhas a partir das quais é permitido calcular essa média tornam-se uma questão viva mais adiante neste capítulo.
De um veredicto para uma probabilidade
Ligação para a secção: De um veredicto para uma probabilidadeO perceptron devolvia um sinal. Um sinal não consegue distinguir rejeitar de rejeitar, mas por pouco, e essa diferença é exactamente aquilo de que uma fábrica precisa para decidir que peças devem ser reinspeccionadas primeiro por uma pessoa.
Por isso, siga literalmente a receita do Capítulo 2. Escreva o que afirma sobre como uma etiqueta é produzida, calcule a verosimilhança, calcule o logaritmo, negue-o, e tem uma perda. Para um resultado de sim/não, a afirmação é uma distribuição Bernoulli: há uma probabilidade de a peça estar defeituosa, e
que é apenas uma forma compacta de escrever « se , e se ». Calcule o logaritmo disso e negue-o, e a perda para um exemplo é
Isto é entropia cruzada binária. Não foi escolhida por ser conveniente; é a verosimilhança logarítmica negativa da única distribuição que um lançamento de moeda 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 toda a recta, e uma probabilidade tem de viver em . A função que faz a passagem entre ambos é 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 tabela de preços. Estar certo com 90 % de confiança custa 0,105. Recusar comprometer-se custa 0,693 — que é , o preço de um encolher de ombros. Estar confiantemente errado custa 4,6, quarenta e quatro vezes mais, e o preço sobe sem limite à medida que o modelo fica mais certo de um erro. A entropia cruzada não se limita a contar erros: cobra pela arrogância.
O gradient é previsão menos verdade
Ligação para a secçã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-o para um exemplo. Com e :
Mostrar detalhes
As duas linhas que fazem a confusão cancelar. A sigmoide tem uma derivada invulgarmente agradável, . E a perda diferencia-se para
Multiplique as duas pela regra da cadeia e aparece uma vez em cima e uma vez em baixo. Cancela exactamente, e é o que sobrevive. Esse cancelamento não é uma coincidência — é o que acontece sempre que a perda é a verosimilhança logarítmica negativa de uma distribuição e a função de saída é aquela que essa distribuição usa naturalmente. Esse emparelhamento tem um nome — um modelo linear generalizado — e o gradient limpo é a sua impressão digital.1
Portanto, a actualização é previsão menos verdade, vezes a entrada. Nada mais. Eis o trainer inteiro, que é a descida 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 directamente faz overflow para muito negativo; o ramo escolhe a forma algebricamente idêntica que mantém o expoente negativo. É a caixa de vírgula flutuante do Capítulo 2 a cobrar a sua primeira dívida, e vai cobrar uma maior daqui a duas secções.
Porque não erro quadrático, e porque a resposta é sobre o gradient
Ligação para a secção: Porque não erro quadrático, e porque a resposta é sobre o gradientA explicação habitual para preferir entropia cruzada ao erro quadrático é o argumento da verosimilhança acima: o erro quadrático é o que se obtém ao assumir ruído gaussiano, as etiquetas não são gaussianas, logo não o faça. Está correcto e não convence ninguém, porque é possível escrever sobre uma sigmoide e isso treina.
O argumento que pega é sobre o gradient. Ponha erro quadrático em cima de uma sigmoide e a regra da cadeia dá
Esse extra é o que antes cancelava. Agora não cancela, e vai para zero sempre que o modelo está confiante — incluindo quando o modelo está confiantemente errado. Avalie ambos para algumas pontuações, num exemplo cuja etiqueta verdadeira é 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 do que o da entropia cruzada. Quanto pior o erro, menos o modelo aprende com ele. O gradient da entropia cruzada, por sua vez, satura em : maximamente errado produz um sinal maximamente grande, e não maior.
Corra a corrida. Dois mil pontos equilibrados, pesos iniciais idênticos escolhidos para estarem confiantemente errados (), learning rate idêntica, só a perda difere. Ambas as execuções são avaliadas com entropia cruzada para que as colunas sejam comparáveis.
| época | perda de entropia cruzada | exatidão | perda de erro quadrático | exatidão |
|---|---|---|---|---|
| 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 acaba na época 50. O erro quadrático ainda está com 24 % de exatidão na época 100 — e não tinha saído dos 23 % na época 10 — pior do que adivinhar, porque começou confiantemente errado e o gradient que o salvaria foi multiplicado por 0,0007. Escapa por volta da época 500 e aterra no mesmo sítio. Portanto, o resumo honesto é que erro quadrático sobre uma sigmoide não é incorrecto; é lento exactamente onde a velocidade mais importa. Num modelo de dois parâmetros perde 450 épocas. Numa rede com cem camadas, onde alguma unidade em algum ponto está sempre confiantemente errada, perde a execução de treino.
Entropia, entropia cruzada e KL, numa página
Ligação para a secção: Entropia, entropia cruzada e KL, numa páginaTrês quantidades, necessárias a sério no Capítulo 8 para perplexidade e no Capítulo 11 para a penalização que mantém uma policy com fine-tuning perto da sua referência. São mais fáceis do que a sua reputação.2
Entropia é o número médio de bits que tem de gastar para comunicar uma amostra de uma distribuição, se usar o melhor código possível para ela:
Entropia cruzada é o que 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 é :
Verifique as três na linha:
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 são visíveis ali. Primeiro, um modelo que simplesmente comunica a taxa de base de treino, 1,69 %, atinge uma entropia cruzada de 0,1330 bits, quase exactamente a entropia das etiquetas de teste — como deve, visto que tem a distribuição certa e nenhuma outra informação. A entropia é o chão que a ignorância do indivíduo compra. Segundo, um modelo que encolhe os ombros e diz 0,5 paga exactamente 1 bit, e a diferença entre os dois, 0,8671 bits, é precisamente a divergência KL. não é uma identidade para memorizar; é uma conta que pode ver a ser somada.
E a ligação de volta ao treino: quando a etiqueta é uma única classe conhecida, a distribuição «verdadeira» é one-hot, a sua entropia é zero, e a entropia cruzada é igual à divergência KL. Minimizar a entropia cruzada e puxar a distribuição do modelo para a verdade são o mesmo acto.
Mais de duas respostas: softmax, e o deslocamento que não custa nada
Ligação para a secção: Mais de duas respostas: softmax, e o deslocamento que não custa nadaDefeituosa não é uma coisa só. Na moldagem, uma peça pode sair como short shot (material insuficiente), flash (material a mais, espremido para fora do molde) ou burn. Quatro resultados, portanto quatro logits, e têm de se tornar quatro probabilidades que somam um. Isso é softmax:
Tem uma propriedade que parece um acidente e é, na verdade, toda a implementação:
para qualquer constante , porque e o cancelam em cima e em baixo. Só as diferenças entre logits significam alguma coisa. O nível absoluto não é informação.
Ainda bem, porque o nível absoluto é o que avaria 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 faz overflow num float de 64 bits, a soma torna-se infinito, e infinito dividido por infinito é nan — não é um erro, não é uma falha, apenas um buraco silencioso onde costumavam estar três probabilidades. Subtrair o maior logit não muda nada matematicamente e muda tudo numericamente, porque o maior expoente passa a ser exactamente . É o truque logsumexp do Capítulo 2 vestido com roupa de trabalho, e todas as implementações sérias o fazem:
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 é outra vez previsão menos verdade, agora com one-hot. O caso binário foi sempre um caso especial.
Treinado em 3.000 peças e testado em 1.000, com três medições cada (largura, peso, temperatura de fusão), atinge 94,00 % de exatidão. Eis o que esse número esconde:
| verdade ↓ / previsto → | ok | short shot | flash | burn | revocação |
|---|---|---|---|---|---|
| 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 |
| precisão | 0,950 | 0,808 | 0,769 | 0,975 |
O modelo encontra menos de metade dos short shots. A exatidão não consegue ver isto, porque 86 % das peças estão boas e acertar nessas chega para sustentar a média. Macro F1 — a média dos F1 por classe, que dá a uma classe rara o mesmo peso que a uma classe comum — é 0,7983, contra um micro F1 de 0,9400 que, por definição, é idêntico à exatidão. Sempre que alguém comunica um único número F1, pergunte qual.
Este é o fim da modelação. O resto do capítulo é sobre os números.
Três modelos, uma exatidão
Ligação para a secção: Três modelos, uma exatidãoPegue no modelo binário treinado e faça duas variantes multiplicando cada logit por uma constante: 0,35 para uma versão hesitante, 4 para uma versão excessivamente confiante. Multiplicar por um número positivo não pode mudar nenhum sinal, por isso os três modelos prevêem exactamente a mesma etiqueta para todas as 4.000 peças de teste. A exatidão não os distingue. A entropia cruzada não tem qualquer dificuldade:
| modelo | exatidão | entropia cruzada | perda média quando acerta | perda média quando erra | pior perda individual |
|---|---|---|---|---|---|
| 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 |
| excessivamente confiante (logits × 4) | 0,9830 | 0,1563 | 0,0009 | 9,1427 | 27,63 |
O modelo hesitante paga uma pequena taxa por cada peça, incluindo os milhares em que acerta. O excessivamente confiante é quase grátis quando acerta e catastrófico quando erra — uma peça nesse conjunto de teste custa-lhe 27,63 nats sozinha. Os dois chegam quase ao mesmo total por caminhos opostos, e o modelo treinado, cujas probabilidades estão calibradas aos dados, fica três vezes abaixo de ambos.
Esta é a forma mais afiada de dizer a diferença entre uma perda e uma métrica. A perda é o que optimiza: tem de ser diferenciável, e vê tudo o que o modelo disse, incluindo quão certo estava. A métrica é aquilo pelo qual é julgado: pode ser uma função degrau, uma regra de negócio, uma contagem de defeitos falhados. Não são o mesmo objecto e nem sempre concordam — é por isso que define ambos antes de começar, e nunca deixa a perda substituir a métrica só porque está no ecrã.
A linha de base estúpida vem primeiro
Ligação para a secção: A linha de base estúpida vem primeiroAntes de qualquer modelo, o requisito: quanto pontua a resposta mais preguiçosa possível? Nesta linha, dizer sempre 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 threshold predefinido 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 %. Bateu a linha de base por 0,15 pontos percentuais, e qualquer relatório que pare na exatidão chamará a isso uma vitória. A matriz de confusão diz o que realmente aconteceu:
| previsto boa | previsto defeituosa | |
|---|---|---|
| realmente boa | 3.924 | 2 |
| realmente defeituosa | 66 | 8 |
Três números dão nome às três formas de ler essa tabela:
- Precisão . Das peças que assinalou, quantas eram realmente defeituosas. Este é o custo de inspecções desperdiçadas.
- Revocação . Das peças defeituosas, quantas apanhou. Este é o custo de enviar uma peça má para um cliente.
- F1 , a média harmónica das duas, que fica perto da menor e por isso recusa ser lisonjeada por apenas uma delas.
Qual importa depende da fábrica, não da matemática: uma inspecção custa alguns segundos e um defeito enviado custa uma chamada de recolha, por isso aqui a revocação domina e 0,108 é um fracasso.
Mas o modelo não é o problema. O threshold é, e o threshold não faz parte do modelo — é uma decisão de negócio aplicada depois a uma probabilidade. Faça uma varredura:
| threshold | TP | FP | FN | exatidão | precisão | revocação | 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 da exatidão para baixo. Cai sempre — de 98,30 % para 65,93 % — enquanto o modelo passa de apanhar 8 defeitos para apanhar 71 de 74. Tudo o que este modelo consegue fazer de útil piora a sua exatidão. Uma equipa a optimizar o número de manchete enviaria a versão que não encontra nada.
Mostrar detalhes
Ponderar classes não cria sinal, move o ponto de operação. O reflexo habitual com classes desequilibradas é ponderar a classe rara na perda. Fazendo isso, com pesos de 1, 10 e 60 nos positivos:
| peso nos positivos | exatidão | precisão | revocação | 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 |
Precisão e revocação mexem muito. A AUC — a probabilidade de o modelo ordenar uma peça defeituosa aleatória acima de uma peça boa aleatória, ignorando inteiramente o threshold — mexe 0,0002, que é nada. A reponderação fez deslizar o mesmo modelo ao longo da mesma curva de compromisso. Muitas vezes é isso que quer, e nunca é informação nova: se a ordenação for má, nenhum esquema de ponderação a salva.
Três divisões, e a fuga que está prestes a encontrar
Ligação para a secção: Três divisões, e a fuga que está prestes a encontrarPorquê três divisões e não duas? Porque no momento em que usa um conjunto de exemplos para escolher qualquer coisa — um threshold, uma learning rate, qual de seis modelos enviar — esse conjunto foi usado para fitting, e a sua pontuação deixa de ser não enviesada.3 Medido nesta linha: varrer o threshold no conjunto de validação escolhe 0,196, e o modelo obtém então F1 = 0,4122 no conjunto de teste intacto. Se a varredura tivesse sido feita directamente no conjunto de teste, o melhor valor alcançável ali teria sido 0,4186 — um número que ninguém tem o direito de comunicar.
A diferença é pequena aqui, 0,006, porque é um hyperparameter varrido uma vez contra 4.000 exemplos de validação. Cresce com cada decisão adicional e com cada redução do conjunto de validação. Note também que a direcção não é garantida numa única execução: o threshold escolhido pontuou 0,3902 na validação e 0,4122 no teste, por isso a validação subestimou-o desta vez. O enviesamento é sistemático ao longo de muitas decisões, não visível numa só.4
Agora o exercício. O registo da linha chega com uma terceira coluna, station_seconds: quanto tempo cada peça passou na estação de inspecção. Acrescentá-la é uma alteração de uma linha no pré-processamento. Eis o que faz:
| modelo | exatidão | precisão | revocação | 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 |
A revocação passa de 10,8 % para 77,0 %. O F1 mais do que quadruplica. E repare no que fez a exatidão: 98,30 % → 99,20 %, um ganho de nove décimas de ponto, que é o tipo de número que num slide de resumo se arredonda para «cerca de 99 % em qualquer caso». A exatidão não viu a falha antes e agora não vê a fraude.
Antes de continuar: o modelo está a fazer batota. Descubra como.
Como caçar uma fuga, pela ordem que a encontra mais depressa.
-
Compare treino e teste. O overfitting aparece como uma grande diferença. Aqui: modelo honesto 0,9838 treino / 0,9830 teste; modelo com fuga 0,9936 treino / 0,9920 teste. Ambas as diferenças ficam abaixo de 0,2 pontos. Uma fuga não se parece com overfitting — a feature com fuga está igualmente disponível no teste, por isso o modelo generaliza lindamente para um mundo que não existe.
-
Treine um modelo por feature, isoladamente. Qualquer coisa que transporte a resposta anunciar-se-á:
feature isolada exatidão revocação 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, ordena 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 escrito. Tempo médio de permanência: 2,23 segundos para peças que passaram, 15,56 segundos para peças que falharam. Claro que sim. Uma peça permanece na estação porque um inspector a retirou da linha — o que acontece depois, e apenas porque, alguém decidiu que estava defeituosa. A coluna não é uma medição da peça. É uma medição do veredicto.
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 é a fuga: o tempo de permanência de uma peça defeituosa é extraído de uma distribuição diferente, porque uma pessoa a retirou da linha. Este é o bug sério mais comum em machine learning aplicado, e tem nome: target leakage — informação nas features de treino que não estaria disponível no momento em que a previsão tem de ser feita.5 Não lança excepções. Produz um número melhor. Todos os incentivos de um projecto apontam para a manter.
A defesa é uma pergunta, feita a cada coluna: no instante em que preciso desta previsão, este valor já existe? Numa linha em produção, station_seconds é desconhecido até depois de a peça ter sido inspeccionada — que é precisamente a coisa que o modelo devia substituir.
De quantos exemplos de teste preciso?
Ligação para a secção: De quantos exemplos de teste preciso?Imagine que avalia um modelo em 20 exemplos e ele acerta 17. Comunica 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 é algures entre 64 % e 95 %. Um modelo genuinamente de 65 % produz este resultado 4,4 % das vezes — uma execução em vinte e três — e se experimentou um punhado de prompts e comunicou o melhor, fabricou essa execução. Dezassete em vinte não distinguem um modelo de 85 % de um de 65 %.
Duas formas de pôr um intervalo numa taxa, e ambas pertencem ao seu 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)Use Wilson6 para uma taxa de sucesso simples; comporta-se bem em qualquer e não precisa de aleatoriedade. Note acima que em a extremidade superior do bootstrap é 1,0000 — ao reamostrar 20 pontos, é fácil obter 20 correctos, por isso não consegue representar um intervalo mais estreito do que a sua própria granularidade. Use o bootstrap7 quando não há fórmula, que é a maioria dos casos interessantes: F1, macro-médias, BLEU, pass@1, a pontuação de um juiz baseado em rubrica. Nesta linha, o F1 de 0,4122 do modelo ajustado carrega um intervalo bootstrap de [0,3009, 0,5156] — que é o número que deve aparecer no relatório, porque a estimativa pontual sozinha convida a uma comparação que não consegue sustentar.
Mais uma medição, porque muda a forma como deve comparar dois modelos. Dois modelos avaliados 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)Os seus intervalos sobrepõem-se, e a regra popular — barras de erro sobrepostas significam que não há diferença significativa — chamaria a comparação inconclusiva. Não é. Os dois modelos correram nos mesmos exemplos, por isso a quantidade certa é a diferença por exemplo, cujo intervalo é [0.0260, 0.0680], confortavelmente acima de zero. Discordam em apenas 31 de 500 itens, e A vence 27 dessas discordâncias; os exemplos partilhados, fáceis e difíceis, cancelam-se em vez de acrescentarem ruído. Compare modelos de forma emparelhada, e chega à mesma conclusão com uma fracção dos dados.
Para onde isto segue
Ligação para a secção: Para onde isto segueAgora tem um modelo que emite 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 alguma dessas coisas funciona. O intervalo Wilson de dez linhas acima é reutilizado literalmente: transporta as variantes de prompt no Capítulo 15, as tabelas de retrieval no Capítulo 19, e o golden set no Capítulo 29. O bootstrap é aquilo a que recorre quando não existe fórmula.
Mas o modelo ainda é uma camada. Desenha uma linha, e o Capítulo 1 provou com quatro linhas de XOR que uma linha não chega. A correcção é empilhar: uma primeira camada que dobra o espaço, uma segunda que desenha a linha no espaço dobrado.
É aí que o gradient limpo deste capítulo deixa de chegar. Tudo acima funcionou porque podia ser escrito à mão, uma vez, para um modelo com uma camada entre a entrada e a perda. Ponha 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 na saída de todo — um peso cuja influência chega apenas através de outra camada, possivelmente por vários caminhos ao mesmo tempo?
Essa derivada existe. Calculá-la à mão é impossível para qualquer coisa maior do que um brinquedo, e calculá-la um parâmetro de cada vez é impossível noutra escala. O que é necessário é um procedimento que obtenha todas as derivadas na rede a partir de uma única passagem para trás sobre o mesmo grafo que a passagem para a frente acabou de percorrer.
Isso é o Capítulo 5, e é o motor em que o resto deste curso corre.
Fontes e método
Ligação para a secção: Fontes e métodoTambém vale a pena ler em paralelo 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 pela ordem seguida neste capítulo; 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) — porque a AUC citada acima não deve ser o único número independente do threshold que consulta quando 1,7 % das peças estão defeituosas.
Referências
Ligação para a secçã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 a sua ligação canónica, 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, em vez de 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 é sobre validação; a aula 17, sobre os três princípios de aprendizagem, é onde data snooping é nomeado. Entre as duas, são a fonte da disciplina neste capítulo: cada olhar para um conjunto de dados é uma decisão de fitting, tenha ou não corrido um optimizador. ↩
-
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 bias–variance e para reamostragem. O volume complementar é onde a armadilha de selecção é afirmada directamente: 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 ganhas por um modelo que tinha aprendido um artefacto da forma como os dados foram reunidos. ↩
-
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 a predefinição certa para uma proporção. O intervalo de manual é o que deve evitar: dá disparates perto de 0 e 1, e cobre de menos 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 pôr um intervalo em qualquer estatística que consiga calcular, incluindo as que não têm teoria de amostragem. ↩