- Mostra que substituir Layer Norm/RMSNorm, usados quase como obrigatórios em Transformers, por Dynamic Tanh (DyT) pode alcançar desempenho semelhante ou melhor do que modelos tradicionais com normalização
- DyT é uma operação elemento a elemento na forma
DyT(x) = tanh(αx), surgida da observação de que a Layer Normalization dentro de Transformers frequentemente cria um mapeamento de entrada e saída em forma de S, parecido com tanh
- A implementação cabe em poucas linhas de PyTorch e aplica escala e viés à saída de
tanh(alpha * x) com alpha, weight e bias treináveis
- A avaliação cobre amplamente modelagem de visão, linguagem, fala e sequências de DNA, incluindo ViT, ConvNeXt, MAE, DINO, DiT, LLaMA, wav2vec 2.0, HyenaDNA e Caduceus
- Mesmo sem ajuste adicional de hiperparâmetros, os resultados em várias configurações foram semelhantes ou melhores que os de modelos equivalentes baseados em normalização, levando a repensar a premissa de que camadas de normalização são indispensáveis
Onde o Dynamic Tanh muda o jogo
- DyT é uma camada simples que substitui Layer Norm ou RMSNorm em blocos Transformer
- A operação central é
DyT(x) = tanh(αx), aplicada elemento a elemento
- Mostra que Transformers sem camadas de normalização podem ter desempenho semelhante ou até superior ao de Transformers tradicionais com normalização
- A ideia parte da observação de que a Layer Normalization em Transformers frequentemente produz uma relação entrada-saída parecida com uma função tanh escalada
Como é implementado
- O módulo DyT pode ser implementado de forma curta em PyTorch
class DyT(nn.Module):
def __init__(self, num_features, alpha_init_value=0.5):
super().__init__()
self.alpha = nn.Parameter(torch.ones(1) * alpha_init_value)
self.weight = nn.Parameter(torch.ones(num_features))
self.bias = nn.Parameter(torch.zeros(num_features))
def forward(self, x):
x = torch.tanh(self.alpha * x)
return x * self.weight + self.bias
alpha é um parâmetro treinável, com valor inicial definido como 0.5
weight e bias também são parâmetros treináveis, aplicados à saída de tanh(alpha * x)
Observação vinda da Layer Normalization
- A Layer Normalization (LN) em Transformers gera um mapeamento de entrada e saída próximo de uma função
tanh escalada
- Nas camadas iniciais, esse mapeamento tende a ser mais próximo de linear
- À medida que as camadas se aprofundam, a curva em S característica da função
tanh aparece com mais clareza
- Os alvos da observação incluem camadas LN selecionadas de Vision Transformer (ViT), do modelo de fala wav2vec 2.0 e de Diffusion Transformer (DiT)
Escopo da avaliação e resultados
- O DyT foi avaliado em várias arquiteturas e tarefas
- Visão com aprendizado supervisionado: ViT, ConvNeXt
- Visão com aprendizado autossupervisionado: MAE, DINO
- Modelos de difusão: DiT
- Grandes modelos de linguagem: LLaMA
- Fala com aprendizado autossupervisionado: wav2vec 2.0
- Modelagem de sequências de DNA: HyenaDNA, Caduceus
- Em todos os casos, Transformers com DyT mostraram desempenho semelhante ou melhor que seus equivalentes baseados em normalização
- O escopo da avaliação é amplo, cobrindo de reconhecimento a geração, de aprendizado supervisionado a autossupervisionado, e de visão computacional a modelos de linguagem
Materiais de referência
- Download Paper: artigo com todos os detalhes da pesquisa
- View on GitHub: repositório para verificar detalhes da implementação
- View Summary: resumo breve dos resultados da pesquisa
Transformers without Normalization foi aceito como artigo da CVPR 2025
1 comentários
Comentários do Hacker News
Ajustar o alpha quase não ajudou, então talvez seja necessário um tuning considerável de hiperparâmetros ou uma inicialização mais sofisticada. Tentei tanto a inicialização padrão do PyTorch quanto a inicialização ortogonal, mas não houve diferença
Ou talvez o otimizador escalar que eu uso não combine bem com isso. Uso um otimizador escalar customizado que converge mais rápido que o Adam, mas na camada DyT ele pareceu ficar só no nível do Adam
Também pode ser o tipo de coisa que só alcança os outros depois de dezenas de bilhões de tokens, mas não tenho orçamento para testar por tanto tempo
Se for possível substituir essas camadas, isso ajudaria bastante a reduzir o custo computacional
O tanh também deve ter outros efeitos. Às vezes a normalização está resolvendo problemas de condicionamento. Ainda assim, é bom ter mais alternativas
Recomendo o artigo original do ResNet, de Kaiming He e outros, além dos trabalhos posteriores
Para uma abordagem moderna sobre RNNs, vale ler o paper da DeepMind: https://arxiv.org/abs/2303.06349
A ideia central é que o maior autovalor, isto é, o raio espectral, deve ficar perto de 1. Isso significa que, ao aplicar repetidamente a transformação linear, as ativações não crescem nem diminuem
y = x + f(x)LNinputeLNoutput, parece colocar peso e viés também depois detanh(a*x)Para ver a similaridade, não seria melhor comparar com o resultado da LayerNorm sem peso e viés?
Se o resultado final for bom, tudo bem, mas olhando separadamente só para a parte que está sendo substituída, talvez desse para entender melhor o que está acontecendo