Saltar para o conteúdo
4/30Capítulo 4 de 30

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 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.

belt.pyPYTHON
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:]
TEXT
N = 20000  defects = 337  base rate = 0.0169
defects per split = 203 60 74

Trê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.

O 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 pp de a peça estar defeituosa, e

P(yp)=py(1p)1yP(y \mid p) = p^{\,y}\,(1-p)^{\,1-y}

que é apenas uma forma compacta de escrever «pp se y=1y = 1, e 1p1-p se y=0y = 0». Calcule o logaritmo disso e negue-o, e a perda para um exemplo é

L=[ylogp+(1y)log(1p)]L = -\big[\,y \log p + (1 - y)\log(1 - p)\,\big]

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 pp. O modelo calcula uma soma ponderada s=wx+bs = \mathbf{w}\cdot\mathbf{x} + b, que é um número real e percorre toda a recta, e uma probabilidade tem de viver em (0,1)(0,1). A função que faz a passagem entre ambos é a sigmoide logística:

σ(s)=11+es\sigma(s) = \frac{1}{1 + e^{-s}}
TEXT
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

Leia 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 é log2\log 2, 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 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 s=wx+bs = \mathbf{w}\cdot\mathbf{x} + b e p=σ(s)p = \sigma(s):

Ls=py,Lw=(py)x,Lb=py\frac{\partial L}{\partial s} = p - y, \qquad \frac{\partial L}{\partial \mathbf{w}} = (p - y)\,\mathbf{x}, \qquad \frac{\partial L}{\partial b} = p - y
Mostrar detalhes

As duas linhas que fazem a confusão cancelar. A sigmoide tem uma derivada invulgarmente agradável, σ(s)=σ(s)(1σ(s))=p(1p)\sigma'(s) = \sigma(s)\,(1 - \sigma(s)) = p(1-p). E a perda diferencia-se para

Lp=yp+1y1p=pyp(1p)\frac{\partial L}{\partial p} = -\frac{y}{p} + \frac{1-y}{1-p} = \frac{p - y}{p\,(1-p)}

Multiplique as duas pela regra da cadeia e p(1p)p(1-p) aparece uma vez em cima e uma vez em baixo. Cancela exactamente, e pyp - y é 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:

logistic.pyPYTHON
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, b

O np.where em sigmoid não é cosmético. Calcular 1/(1+es)1/(1+e^{-s}) directamente faz overflow para ss 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 gradient

A 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 L=(py)2L = (p - y)^2 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á

Ls=2(py)p(1p)\frac{\partial L}{\partial s} = 2\,(p - y)\,p\,(1-p)

Esse p(1p)p(1-p) 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 ssppentropia cruzada L/s\partial L/\partial serro quadrático L/s\partial L/\partial srazão
8-80,0003350.999665-0.9996650.000670-0.0006701.491
4-40,0179860.982014-0.9820140.034690-0.03469028,3
2-20,1192030.880797-0.8807970.184956-0.1849564,8
000,5000000.500000-0.5000000.250000-0.2500002,0
+2+20,8807970.119203-0.1192030.025031-0.0250314,8

Em s=8s = -8 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 1-1: 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 (w=[6,6]\mathbf{w} = [-6, -6]), 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.

épocaperda de entropia cruzadaexatidãoperda de erro quadráticoexatidão
15,48650,23005,94990,2290
101,55250,24605,90420,2290
500,46420,77805,69130,2320
1000,46390,77705,39550,2410
2000,46390,77704,63110,2745
5000,46390,77700,52910,7660
1.0000,46390,77700,46400,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.

Trê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:

H(p)=ipilog2piH(p) = -\sum_i p_i \log_2 p_i

Entropia cruzada é o que gasta quando usa um código construído para qq em dados que, na verdade, vêm de pp:

H(p,q)=ipilog2qiH(p, q) = -\sum_i p_i \log_2 q_i

Divergência KL é o excesso — o desperdício, em bits, causado por acreditar em qq quando a verdade é pp:

DKL(pq)=H(p,q)H(p)D_{\mathrm{KL}}(p \parallel q) = H(p,q) - H(p)

Verifique as três na linha:

TEXT
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

Duas 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. H+DKL=H(p,q)H + D_{\mathrm{KL}} = H(p,q) 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 nada

Defeituosa 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:

softmax(z)i=ezijezj\operatorname{softmax}(\mathbf{z})_i = \frac{e^{z_i}}{\sum_j e^{z_j}}

Tem uma propriedade que parece um acidente e é, na verdade, toda a implementação:

softmax(z+c)=softmax(z)\operatorname{softmax}(\mathbf{z} + c) = \operatorname{softmax}(\mathbf{z})

para qualquer constante cc, porque ezi+c=ecezie^{z_i + c} = e^{c} e^{z_i} e o ece^c 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:

TEXT
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

e800e^{800} 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 e0=1e^0 = 1. É o truque logsumexp do Capítulo 2 vestido com roupa de trabalho, e todas as implementações sérias o fazem:

softmax.pyPYTHON
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, b

O gradient é outra vez previsão menos verdade, agora com YY 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 →okshort shotflashburnrevocação
ok8505900,984
short shot2221000,488
flash2003010,588
burn300390,929
precisão0,9500,8080,7690,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.

Pegue 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:

modeloexatidãoentropia cruzadaperda média quando acertaperda média quando errapior perda individual
hesitante (logits × 0,35)0,98300,15490,13691,19902,80
como treinado0,98300,05640,01472,46897,82
excessivamente confiante (logits × 4)0,98300,15630,00099,142727,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ã.

Antes de qualquer modelo, o requisito: quanto pontua a resposta mais preguiçosa possível? Nesta linha, dizer sempre que está boa:

TEXT
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:

TEXT
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 boaprevisto defeituosa
realmente boa3.9242
realmente defeituosa668

Três números dão nome às três formas de ler essa tabela:

  • Precisão =TP/(TP+FP)=8/10=0.800= \mathrm{TP}/(\mathrm{TP}+\mathrm{FP}) = 8/10 = 0.800. Das peças que assinalou, quantas eram realmente defeituosas. Este é o custo de inspecções desperdiçadas.
  • Revocação =TP/(TP+FN)=8/74=0.108= \mathrm{TP}/(\mathrm{TP}+\mathrm{FN}) = 8/74 = 0.108. Das peças defeituosas, quantas apanhou. Este é o custo de enviar uma peça má para um cliente.
  • F1 =2PR/(P+R)=0.190= 2PR/(P+R) = 0.190, 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:

thresholdTPFPFNexatidãoprecisãorevocaçãoF1
0,50082660,98300,8000,1080,190
0,2002728470,98120,4910,3650,419
0,10042118320,96250,2630,5680,359
0,05054236200,93600,1860,7300,297
0,0206757070,85580,1050,9050,188
0,005711.36030,65930,0500,9590,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 positivosexatidãoprecisãorevocaçãoF1AUC
10,98300,8000,1080,1900,9363
100,96050,2530,5810,3520,9361
600,82900,0910,9190,1660,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.

Porquê 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:

modeloexatidãoprecisãorevocaçãoF1entropia cruzadaAUC
largura + peso0,98300,8000,1080,1900,05640,9363
+ station_seconds0,99200,7920,7700,7810,02360,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.

  1. 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.

  2. Treine um modelo por feature, isoladamente. Qualquer coisa que transporte a resposta anunciar-se-á:

    feature isoladaexatidãorevocaçãoF1AUC
    largura0,98150,0140,0260,8691
    peso0,98150,0000,0000,7914
    station_seconds0,98500,4050,5000,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.

  3. 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.

the planted leakPYTHON
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.

Imagine que avalia um modelo em 20 exemplos e ele acerta 17. Comunica 85 %.

TEXT
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

A 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:

uncertainty.pyPYTHON
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 nn e não precisa de aleatoriedade. Note acima que em n=20n = 20 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:

TEXT
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.

Agora 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 L/s=py\partial L/\partial s = p - y 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.


També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.

  1. Ma, T. e Ng, A. CS229 Lecture Notes, Stanford University, capítulos 2 e 3. Onde o cancelamento que produz pyp - y 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.

  2. 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.

  3. 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.

  4. 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.

  5. 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.

  6. 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 p^±zp^(1p^)/n\hat{p} \pm z\sqrt{\hat{p}(1-\hat{p})/n} é o que deve evitar: dá disparates perto de 0 e 1, e cobre de menos em nn pequeno.

  7. 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.

Pronto para deixar a LIA escolher?

Construa com todos os modelos de IA num só sítio — comece grátis hoje.