Pular para o conteúdo
6/30Capítulo 6 de 30

Fazer a rede treinar — e fazê-la generalizar

Uma rede de seis camadas com perda travada em ln 2, corrigida uma medida por vez. Depois, double descent: 5.000 parâmetros em 40 pontos.

Nesta página

A rede do Capítulo 5 funciona. Ela tem nove parâmetros, aprende XOR, e seus gradientes batem com o PyTorch até dezesseis casas decimais.

Torne-a profunda, com seis camadas, e ela para de aprender por completo. Não devagar — por completo. Aqui está uma rede de seis camadas em um problema de classificação de duas espirais, treinada por 5000 passos:

TEXT
step    1: loss 0.693147
step 5000: loss 0.693147
accuracy: 50.0 %

Esse número não é arbitrário. ln2=0.693147\ln 2 = 0.693147 é a entropia cruzada binária de um modelo que produz probabilidade 0.50.5 para tudo, e 50% é um cara ou coroa em um conjunto de dados balanceado. Depois de cinco mil passos, a rede não moveu nem um único dígito. Nada quebrou, nada avisou, e os gradientes ainda estão exatamente certos.

Este capítulo é sobre a lacuna entre uma rede que executa e uma rede que funciona. Ele tem duas metades que parecem assuntos diferentes e são o mesmo trabalho: fazer a perda cair, e fazer com que ela caia em dados que o modelo nunca viu.

Comece olhando, em vez de adivinhar. Passe um batch de entradas pela rede e imprima o desvio padrão das ativações em cada camada, e depois o desvio padrão dos gradientes dos pesos:

profile.pyPYTHON
def profile(model, x):
    h = x
    for layer in model:
        h = layer(h)
        if isinstance(layer, (nn.Tanh, nn.ReLU)):
            print(f"activation std: {h.std().item():.4f}")
    model(x).sum().backward()
    for p in model.parameters():
        if p.dim() == 2:
            print(f"gradient std: {p.grad.std().item():.2e}")

Três inicializações, mesma arquitetura, seis camadas de tanh\tanh:

inicializaçãodesvio padrão da ativação, camadas 1→6
normal, std 0.010.010.0145 · 0.0016 · 0.0002 · 0.0000 · 0.0000 · 0.0000
normal, std 110.6573 · 0.9296 · 0.9585 · 0.9634 · 0.9637 · 0.9625
Xavier0.1579 · 0.1493 · 0.1353 · 0.1333 · 0.1325 · 0.1403
inicializaçãodesvio padrão do gradiente, primeira camada → última
normal, std 0.010.013.20e-06 · 4.97e-07 · … · 6.40e-06
normal, std 111.94e+03 · 2.28e+02 · 1.22e+02 · 4.43e+01 · 1.85e+01 · 7.30e+00
Xavier2.31e+00 · 4.50e-01 · 4.26e-01 · 3.89e-01 · 4.39e-01 · 4.73e-01

A primeira linha é a rede acima, e ela não está aprendendo devagar — não resta sinal nenhum. Na camada quatro, o desvio padrão da ativação sofreu underflow para zero em quatro casas decimais. Toda entrada produz a mesma saída, a saída é uma constante, e o gradiente de uma constante é nada. Os pesos foram inicializados pequenos "por segurança", e pequeno foi fatal.

A segunda linha é a falha oposta, e vale entendê-la porque ela é contraintuitiva. As ativações parecem saudáveis — em torno de 0,96 —, mas isso é tanh\tanh saturada, presa perto do seu limite, exatamente o regime que o Capítulo 5 mediu como perdendo um fator de quase dez mil no gradiente. E, ainda assim, os gradientes são enormes: 1940 na primeira camada. As duas coisas são verdadeiras ao mesmo tempo. Cada passo para trás multiplica por WW^\top, e com 128 entradas em variância unitária esse fator tem um ganho de cerca de 12811\sqrt{128} \approx 11, o que sobrecarrega o encolhimento da tanh\tanh saturada. Os gradientes crescem geometricamente no caminho de volta. Esse é o gradiente explosivo, e ele produz valores de perda de nan em poucos passos em qualquer execução de treinamento real.

A terceira linha é o que você quer: ativações com escala aproximadamente constante ao longo da profundidade, gradientes com escala aproximadamente constante ao longo da profundidade. Nada morre, nada explode.

Inicializar bem corrige a escala no passo zero. Não a mantém fixa: os pesos se movem, e no passo cinco mil o argumento cuidadoso da variância já não se aplica.

Camadas de normalização impõem a escala continuamente. Dado um vetor de ativações, subtraia uma média, divida por um desvio padrão, depois aplique uma escala aprendida γ\gamma e um deslocamento β\beta para que a camada possa desfazer a normalização se isso acabar sendo o que ela quer:

h^=hμσ2+ϵ,y=γh^+β\hat{h} = \frac{h - \mu}{\sqrt{\sigma^2 + \epsilon}}, \qquad y = \gamma\hat{h} + \beta

A única pergunta real é sobre o que você faz a média. Batch normalization3 calcula μ\mu e σ\sigma ao longo da dimensão do batch, uma estatística por feature. Layer normalization4 calcula essas estatísticas ao longo das features, uma estatística por exemplo.

Essa escolha parece pequena e decide quase tudo depois:

BatchNorm faz a saída de cada exemplo depender dos outros exemplos que por acaso estavam no seu batch. Durante o treinamento, isso é um regularizador leve. Na inferência, não há batch, então ela precisa manter uma média móvel das estatísticas coletadas durante o treinamento — o que significa que a camada se comporta de forma diferente nos modos de treinamento e avaliação, e esquecer de trocar os modos é um dos bugs mais comuns na área. Ela também degrada com batches pequenos, e é estranha com sequências de comprimento variável, porque "a média sobre o batch na posição 40" é calculada a partir de quantas sequências por acaso forem longas o suficiente.

LayerNorm normaliza cada exemplo por conta própria. Sem dependência de batch, sem estatísticas móveis, comportamento idêntico em treinamento e inferência, indiferente ao tamanho do batch, indiferente ao comprimento da sequência. Cada uma dessas propriedades vira requisito, não luxo, quando você está gerando um token por vez para um usuário, que é onde o Capítulo 13 chega.

É por isso que LayerNorm é a que você encontrará de novo no Capítulo 9 sem alterações: o bloco transformer a usa, e a usa pelos motivos da coluna da direita, não porque ela funciona melhor em abstrato.

Corrigir uma coisa por vez, que é a skill de verdade

Link para a seção: Corrigir uma coisa por vez, que é a skill de verdade

Quatro possíveis correções para a rede morta: inicialização Xavier, LayerNorm, conexões residuais e Adam em vez de SGD. A tentação é aplicar as quatro e seguir em frente. Faça isso e você nunca saberá qual delas importou, e na próxima vez que acontecer você não terá método — só um ritual.

Então aplique uma por vez. Mesma seed, mesmos dados, mesma arquitetura, 800 passos:

o que foi adicionadoperda finalacurácia
nada0.693150,0%
inicialização Xavier0.569260,4%
LayerNorm0.623061,5%
conexões residuais0.665156,6%
Adam0.678758,7%
as quatro0.0000100,0%

Leia essa tabela como você a leria às 2 da manhã, e a conclusão é: nada funciona sozinho, tudo funciona junto, portanto deep learning é alquimia. Essa conclusão está errada, e descobrir por quê é a coisa mais útil deste capítulo.

Dê a cada execução seis vezes o orçamento — 5000 passos em vez de 800 — e tudo muda:

o que foi adicionadoperda final @ 5000acurácia
nada0.693150,0%
inicialização Xavier0.0007100,0%
LayerNorm0.0002100,0%
conexões residuais0.665356,7%
Adam0.690853,4%
Xavier + Adam0.0000100,0%
Xavier + LayerNorm0.0001100,0%

Agora o quadro é nítido, e é um diagnóstico em vez de um ritual.

A inicialização sozinha corrige. A normalização sozinha corrige. Cada uma trata a doença real — o sinal para frente colapsando para zero — e qualquer uma das duas é suficiente. Em 800 passos, elas apenas pareciam crédito parcial, porque tinham resolvido o problema e ainda estavam saindo do buraco.

Conexões residuais e Adam não corrigem, em nenhum orçamento. Não porque sejam ruins, mas porque tratam outra doença. Uma conexão residual dá ao gradiente um caminho em volta de uma camada bloqueante; isso vale muito quando o gradiente é o problema, e nada quando o sinal para frente já é zero, porque um atalho em volta de uma camada morta ainda carrega um valor morto. Adam redimensiona o passo de cada parâmetro pelo seu próprio histórico de gradientes; isso ajuda quando gradientes têm magnitudes muito diferentes, e não consegue ressuscitar uma rede cuja saída não depende da entrada.

E "nada" continua sendo exatamente 0.6931 depois de cinco mil passos. Não 0.6929. Não é lenta; está morta, e essa distinção fica visível de um jeito que antes não ficava, porque você tem a linha que diz que uma correção funciona para comparar.

Daqui em diante, este curso usa PyTorch. Isso deve ser merecido, não apenas anunciado, então aqui está exatamente o que ele faz e que você já sabe fazer.

Um otimizador é uma regra para transformar gradientes em atualizações de parâmetros. O gradient descent simples usa o gradiente. Momentum usa uma média móvel dele, o que suaviza o ruído e ganha velocidade nas direções que permanecem consistentes:

optim_by_hand.pyPYTHON
v = beta * v + p.grad          
p -= lr * v                    

Adam5 mantém duas médias móveis — do gradiente e do gradiente ao quadrado — e divide uma pela raiz quadrada da outra, de modo que cada parâmetro recebe um passo escalado para a magnitude recente do seu próprio gradiente:

optim_by_hand.pyPYTHON
m = b1 * m + (1 - b1) * g          # mean of the gradient          
v = b2 * v + (1 - b2) * g * g      # mean of the squared gradient  
m_hat = m / (1 - b1 ** t)          # bias correction: both averages start at zero
v_hat = v / (1 - b2 ** t)
p -= lr * m_hat / (v_hat.sqrt() + eps)   

Dez linhas. Rode ambos contra torch.optim no mesmo problema por 50 passos:

TEXT
SGD+momentum   by hand [2.7781870365142822, -1.0304985046386719]
               torch   [2.7781870365142822, -1.0304983854293823]   max |diff| = 1.19e-07
Adam           by hand [0.4893140196800232, -0.46317872405052185]
               torch   [0.48931416869163513, -0.46317875385284424]   max |diff| = 1.49e-07

Idênticos até a precisão de float32. torch.optim.Adam são aquelas cinco linhas, mais décadas de cuidado com casos de borda e um kernel em C++. Essa é a troca que você fará daqui em diante: não magia no lugar do entendimento, mas velocidade no lugar de linhas que você já escreveu.

A explicação usual de Adam é "taxas de aprendizado adaptativas por parâmetro", o que é uma descrição, não um motivo. O motivo é geometria, e ela pode ser medida.

Pegue uma perda cuja curvatura difere entre direções: íngreme em uma, rasa em outra. SGD tem uma única taxa de aprendizado global, então precisa escolher um valor pequeno o bastante para ser estável na direção mais íngreme — e esse valor então é pequeno demais para a direção rasa, onde o progresso se arrasta. É isso que causa a imagem clássica do gradient descent ziguezagueando por um vale estreito.

Duas razões de curvatura, três otimizadores, 300 passos, e cada otimizador recebendo a melhor taxa de aprendizado de uma varredura para que ninguém fique em desvantagem:

razão de curvaturaSGDSGD + momentumAdam
10 : 1erro 0.000002erro 0.000000erro 0.000000
1000 : 1erro 1.925485erro 0.001432erro 0.000000
divergiu em (1000:1)4 de 8 taxas4 de 8 taxas0 de 6 taxas

Com uma razão de dez, tudo funciona e não há o que discutir. Com mil, SGD simples não consegue chegar à resposta em nenhuma taxa de aprendizado testada — seu melhor resultado ainda é um erro de 1,93 — e diverge diretamente em metade das taxas. Adam chega exatamente ao alvo e não diverge em nenhuma delas.

Essa última coluna é o motivo prático pelo qual Adam é o padrão. Não é que Adam encontre soluções melhores; em problemas bem condicionados, SGD ajustado muitas vezes iguala ou supera. É que Adam é muito menos sensível à taxa de aprendizado que você escolheu, e redes reais têm razões de curvatura muito piores do que mil entre seus milhões de parâmetros.

Mais duas peças pertencem aqui, e ambas cabem em uma linha. Gradient clipping redimensiona o vetor de gradiente sempre que sua norma excede um limite, o que transforma a linha "perda de repente salta para um valor enorme" da tabela de diagnóstico em um não evento. E cronogramas de taxa de aprendizado: um warmup curto a partir de quase zero nos primeiros poucos centenas de passos, porque as estimativas de variância de Adam são lixo até terem visto alguns gradientes, e um passo de tamanho cheio dado sobre lixo pode destruir uma inicialização; depois cosine decay em direção a zero, porque terminar uma execução com o mesmo tamanho de passo com que você começou significa tremer em torno do mínimo em vez de se assentar nele.

A segunda metade: o modelo que ajusta perfeitamente e não prevê nada

Link para a seção: A segunda metade: o modelo que ajusta perfeitamente e não prevê nada

Tudo até aqui foi sobre fazer a perda cair. Agora vem a metade mais difícil, porque a perda cair não é o objetivo — é uma aproximação do objetivo, e essa aproximação falha de uma forma específica e famosa.

Doze pontos de uma função suave com um pouco de ruído. Ajuste polinômios de grau crescente:

grauRMSE de treinoRMSE de teste
10.7644990.6985
30.2526050.3031
50.1644370.1568
90.0889600.2347
110.0000001.2094

Grau 11 passando por 12 pontos atravessa cada um deles exatamente — erro de treino zero em seis casas decimais — e é oito vezes pior que o grau 5 em dados que não viu. Peça ao grau 3 e ao grau 11 que prevejam em x=3.25x = 3.25, logo fora do intervalo de treinamento:

TEXT
degree  3: predicts   -1.053   (truth -0.012)
degree 11: predicts  +61.224   (truth -0.012)

Sessenta e um, onde a resposta é aproximadamente zero. O modelo não aprendeu a função; aprendeu os doze pontos, e entre eles faz o que a aritmética exigir.

Isso é overfitting, e seu oposto — grau 1, que não consegue representar a curva de jeito nenhum e é ruim em todos os lugares — é underfitting. A explicação clássica divide o erro esperado de um modelo em três partes: viés, o erro por o modelo ser rígido demais para representar a verdade; variância, o erro por o modelo ser flexível demais e perseguir o ruído nesta amostra específica; e ruído irredutível, que nada corrige. Modelos simples são enviesados, modelos flexíveis têm alta variância, e a prescrição clássica é encontrar o ponto ideal no meio — grau 5 na tabela acima.

As ferramentas padrão atacam o termo da variância:

  • Regularização L2 (weight decay) adiciona λw2\lambda \lVert w \rVert^2 à perda, puxando pesos em direção a zero e tornando a função mais suave. Na tabela acima, o maior coeficiente do grau 11 causa o estrago; penalizar o tamanho o desarma.
  • L1 adiciona λwi\lambda \sum |w_i| em vez disso. A diferença não é cosmética: o gradiente de L2 é proporcional ao peso e, portanto, encolhe conforme o peso encolhe, aproximando-se de zero sem chegar lá, enquanto o gradiente de L1 é uma constante ±λ\pm\lambda que continua empurrando até o fim. Portanto, L1 produz pesos que são exatamente zero — ela seleciona features. L2 produz pesos pequenos. Use L2 quando quiser suavidade, L1 quando quiser esparsidade.
  • Dropout7 zera um subconjunto aleatório de ativações em cada passo de treinamento, para que nenhuma unidade possa depender de uma unidade específica estar presente.
  • Early stopping observa a perda de validação e para quando ela começa a subir.
  • Data augmentation fabrica mais exemplos de treinamento a partir dos que você tem, atacando o problema na fonte: overfitting é falta de dados tanto quanto é excesso de parâmetros.
  • Cross-validation divide os dados de kk maneiras e treina kk vezes, o que compra uma estimativa confiável do erro de teste quando você tem dados demais em falta para separar um conjunto retido.

Double descent, ou por que a seção anterior não é a história toda

Link para a seção: Double descent, ou por que a seção anterior não é a história toda

Agora o fato que quebra o quadro.

A história de viés-variância diz que, depois do ponto ideal, mais parâmetros significam pior generalização. Modelos de linguagem modernos têm muito mais parâmetros do que as regras clássicas permitiriam para os dados que veem, e generalizam de forma excelente. As duas afirmações são verdadeiras, e conciliá-las é a coisa mais útil deste capítulo.

Quarenta pontos de treinamento, entradas de vinte dimensões, features ReLU aleatórias, e o número de features PP varrido de 2 a 5000 — com a solução de norma mínima escolhida sempre que houver muitas que se ajustem:

PPP/nP/nRMSE de treinoRMSE de testew\lVert w \rVert
100.250.88221.25201.89
200.500.59621.16342.59
300.750.38961.53234.15
380.950.17693.716310.25
401.000.00005.814014.83
421.050.00003.16239.35
601.500.00001.10582.78
2005.000.00000.66380.98
150037.500.00000.58590.33
5000125.000.00000.56640.18

Leia em três partes. Até P/n=0.5P/n = 0.5, a história clássica vale exatamente: o erro cai, depois começa a subir. Em P=n=40P = n = 40 — o limiar de interpolação, onde o modelo tem exatamente parâmetros suficientes para passar por todos os pontos de treinamento — o erro de teste atinge o pico, em 5,81, cinco vezes pior que o modelo pequeno. Esse pico é o alerta clássico, e ele é real.

Então ele desce de novo. E continua descendo, passando de P=5nP = 5n, passando de P=37nP = 37n, até P=125nP = 125n, onde o erro de teste de 0,5664 é melhor do que o melhor modelo subparametrizado já conseguiu. Um modelo com 5000 parâmetros ajustado a 40 pontos é o melhor modelo da tabela.

Isso é double descent,89 e o mecanismo está visível na última coluna. Quando P>nP > n, há infinitas configurações de parâmetros que ajustam os dados de treinamento exatamente, e qual delas você obtém depende de como escolhe. A solução de norma mínima escolhe a menor, e w\lVert w \rVert mostra o que isso significa: ela atinge o pico em 14,83 exatamente no limiar — onde há exatamente uma solução interpoladora e você fica preso a ela, por mais extrema que seja — e depois cai monotonicamente conforme PP cresce, porque mais parâmetros significam mais soluções interpoladoras para escolher, o que significa que a menor disponível fica menor. Em P=5000P = 5000, a norma é 0,18, oitenta vezes menor do que no limiar.

Então os parâmetros extras não estão adicionando complexidade. Estão adicionando escolha, e a regra de seleção gasta essa escolha em simplicidade. A regularização não está na função de perda; está no algoritmo. Gradient descent a partir de uma inicialização pequena tem um viés documentado em direção a soluções de norma pequena, e é por isso que esse comportamento aparece em redes reais treinadas da forma comum, não apenas na álgebra linear acima.

A consequência prática, da qual o Capítulo 10 depende: "o modelo tem mais parâmetros do que dados, então vai sofrer overfit" não é um argumento válido. Era uma boa regra quando os modelos viviam à esquerda do limiar. Tudo que é interessante agora vive muito à direita dele, onde a regra se inverte.

As ferramentas deste capítulo bastam para treinar uma rede que funciona em dados que você consegue colocar em uma tabela: linhas de números, uma coluna de rótulos.

Linguagem não é isso. Antes que um modelo possa prever a próxima palavra, algo precisa decidir o que uma "palavra" sequer é — e a resposta não são letras nem palavras, mas um vocabulário que o modelo aprende a partir dos bytes brutos dos dados de treinamento. Essa decisão, tomada uma vez antes do treinamento começar, determina quantas coisas o modelo pode dizer, quanto custa uma solicitação, e por que modelos que conseguem passar em uma prova de direito não conseguem contar de forma confiável as letras em strawberry.

O Capítulo 7 constrói um tokenizer.


Para as conexões residuais usadas acima, He et al., Deep Residual Learning for Image Recognition (arXiv:1512.03385). Building makemore Part 3: Activations & Gradients, BatchNorm, de Andrej Karpathy, percorre o diagnóstico por histograma de ativações em um modelo real e é o melhor tratamento prático da primeira metade deste capítulo. As aulas 8 e 11–13 de Learning From Data, de Yaser Abu-Mostafa, apresentam corretamente a teoria clássica de generalização, incluindo as partes que este capítulo comprimiu em um parágrafo.

  1. Glorot, X. and Bengio, Y. Understanding the difficulty of training deep feedforward neural networks. AISTATS (2010). O argumento de preservação de variância reproduzido na caixa acima.

  2. He, K., Zhang, X., Ren, S. and Sun, J. Delving Deep into Rectifiers: Surpassing Human-Level Performance on ImageNet Classification. arXiv:1502.01852 (2015).

  3. Ioffe, S. and Szegedy, C. Batch Normalization: Accelerating Deep Network Training by Reducing Internal Covariate Shift. arXiv:1502.03167 (2015). Observe que a explicação de "internal covariate shift" no título foi substancialmente contestada desde então; a camada funciona, mas a explicação original do motivo é disputada.

  4. Ba, J. L., Kiros, J. R. and Hinton, G. E. Layer Normalization. arXiv:1607.06450 (2016).

  5. Kingma, D. P. and Ba, J. Adam: A Method for Stochastic Optimization. arXiv:1412.6980 (2014).

  6. Loshchilov, I. and Hutter, F. Decoupled Weight Decay Regularization. arXiv:1711.05101 (2017).

  7. Srivastava, N., Hinton, G., Krizhevsky, A., Sutskever, I. and Salakhutdinov, R. Dropout: A Simple Way to Prevent Neural Networks from Overfitting. JMLR 15, pp. 1929–1958 (2014).

  8. Belkin, M., Hsu, D., Ma, S. and Mandal, S. Reconciling modern machine-learning practice and the classical bias–variance trade-off. PNAS 116(32), pp. 15849–15854 (2019). O artigo que deu nome ao fenômeno.

  9. Nakkiran, P., Kaplun, G., Bansal, Y., Yang, T., Barak, B. and Sutskever, I. Deep Double Descent: Where Bigger Models and More Data Hurt. arXiv:1912.02292 (2019). Mostra o efeito em redes profundas reais, e ao longo do eixo de tempo de treinamento, além do eixo de tamanho do modelo.

Pronto para deixar a LIA escolher por você?

Crie com todos os modelos de IA em um só lugar — comece grátis hoje.