AdamW é um algoritmo de otimização usado em aprendizado de máquina e aprendizado profundo para treinar redes neurais. É uma variante do otimizador Adam que desacopla a decadência de peso das atualizações adaptativas da taxa de aprendizado. O algoritmo foi introduzido em 2017 por Ilya Loshchilov e Frank Hutter, e tornou-se uma escolha padrão para treinar modelos grandes, incluindo transformadores e modelos de linguagem de grande porte. Ao separar a decadência de peso das atualizações de parâmetros baseadas em gradiente, o AdamW aborda um problema conhecido no otimizador Adam original, onde a decadência de peso era aplicada de uma forma que interferia com as taxas de aprendizado adaptativas, levando a convergência e generalização subótimas.
A principal motivação para o AdamW veio da observação de que, no Adam, o termo de regularização L2 (frequentemente usado como decadência de peso) é dividido pela raiz quadrada da média móvel exponencial dos gradientes ao quadrado. Esse acoplamento faz com que a decadência de peso efetiva varie entre parâmetros e ao longo do tempo, o que pode dificultar a otimização. O AdamW propõe uma correção simples: aplicar a decadência de peso diretamente aos parâmetros após a atualização do gradiente, independente da escala adaptativa da taxa de aprendizado. Esse desacoplamento mostrou-se capaz de melhorar o desempenho do treinamento em uma variedade de tarefas, incluindo classificação de imagens e modelagem de linguagem.
Contexto: Otimização em Aprendizado de Máquina
Treinar uma rede neural envolve minimizar uma função de perda, tipicamente usando variantes da descida do gradiente estocástica (SGD). Na SGD, os parâmetros do modelo são atualizados iterativamente movendo-se na direção do gradiente negativo da perda, calculado em um subconjunto aleatório dos dados de treinamento. O tamanho do passo, ou taxa de aprendizado, controla o quão grande é cada atualização. Ao longo dos anos, muitas melhorias foram propostas para acelerar a convergência e melhorar o desempenho final, como momentum, taxas de aprendizado adaptativas e decadência de peso.
A decadência de peso é uma técnica de regularização que penaliza valores grandes de parâmetros adicionando um termo proporcional à soma dos pesos ao quadrado à função de perda. Na SGD padrão, a decadência de peso é equivalente à regularização L2, mas essa equivalência se quebra em métodos adaptativos como o Adam. O AdamW foi projetado para restaurar o comportamento pretendido da decadência de peso em otimizadores adaptativos.
O Otimizador Adam
O Adam (Estimativa Adaptativa de Momentos) foi introduzido por Diederik Kingma e Jimmy Ba em 2014. Ele mantém taxas de aprendizado por parâmetro ao manter uma média exponencialmente decrescente de gradientes passados (primeiro momento) e gradientes ao quadrado passados (segundo momento). A regra de atualização do Adam é:
\[ \theta_{t+1} = \theta_t - \eta \frac{\hat{m}_t}{\sqrt{\hat{v}_t} + \epsilon} \]
onde \(\hat{m}_t\) e \(\hat{v}_t\) são estimativas corrigidas por viés dos primeiro e segundo momentos, \(\eta\) é a taxa de aprendizado, e \(\epsilon\) é uma constante pequena para estabilidade numérica. O Adam tornou-se popular devido à sua robustez a hiperparâmetros e convergência rápida, especialmente para treinar redes profundas.
No entanto, na implementação original do Adam, a decadência de peso era implementada como regularização L2, que adiciona um termo \(\frac{\lambda}{2} \|\theta\|^2\) à perda. Esse termo é então incluído no gradiente, e como o Adam normaliza o gradiente pelo segundo momento, a decadência de peso efetiva torna-se \(\lambda / \sqrt{\hat{v}_t}\). Isso significa que parâmetros com gradientes grandes recebem menos regularização, o que pode levar a overfitting e má generalização.
O Algoritmo AdamW
O AdamW modifica a regra de atualização removendo o termo de regularização L2 do gradiente e, em vez disso, aplicando a decadência de peso diretamente aos parâmetros após a atualização adaptativa. A regra de atualização torna-se:
\[ \theta_{t+1} = \theta_t - \eta \frac{\hat{m}_t}{\sqrt{\hat{v}_t} + \epsilon} - \eta \lambda \theta_t \]
onde \(\lambda\) é o coeficiente de decadência de peso. Esse desacoplamento garante que a decadência de peso seja aplicada uniformemente a todos os parâmetros, independentemente de suas magnitudes de gradiente. Os autores argumentaram que isso leva a melhor generalização e treinamento mais estável, especialmente ao usar taxas de aprendizado grandes.
Em seu artigo, Loshchilov e Hutter demonstraram que o AdamW supera o Adam com regularização L2 em várias tarefas de referência, incluindo classificação de imagens no CIFAR-10 e modelagem de linguagem no Penn Treebank. Eles também mostraram que o AdamW é mais robusto à escolha da taxa de aprendizado e dos hiperparâmetros de decadência de peso.
Impacto no Aprendizado Profundo
O AdamW teve um impacto significativo no campo da inteligência artificial. Ele agora é o otimizador padrão em muitos frameworks populares de aprendizado profundo, como PyTorch e TensorFlow, e é amplamente usado no treinamento de transformadores e modelos de linguagem de grande porte. Por exemplo, muitos modelos desenvolvidos por organizações como OpenAI, Anthropic e Google DeepMind usam AdamW como parte de seus pipelines de treinamento.
A decadência de peso desacoplada no AdamW tem sido particularmente benéfica para treinar modelos em grande escala, onde a regularização é crucial para prevenir overfitting em conjuntos de dados massivos. Também foi mostrado que melhora a velocidade de convergência e o desempenho final em comparação com outros otimizadores como SGD com momentum ou Adam com regularização L2.
Comparação com Outros Otimizadores
O AdamW é frequentemente comparado com outros otimizadores adaptativos como AdaGrad, RMSProp e Adam. Enquanto AdaGrad e RMSProp ajustam as taxas de aprendizado com base em gradientes históricos, o Adam combina tanto momentum quanto taxas de aprendizado adaptativas. O AdamW melhora o Adam ao corrigir o problema da decadência de peso, tornando-o mais confiável para uma ampla gama de tarefas.
Outro otimizador relacionado é a SGD com momentum, que é mais simples, mas frequentemente requer ajuste cuidadoso do cronograma de taxa de aprendizado. O AdamW oferece um bom equilíbrio entre facilidade de uso e desempenho, razão pela qual se tornou uma escolha preferida para muitos profissionais. No entanto, alguns estudos mostraram que a SGD com momentum pode alcançar melhor generalização em certas tarefas se devidamente ajustada, mas o AdamW permanece competitivo e mais robusto a escolhas de hiperparâmetros.
Considerações Práticas
Ao usar AdamW, há algumas considerações práticas. O coeficiente de decadência de peso \(\lambda\) é tipicamente definido para um valor pequeno, como 0,01 ou 0,1, mas pode precisar ser ajustado para tarefas específicas. A taxa de aprendizado é frequentemente definida para um valor como 1e-4 ou 3e-4 para treinar transformadores. Além disso, o AdamW frequentemente se beneficia de um cronograma de aquecimento da taxa de aprendizado, onde a taxa aumenta gradualmente de um valor pequeno para o valor alvo ao longo dos primeiros milhares de passos.
Na prática, o AdamW é implementado na maioria das bibliotecas de aprendizado profundo, então os usuários podem simplesmente especificar o otimizador e definir o parâmetro de decadência de peso. Por exemplo, em PyTorch, pode-se usar torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0.01). Essa simplicidade contribuiu para sua adoção generalizada.
Extensões e Variantes
Várias extensões e variantes do AdamW foram propostas. Por exemplo, o AdamW com decadência de peso desacoplada foi combinado com outras técnicas como Lookahead, que mantém um conjunto mais lento de parâmetros para atualizações mais suaves. Outra variante é o AdamW com cronogramas de taxa de aprendizado de anelamento de cosseno, que mostrou melhorar o desempenho em tarefas de classificação de imagens.
No contexto de modelos de linguagem de grande porte, o AdamW é frequentemente usado com treinamento de precisão mista e acumulação de gradiente para lidar com tamanhos de lote grandes. Alguns frameworks também oferecem implementações fundidas do AdamW que reduzem o uso de memória e melhoram a eficiência computacional, o que é importante ao treinar modelos com bilhões de parâmetros.
Conclusão
O AdamW tornou-se uma ferramenta fundamental no kit de ferramentas de aprendizado de máquina. Ao desacoplar a decadência de peso das atualizações adaptativas de gradiente, ele aborda uma falha sutil, mas importante, no otimizador Adam original, levando a melhor generalização e treinamento mais estável. Sua simplicidade e eficácia tornaram-no a escolha padrão para muitas aplicações de aprendizado profundo, desde classificação de imagens até processamento de linguagem natural. À medida que o campo continua a evoluir, o AdamW permanece um método de otimização confiável e amplamente usado.
Referências
- Loshchilov, I., & Hutter, F. (2017). Decoupled Weight Decay Regularization. arXiv:1711.05101.
- Kingma, D. P., & Ba, J. (2014). Adam: A Method for Stochastic Optimization. arXiv:1412.6980.