Neste post, vou explicar o que é atenção e como funcionam os mecanismos apresentados no artigo “Attention is All You Need”. No caminho, vamos ver como funcionam os mecanismos Scaled Dot-Product Attention e Multihead Attention, e ao final, também vamos implementar esses mecanismos do zero em PyTorch com foco em eficiência computacional.
O que é atenção?
No contexto de redes neurais para transdução de sequências, atenção refere-se à capacidade do modelo de considerar o contexto de cada elemento da sequência ao gerar um novo valor para cada elemento de entrada. Ou seja, o modelo pode “prestar atenção” em diferentes partes da sequência para produzir saídas mais precisas e contextualizadas.
Nesse processo, os pesos de atenção formam uma sequência intermediária que indica o quanto cada elemento da entrada deve ser considerado na geração de cada novo elemento. Ao aplicar esses pesos sobre os elementos originais, é gerada uma nova sequência de elementos para representar o conteúdo da sequência recebida, dessa vez representando melhor o contexto de cada elemento.
O uso de mecanismos de atenção dessa forma permite que os modelos compreendam melhor o contexto em que cada trecho está inserido. Isso torna o treinamento mais rápido e eficiente por realizar menos operações, além de ajudar o modelo a lidar com trechos ambíguos ao considerar o contexto completo ao interpretar cada elemento da sequência.
Por exemplo, considere o texto:
João pensou: O trem estava cheio, mas ele conseguiu uma cadeira livre.
Para simplificar, vamos ignorar os tokens especiais e assumir que cada palavra já foi convertida em um embedding (representado por ). Assim, o texto se transforma na seguinte sequência:
Então, um mecanismo de atenção recebe essa sequência e gera outra de mesmo comprimento, onde o embedding que representa a palavra “ele” será composto de embeddings mais próximos do trecho que compõe a palavra “João” do que da palavra “trem”:
Self-Attention
Todos os mecanismos utilizados na arquitetura dos Transformers funcionam combinando os próprios elementos da sequência recebida entre si para calcular a atenção de cada elemento. Ou seja, a atenção atribuída a cada elemento da sequência é baseada apenas nos elementos da sequência recebida, sem depender de informações externas.
No entanto, nem todos os mecanismos de atenção funcionam dessa forma. Alguns dependem de informações externas, como variáveis adicionais ou outros modelos, para calcular a atenção. Por isso, dizemos que os mecanismos explicados aqui utilizam Self-Attention: a atenção de cada elemento é determinada apenas com base nos próprios elementos da sequência, sem depender de fontes externas.
Queries, Keys e Values
Para explicar conceitualmente o funcionamento desses mecanismos, é importante destacar que, embora os termos dessas operações sejam numericamente iguais no início, cada um deles recebe um nome abstrato diferente para facilitar a compreensão do papel que desempenham no mecanismo de atenção.
De forma simplificada, um mecanismo de atenção pode ser comparado a um dicionário em Python: chaves (keys ou ) são associadas a valores (values ou ), e é possível recuperar um valor a partir de uma consulta (query ou ). No contexto do mecanismo de atenção, os papéis de query, key e value são desempenhados pelos próprios elementos da sequência recebida, conforme descrito a seguir:
- Query: O elemento da sequência para o qual queremos representar de outra forma. No exemplo anterior, seria a palavra “ele”.
- Keys: Todos os elementos da sequência recebida, que funcionam como possíveis referências para determinar o contexto da query.
- Value: Um novo elemento que representa o significado da query, calculado a partir do elemento original e das keys.
- Se a query for ambígua, o value será ajustado para refletir melhor seu significado no contexto. No exemplo anterior, para “ele”, o value ficará mais próximo do elemento que representa “João”.
- Se não houver ambiguidade, o value pode ser igual ou muito próximo ao embedding original da query. No exemplo anterior, para “João”, o value praticamente não muda.
O exemplo a seguir ilustra essa analogia:
sequence = keys = values = [ 'João', 'pensou', 'O', 'trem', 'estava', 'cheio', 'mas', 'ele', ...]
mechanism = AttentionMechanism(keys, values)query = 'ele'
assert mechanism[query] == 'João'No entanto, diferentemente do exemplo acima, retornar apenas um elemento das keys para uma query geralmente não é suficiente para capturar corretamente o contexto, especialmente em situações ambíguas. Por exemplo:
Vi João e Maria ontem. Eles estavam juntos.
Nesse caso, a palavra “Eles” é ambígua, pois pode se referir tanto a “João” quanto a “Maria”. Assim, um mecanismo de atenção não atribui todo o peso a um único value, mas distribui a atenção em diferentes proporções entre os possíveis valores, indicando o quanto cada um deve ser considerado:
sequence = keys = values = [ 'Vi', 'João', 'e', 'Maria', 'ontem.', 'Eles', ...]
mechanism = AttentionMechanism(keys, values)query = 'ele'
assert mechanism[query] == { 'Vi': 0.01, 'João': 0.45, 'e': 0.01, 'Maria': 0.45, 'ontem.': 0.01, 'Eles': 0.01, ...}Essas frações são chamadas de pesos de atenção e formam uma sequência normalizada, ou seja, a soma de todos os elementos é igual a 1. Como os pesos são normalizados e, nos Transformers, as sequências de entrada são representadas por embeddings, é possível calcular uma média ponderada desses embeddings usando os pesos de atenção para gerar uma nova representação. Essa média será a saída do mecanismo de atenção.
A normalização garante que a quantidade total de atenção distribuída entre os elementos seja limitada, permitindo que o mecanismo funcione corretamente. Assim, se o peso de um elemento aumenta em relação aos outros, pelo menos um dos demais pesos precisa diminuir proporcionalmente para que a somatória continue sendo igual a 1.
Como nos Transformers é necessário transformar cada elemento da sequência recebida, o mecanismo de atenção pode ser acelerado ao ser aplicado em batch, utilizando toda a sequência como queries. Assim, as queries correspondem à própria sequência recebida, permitindo processar todos os elementos simultaneamente.
Por isso, embora queries, keys e values sejam inicialmente derivados da mesma sequência, é importante destacar que cada um desempenha um papel específico e distinto para o mecanismo de atenção.
Mecanismos de atenção
Dot-Product Attention (DPA)
Nesse mecanismo, os pesos de atenção de cada elemento são calculados usando o produto interno (dot product) entre as queries e as keys correspondentes.
Ao calcular a nova sequência em batch, o DPA pode ser definido da seguinte forma:
No entanto, dessa forma, o produto interno entre dois vetores pode assumir valores de a . Isso significa que o mecanismo de atenção pode, potencialmente, atribuir atenção excessiva a determinados elementos, o que pode enviesar o modelo e prejudicar seu funcionamento. Para evitar esse problema e garantir que os pesos de atenção sejam normalizados, aplica-se a função softmax sobre esses valores.
Scaled Dot-Product Attention (SDPA)
Essa variação do DPA inclui um termo de normalização nos pesos para estabilizar os gradientes durante o backpropagation.
A normalização pelo fator tem base empírica. Antes dos Transformers, experimentos já mostravam que normalizar a variância dos gradientes gerados por camadas ocultas ajuda a evitar problemas como gradient vanishing e neurônios mortos durante o treinamento. Por isso, essa prática foi incorporada à arquitetura dos Transformers.
def apply_scaled_dot_product_attention( queries: torch.Tensor, keys: torch.Tensor, values: torch.Tensor,) -> torch.Tensor: keys = keys.transpose(2, 3) scores = queries @ keys / (split_embed_dim**0.5)
if mask is not None: scores = scores.masked_fill(mask, float("-inf"))
weights = F.softmax(scores, dim=3) outputs = weights @ values
return outputsProjeções lineares
Uma das maneiras para melhorar a performance dos Transformers é aplicar projeções lineares separadas às queries, keys e values antes do seu uso nos mecanismos de atenção, multiplicando cada uma por uma matrizes de parâmetros treinável específica (, e ). Isso faz com que queries, keys e values passem a pertencer a espaços diferentes, permitindo que o modelo aprenda, durante o treinamento, como transformar os embeddings originais em representações mais adequadas ao contexto de cada token. Essas projeções são otimizadas para melhorar a capacidade do Transformer de capturar relações contextuais relevantes entre os elementos da sequência.
queries_projection = nn.Linear(embed_dim, embed_dim, bias=False)keys_projection = nn.Linear(embed_dim, embed_dim, bias=False)values_projection = nn.Linear(embed_dim, embed_dim, bias=False)
queries = queries_projection(queries)keys = keys_projection(keys)values = values_projection(values)
embeddings = apply_scaled_dot_product_attention(queries, keys, values)Multihead Attention (MHA)
Essa variação do SDPA aplica o mecanismo de atenção vezes em paralelo, cada uma com diferentes projeções lineares dos elementos da sequência. O número de cabeças é um hiperparâmetro do modelo.
Como os pesos no SDPA são normalizados, cada cabeça de atenção não pode “prestar atenção” igualmente em todos os elementos ao mesmo tempo. O objetivo do uso de MHA é permitir que cada cabeça (ou seja, cada SDPA sendo realizado em paralelo) foque em diferentes padrões ou relações no contexto, enriquecendo a representação aprendida pelo modelo durante o treinamento.
Usando uma analogia, o MHA seria como ler um texto vezes, focando em partes diferentes a cada leitura para compreender melhor o contexto de cada palavra. No mecanismo, porém, todas essas “leituras” acontecem simultaneamente.
Embora esse mecanismo otimize o modelo, se implementado literalmente, seria necessário calcular o SDPA vezes, tornando o MHA vezes mais lento, o que pode tornar o mecanismo inviável computacionalmente. Para evitar esse problema, utiliza-se a adaptação a seguir no algoritmo, que permite que a sua complexidade computacional não cresça com o número de cabeças, mantendo a escalabilidade do modelo:
- Dividir as queries, keys e values em partes, transformando as dimensões do batch de para .
- Transpor o tensor para que a dimensão das cabeças venha antes da dimensão das sequências, mudando de para .
- Aplicar o SDPA separadamente em cada cabeça, usando projeções diferentes para cada uma.
- Concatenar os embeddings resultantes de todas as cabeças.
- Transpor o tensor para restaurar a ordem original das dimensões, voltando para .
- Aplicar uma projeção linear ao tensor.
Após a etapa 3, o funcionamento do algoritmo pode ser representado pela seguinte equação:
Assim, o SDPA é aplicado vezes, mas como cada cabeça opera em uma dimensão menor, cada operação é proporcionalmente mais rápida. Isso garante que a complexidade computacional do MHA permaneça equivalente à do SDPA, mesmo com múltiplas cabeças.
Na prática, o MHA ainda é um pouco mais lento que o SDPA devido à projeção linear ao final, mas a complexidade computacional não cresce com o número de cabeças ou o tamanho das sequências.
n_heads = 8split_embed_dim = embed_dim // n_headsoutputs_projection = nn.Linear(embed_dim, embed_dim, bias=False)
def split_embeddings( embeddings: torch.Tensor, batch_size: int, n_tokens: int,) -> torch.Tensor: splitted = embeddings.view(batch_size, n_tokens, n_heads, split_embed_dim) splitted = splitted.transpose(1, 2)
return splitted
def join_embeddings( embeddings: torch.Tensor, batch_size: int, n_tokens: int,) -> torch.Tensor: joined = embeddings.transpose(1, 2) joined = joined.contiguous() joined = joined.view(batch_size, n_tokens, embed_dim) return joined
def apply_multihead_attention( queries: torch.Tensor, keys: torch.Tensor, values: torch.Tensor,) -> torch.Tensor: batch_size, n_tokens, _ = queries.size()
queries = queries_projection(queries) keys = keys_projection(keys) values = values_projection(values)
queries = split_embeddings(queries, batch_size, n_tokens) keys = split_embeddings(keys, batch_size, n_tokens) values = split_embeddings(values, batch_size, n_tokens)
keys = keys.transpose(2, 3) scores = queries @ keys / (split_embed_dim**0.5) weights = F.softmax(scores, dim=3)
outputs = weights @ values outputs = join_embeddings(outputs, batch_size, n_tokens) outputs = outputs_projection(outputs)
return outputs
class MultiheadAttention(nn.Module): def __init__(self: Self, embed_dim: int, n_heads: int) -> None: super().__init__()
self.embed_dim = embed_dim self.n_heads = n_heads self.split_embed_dim = self.embed_dim // self.n_heads
self.queries_projection = nn.Linear(self.embed_dim, self.embed_dim, bias=False) self.keys_projection = nn.Linear(self.embed_dim, self.embed_dim, bias=False) self.values_projection = nn.Linear(self.embed_dim, self.embed_dim, bias=False) self.outputs_projection = nn.Linear(self.embed_dim, self.embed_dim, bias=False)
def split_embeddings( self: Self, embeddings: torch.Tensor, batch_size: int, n_tokens: int, ) -> torch.Tensor: splitted = embeddings.view( batch_size, n_tokens, self.n_heads, self.split_embed_dim, )
splitted = splitted.transpose(1, 2)
return splitted
def join_embeddings( self: Self, embeddings: torch.Tensor, batch_size: int, n_tokens: int, ) -> torch.Tensor: joined = embeddings.transpose(1, 2) joined = joined.contiguous() joined = joined.view(batch_size, n_tokens, self.embed_dim)
return joined
def forward( self: Self, queries: torch.Tensor, keys: torch.Tensor, values: torch.Tensor, ) -> torch.Tensor: batch_size, n_tokens, _ = queries.size()
queries = self.queries_projection(queries) keys = self.keys_projection(keys) values = self.values_projection(values)
queries = self.split_embeddings(queries, batch_size, n_tokens) keys = self.split_embeddings(keys, batch_size, n_tokens) values = self.split_embeddings(values, batch_size, n_tokens)
keys = keys.transpose(2, 3)
scores = queries @ keys / (self.split_embed_dim**0.5) weights = F.softmax(scores, dim=3)
outputs = weights @ values outputs = self.join_embeddings(outputs, batch_size, n_tokens) outputs = self.outputs_projection(outputs)
return outputsConclusão
Os mecanismos de atenção são uma maneira eficiente de ajustar as representações dos valores para que reflitam melhor seu significado no contexto. A eficiência do MHA nessa tarefa faz toda a diferença quando combinada com outras técnicas para criar Transformers capazes de realizar bem várias tarefas.
Apesar disso, a atenção sozinha pode ter algumas limitações. No próximo post, vamos ver como o Positional Encoding pode ajudar a representar a posição dos elementos em uma sequência e superar esses problemas.