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

Attention e o bloco transformer, derivados de uma média

Parta da média do context, veja onde ela falha e derive a fórmula de attention como reparo.

Nesta página

Você chega aqui com um tokenizer do Capítulo 7, uma tabela de embedding do Capítulo 8 e o objetivo que vem com eles: dados os tokens até agora, atribuir uma probabilidade ao próximo.

O que falta é o meio. Para prever o token tt, o modelo precisa de um vetor que resuma tudo o que veio antes, e nada do que você construiu produz isso. O embedding do token t1t-1 não é isso — esse é um modelo de bigramas, e ele não consegue saber que a frase começou com uma pergunta. Uma concatenação de todos os embeddings anteriores também não é isso: a quantidade deles muda a cada passo, e uma matriz de pesos fixa não consegue receber uma entrada de comprimento variável.

Então: um vetor de tamanho fixo, resumindo um número variável de vetores. Esse é o problema inteiro, e attention é o que você obtém ao resolvê-lo da maneira mais preguiçosa possível e depois consertar as duas coisas que quebram.

A resposta que a área tinha, e por que não vamos construí-la

Link para a seção: A resposta que a área tinha, e por que não vamos construí-la

De 1997 até por volta de 2017, o resumo era um estado recorrente: manter um vetor h\mathbf{h} e atualizá-lo a cada token, ht=f(ht1,xt)\mathbf{h}_t = f(\mathbf{h}_{t-1}, \mathbf{x}_t). Tamanho fixo, entrada variável, exatamente o formato certo.

Ele falhou de três maneiras, e a arquitetura deste capítulo responde às três. Fazer backpropagation por TT passos multiplica TT jacobianas, então o gradient desaparece ou explode — a doença que o Capítulo 5 mediu dentro de um único nó tanh\tanh. A LSTM1 foi projetada exatamente contra isso e empurrou o alcance utilizável de dezenas de passos para centenas, sem mudar o fato de que a informação do token 5 chega ao token 500 apenas se sobreviver a 495 atualizações sequenciais. Toda a fonte precisava caber em um vetor: em tradução sequence-to-sequence2, um encoder comprime a entrada em seu estado final. Bahdanau, Cho e Bengio deram nome a esse gargalo e o corrigiram em 2014, três anos antes do transformer, ao permitir que o decoder fizesse uma soma ponderada de todos os estados do encoder com pesos que ele próprio calculava.3 Tudo abaixo é essa ideia, aplicada por uma sequência a si mesma, com a recorrência removida. E a atualização é sequencial por construção: ht\mathbf{h}_t precisa de ht1\mathbf{h}_{t-1}, e uma GPU com dez mil núcleos não pode fazer nada com isso. A arquitetura que venceu não é obviamente mais inteligente; é aquela cujo passo caro é uma multiplicação de matrizes.

O outro viés indutivo clássico, convolução — deslizar um filtro pequeno por toda a entrada, de modo que uma feature detectada em qualquer lugar seja detectada em todos os lugares — também não será construído aqui; ele é quase exatamente certo para imagens e fica delegado a um curso de visão. Nem recorrência nem convolução reaparecem depois desta página, e é por isso que nenhuma das duas recebe um capítulo: o Capítulo 1 prometeu que as omissões seriam declaradas, não silenciosas.

A função mais óbvia de um número variável de vetores que retorna um vetor é a média:

ct=1ti=1txi\mathbf{c}_t = \frac{1}{t}\sum_{i=1}^{t} \mathbf{x}_i

Qualquer número de entradas, saída de tamanho fixo, diferenciável, de graça. Tabela de embedding mais essa média mais uma camada linear para o vocabulário é um modelo de linguagem completo em quinze linhas. Também é péssimo, e a maneira como é péssimo é toda a derivação.

O corpus abaixo é um megabyte de Shakespeare, 1.115.394 caracteres, passando por um tokenizer BPE em nível de byte do tipo construído no Capítulo 7, com vocabulário de 1024: 459.760 tokens com 2,43 caracteres cada, dividido 90/10. Todo modelo tem largura 128, vê 128 tokens e treina por 3000 passos de AdamW em 10310^{-3} com um batch de 64. A perplexidade é medida na divisão reservada.4

modeloparâmetrosperplexidade de validação
apenas o token atual, sem context nenhum263.16859,71
mais a média uniforme de tudo antes dele263.168248,07
mais learned position embeddings279.552245,93
média uniforme adicionada ao token em vez de substituí-lo263.16860,45

Leia a segunda linha duas vezes. Fazer a média do context não ajuda um pouco; deixa o modelo quatro vezes pior do que ignorar o context completamente. Há duas razões, ambas demonstráveis em vez de empíricas.

A média não consegue ver ordem. A adição é comutativa, então embaralhar a window deixa o resumo inalterado — não aproximadamente:

order.pyPYTHON
A = torch.tril(torch.ones(T, T))
A = A / A.sum(1, keepdim=True)          # rows of the averaging matrix
y = x[torch.randperm(T)]                # the same tokens, shuffled
print((A[-1] @ x - A[-1] @ y).abs().max().item())
TEXT
2.9802322387695312e-08

Ruído de ponto flutuante em uma soma reordenada: os dois resumos são o mesmo vetor. Um modelo cuja única visão do context é uma média não consegue distinguir the dog bit the man de the man bit the dog. A terceira linha prova que isso não se conserta adicionando posições às entradas — um learned position embedding em cada token antes da média comprou 2,14 pontos de 188. As posições entram na soma, e a soma as esquece.

E a média afoga o presente. Na posição 100, o token atual é um centésimo do resumo. Isso tem um conserto barato que você já possui: mantenha o token e adicione o resumo a ele — uma conexão residual, do Capítulo 6, e a quarta linha mostra o que ela faz. Com a diluição corrigida, a média uniforme não contribui com nada: 60,45 contra um baseline de 59,71. Todo token está ali, com o mesmo peso, e pesos iguais são o mesmo que ausência de informação.

O problema não é fazer média. São os pesos.

A média é uma multiplicação de matrizes, e a máscara é um softmax

Link para a seção: A média é uma multiplicação de matrizes, e a máscara é um softmax

Fazer média sobre um prefixo crescente parece um loop. É uma multiplicação por uma matriz triangular inferior cujas linhas somam um — e também, exatamente, um softmax:

mechanics.pyPYTHON
loop = torch.stack([x[:t + 1].mean(0) for t in range(T)])   # the obvious version

A = torch.tril(torch.ones(T, T))
A = A / A.sum(1, keepdim=True)
mat = A @ x                                                  # the same thing

S = torch.zeros(T, T).masked_fill(torch.tril(torch.ones(T, T)) == 0, float("-inf"))
soft = F.softmax(S, dim=-1) @ x                              # and the same thing again
TEXT
loop vs matmul   max |diff| = 5.960464477539063e-08
loop vs softmax  max |diff| = 5.960464477539063e-08

the averaging matrix A (rows sum to 1, upper triangle is zero):
  1.000 0.000 0.000 0.000 0.000 0.000
  0.500 0.500 0.000 0.000 0.000 0.000
  0.333 0.333 0.333 0.000 0.000 0.000
  0.250 0.250 0.250 0.250 0.000 0.000
  0.200 0.200 0.200 0.200 0.200 0.000
  0.167 0.167 0.167 0.167 0.167 0.167

Três componentes nomeados de um transformer agora estão na tela. O triângulo é a causal mask, imposta pelo objetivo: se a posição tt pudesse ver a posição t+1t{+}1, a resposta estaria na entrada — o vazamento que o Capítulo 6 mandou auditar, só que dentro da arquitetura. O softmax é como a máscara é implementada: definir entradas proibidas como -\infty as leva exatamente a zero e normaliza o que resta, então mascarar e normalizar são uma única operação. (Use -\infty, não -1e9: é o valor que o mascaramento significa, sobrevive a um cast para float16 como -\infty e evita que você decida se a constante escolhida é grande o bastante para o intervalo em que você por acaso está — que é a caixa de ponto flutuante do Capítulo 2 fazendo uma pergunta que você não precisa responder.) E os scores são o parâmetro livre. A média uniforme é o que você obtém quando todo score permitido é o mesmo número; coloque quaisquer números ali e o softmax os transforma em pesos válidos.

O restante deste capítulo é uma pergunta: de onde vêm esses números?

Eles não podem ser parâmetros simples. Uma matriz T×TT \times T aprendida seria idêntica para toda frase — poderia codificar "olhe quatro tokens para trás", mas nunca "olhe para o substantivo a que este pronome se refere". O peso ligando a posição tt à posição ii precisa depender do que há em ambas as posições, porque relevância é uma relação, não uma propriedade: a palavra it não é intrinsecamente relevante, ela é relevante para alguma coisa.

A função mais barata de dois vetores que retorna um número é o produto escalar do Capítulo 1. Dê um score à posição ii para a posição tt como xtxi\mathbf{x}_t \cdot \mathbf{x}_i e o mecanismo funciona — mal, de duas maneiras que forçam todo o resto. O produto escalar de um vetor consigo mesmo é sua norma ao quadrado, então todo token daria attention principalmente a si mesmo. E a relação seria simétrica: se it dá attention forte a animal, então animal dá attention forte a it, o que é falso na linguagem, onde um adjetivo precisa de seu substantivo muito mais do que o substantivo precisa do adjetivo.

Então dê a cada token dois papéis, como dois mapas lineares aprendidos dele: o que esta posição está procurando, qt=Wqxt\mathbf{q}_t = W_q\mathbf{x}_t, a query; e o que ela oferece para ser encontrada por, ki=Wkxi\mathbf{k}_i = W_k\mathbf{x}_i, a key. Calcule o score qtki\mathbf{q}_t \cdot \mathbf{k}_i e a simetria desaparece, porque WqWkW_q \neq W_k: um token pode anunciar uma coisa e procurar outra.

Uma coisa ainda está errada. A soma ponderada era sobre os próprios xi\mathbf{x}_i, o que força aquilo que é copiado a ser a mesma coisa que é comparada. Comparar quer as features que identificam um token; copiar quer as features úteis downstream. Então aprenda um terceiro mapa, vi=Wvxi\mathbf{v}_i = W_v\mathbf{x}_i, o value, e some esses.

A fórmula agora é contabilidade:

Attention(Q,K,V)=softmax ⁣(QKdk+M)V\mathrm{Attention}(Q, K, V) = \mathrm{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}} + M\right)V

com MM sendo a causal mask, zero na diagonal e abaixo dela, e -\infty acima. Em código, são trinta linhas, vinte das quais são shapes:

attention.pyPYTHON
class Head(nn.Module):
    """One head of causal self-attention."""

    def __init__(self, d_model, d_head, block):
        super().__init__()
        self.q = nn.Linear(d_model, d_head, bias=False)      
        self.k = nn.Linear(d_model, d_head, bias=False)      
        self.v = nn.Linear(d_model, d_head, bias=False)      
        self.d_head = d_head
        self.register_buffer("mask", torch.tril(torch.ones(block, block)).bool())

    def forward(self, x):
        T = x.shape[1]
        q, k, v = self.q(x), self.k(x), self.v(x)
        s = q @ k.transpose(-2, -1) / math.sqrt(self.d_head)          
        s = s.masked_fill(~self.mask[:T, :T], float("-inf"))          
        w = F.softmax(s, dim=-1)                                      
        return w @ v                                                  

Score, máscara, normalização, mistura. Todo o resto é uma projeção.

A divisão pela raiz quadrada, e contra o que ela protege

Link para a seção: A divisão pela raiz quadrada, e contra o que ela protege

Quase toda explicação de dk\sqrt{d_k} diz "para impedir que o softmax sature", o que é verdade e não explica nada. O argumento tem duas linhas da variância do Capítulo 2. Se as entradas de q\mathbf{q} e k\mathbf{k} são independentes, com média zero e variância um, cada produto qjkjq_j k_j tem variância um, e variâncias de coisas independentes se somam:

Var(qk)=j=1dkVar(qjkj)=dk\mathrm{Var}(\mathbf{q}\cdot\mathbf{k}) = \sum_{j=1}^{d_k}\mathrm{Var}(q_j k_j) = d_k

Então os scores têm desvio padrão dk\sqrt{d_k}. Medido em vinte mil pares aleatórios:

TEXT
     d     Var(q.k)         std   sqrt(d)
     4        3.975       1.994     2.000
    16       16.071       4.009     4.000
    64       64.249       8.016     8.000
   256      253.065      15.908    16.000
  1024     1015.562      31.868    32.000

Por que isso importa: o softmax é sensível à escala de uma maneira que uma camada linear não é. Dobrar a entrada de uma camada linear dobra sua saída; multiplicar scores por dez antes de um softmax transforma uma mistura suave em uma escolha dura. Uma linha de 64 scores, com e sem a divisão:

dkd_kmaior peso, sem divisãoentropiatokens efetivosmaior peso, divididoentropiatokens efetivos
40,2052,94419,00,0813,75842,9
160,4381,6925,40,0753,84946,9
640,4890,8742,40,0853,67339,4
2560,99990,00071,00,1433,54734,7
10241,00000,00001,00,1323,64438,3

"Tokens efetivos" é a exponencial da entropia: sobre quantas posições a linha realmente faz média. Sem divisão, em dk=256d_k = 256, uma head recém-inicializada dá attention a exatamente um token de 64, escolhido por nada além do sorteio aleatório.

Isso é ruim no forward e pior no backward, em uma forma que o Capítulo 5 já mediu em um tanh\tanh. Um softmax comprometido com uma entrada quase não tem derivada: a diagonal de sua jacobiana é wi(1wi)w_i(1-w_i), zero nas duas extremidades. Em duas mil linhas aleatórias:

dkd_kiwi(1wi)\sum_i w_i(1-w_i) sem divisãodivididolinhas saturadas (maior peso acima de 0,99)
40,84270,95680,2 % → 0,0 %
640,29400,960917,9 % → 0,0 %
2560,14060,960949,1 % → 0,0 %
10240,06810,961170,4 % → 0,0 %

Em dk=1024d_k = 1024, sete linhas em dez estão congeladas antes de o treinamento começar, e uma head que começa congelada não consegue aprender para onde olhar. Dividida, a quantidade fica plana em 0,96 em toda largura e nada satura.

Agora a parte que ninguém publica: isso muda a perplexidade final? Remova a divisão e treine, em quatro larguras de head:

largura da headsem divisãodividido por dk\sqrt{d_k}dividido por dkd_k
quatro heads, dk=32d_k = 3237,2938,0737,89
uma head, dk=128d_k = 12848,5146,1045,99
uma head, dk=256d_k = 25665,3747,53
uma head, dk=512d_k = 51267,0649,15
uma head, dk=1024d_k = 102476,6959,17

As duas primeiras linhas vêm do orçamento de 3000 passos acima; as três últimas são uma execução mais curta — 1500 passos, batch de 32, uma head, sem normalização antes das projeções — com as duas variantes sob configurações idênticas.

Em dk=32d_k = 32 a divisão não vale nada e a execução sem ela fica muito ligeiramente à frente. Isso não é licença para removê-la, porque em 256 ela vale 18 pontos de perplexidade e em 1024 vale 17. O mecanismo é visível nos próprios scores:

dkd_kstd do score na inicializaçãoapós 1500 passos, sem divisãoapós 1500 passos, divididolinhas saturadas, sem divisãodividido
25610,49121,672,1391,9 %0,8 %
51215,13836,852,6698,7 %1,3 %
102421,155147,463,4499,9 %16,5 %

A head sem divisão não se recupera. Ela dispara: o desvio padrão de seus scores vai de 21 na inicialização para 5147, a entropia de attention cai a zero e 99,9 % das linhas colocam mais de 0,99 de seu peso em um único token. Uma vez que uma head vira seletor duro, seu gradient é quase zero e nada a puxa de volta, então o colapso é estável. A head dividida fica em um desvio padrão de score de 3,44 após o mesmo treinamento, que é uma mistura suave que ainda pode ser alterada.

Vaswani et al. dizem exatamente isso e nada mais — eles suspeitam que os produtos "crescem muito em magnitude para valores grandes de dkd_k" e dividem.5 A palavra grandes carrega o peso, e as tabelas dizem onde grande começa: nada em 32, tudo até 256.

Mais de uma opinião, e os dois terços sobre os quais ninguém fala

Link para a seção: Mais de uma opinião, e os dois terços sobre os quais ninguém fala

Uma head é uma linha de softmax por posição, então ela contém uma resposta para "o que é relevante aqui". Prever a palavra depois de the em the animal that crossed the wet street precisa do slot sintático, do sujeito e do token anterior ao mesmo tempo, e uma distribuição de probabilidade não consegue se concentrar em três lugares. Então execute várias heads em paralelo, cada uma com largura dmodel/hd_{\text{model}}/h, concatene e misture com mais uma matriz WoW_o: você particionou a largura, não adicionou a ela.

Attention também faz exatamente uma coisa — move informação entre posições. Toda operação no código acima é linear ao longo do eixo de features, e o Capítulo 5 provou o que é uma pilha de mapas lineares. Então cada bloco também carrega uma pequena MLP aplicada a cada posição independentemente, expandindo a largura por quatro e voltando, com uma GELU no meio. Vale memorizar a divisão de trabalho: attention mistura entre posições, a rede feed-forward calcula dentro de uma posição.

A escada completa, cada linha adicionando uma peça à linha acima:

modeloparâmetrosperplexidade de validação
média uniforme, adicionada279.55260,45
uma attention head, substituindo o token328.70455,47
uma attention head, adicionada328.70446,10
quatro heads em vez de uma345.21643,21
mais a rede feed-forward476.92839,87
mais LayerNorm — o bloco completo477.69638,07

Pesos aprendidos vencem pesos uniformes por 14 pontos de perplexidade, que é todo o argumento deste capítulo em uma linha. Quatro heads compram mais 3 por 16.512 parâmetros extras. E a mesma head vale 9 pontos a mais adicionada do que substituindo: attention traz informação para dentro, ela não decide o que uma posição é.

Agora onde os parâmetros realmente ficam, o que surpreende quem só viu o diagrama:

larguraheadsattentionfeed-forwardtotal por bloco
128465.664 (33,2 %)131.712 (66,6 %)197.888
768122.360.064 (33,3 %)4.722.432 (66,6 %)7.085.568
40963267.112.960 (33,3 %)134.238.208 (66,7 %)201.367.552

Dois terços de todo bloco transformer são a rede feed-forward, em qualquer escala, porque attention tem quatro matrizes d×dd \times d e a MLP tem o equivalente a oito. O que quer que um modelo saiba, a maior parte dos parâmetros que seguram esse conhecimento está na MLP por posição.

Residuais e LayerNorm, herdados do Capítulo 6

Link para a seção: Residuais e LayerNorm, herdados do Capítulo 6

LayerNorm foi construída e medida no Capítulo 6, e este capítulo a usa como ela foi deixada ali; conexões residuais foram nomeadas e submetidas a ablation ali, e são construídas aqui. As linhas "adicionada, não substituindo" acima são conexões residuais, valendo 188 pontos de perplexidade para a média e 9 para uma head. LayerNorm7 normaliza cada exemplo ao longo de suas features, e o Capítulo 6 deu as razões pelas quais ela, e não BatchNorm, sobreviveu aqui — nenhuma dependência do batch, nenhuma estatística acumulada, idêntica em treinamento e inferência, indiferente ao comprimento da sequência — e cada uma delas vira um requisito quando você gera um token por vez para um usuário, que é onde o Capítulo 13 acaba. Ela custa 768 parâmetros e compra 1,8 ponto de perplexidade.

block.pyPYTHON
class Block(nn.Module):
    def forward(self, x):
        x = x + self.att(self.ln1(x))     
        x = x + self.ff(self.ln2(x))      
        return x

Observe onde a normalização fica: na entrada de cada subcamada, com o caminho residual da entrada à saída nunca normalizado. Isso é pre-norm. O artigo de 2017 faz o oposto, x = LayerNorm(x + Att(x))post-norm, que coloca uma LayerNorm no próprio caminho residual.

Xiong et al. explicaram a diferença pelo gradient na inicialização, que em uma rede post-norm fica mal escalado com a profundidade — a razão pela qual o transformer original precisava de warmup de learning rate para treinar.8 Doze blocos, 1000 passos, learning rate 3×1033 \times 10^{-3}:

TEXT
gradient norm per block at initialisation, before any step
  pre-norm    block 1 0.0498 ... block 12 0.0657   ratio last/first  1.32
  post-norm   block 1 0.0977 ... block 12 0.1613   ratio last/first  1.65

  pre-norm,  no warmup          perplexity   37.82
  pre-norm,  200-step warmup    perplexity   37.62
  post-norm, no warmup          perplexity  308.05
  post-norm, 200-step warmup    perplexity   37.88

Post-norm sem warmup é oito vezes pior, e post-norm com warmup iguala pre-norm exatamente. Warmup não é uma boa prática geral aqui; é um patch para uma organização específica da normalização, e mover a LayerNorm remove a necessidade dele. É por isso que praticamente todo modelo desde 2019 é pre-norm, e por isso o diagrama de 2017 deve ser lido como história, não como especificação.

Remova os position embeddings e o modelo ainda treina; ele simplesmente não consegue saber onde nada está, e isso é uma simetria, não uma falha de treinamento. Nada no score de attention menciona os próprios tt ou ii, então permutar a entrada permuta a saída: self-attention é equivariante a permutações. É a cegueira à ordem da média em um disfarce melhor — a causal mask restaura um pouco de ordem, já que cada posição vê um prefixo diferente, mas dentro de um prefixo todas as ordenações são iguais.

Quatro maneiras de injetar posição, treinadas em windows de 64 tokens e avaliadas em 64, 128 e 256 — além de qualquer comprimento que tenham visto:

posiçõesperplexidade em 64em 128em 256
nenhuma48,7952,6357,52
learned absolute embeddings38,63108,47181,94
senoides fixas42,9695,26152,25
RoPE44,1250,5284,84
ALiBi44,9543,5142,49

Learned absolute embeddings — um vetor por posição, adicionado ao token — vencem no comprimento treinado e depois caem de um penhasco, porque a posição 100 nunca esteve em um batch e seu embedding ainda é o vetor aleatório com que começou. Senoides, a escolha original, são calculadas em vez de aprendidas, a partir de senos e cossenos em frequências espaçadas geometricamente; o artigo de 2017 esperava que isso extrapolasse, e a tabela diz que não — a função está definida na posição 200, mas o modelo nunca aprendeu a lê-la ali. RoPE9 não adiciona nada e, em vez disso, rotaciona query e key por um ângulo proporcional à posição, em fatias bidimensionais; como rotacionar igualmente os dois lados de um produto escalar o deixa inalterado, o score acaba dependendo apenas de tit - i, então a posição se torna relativa de graça e não há tabela a esgotar. Ele degrada, mas degrada. ALiBi10 é o resultado mais simples e mais estranho aqui: uma penalidade linear no score proporcional à distância, com uma inclinação diferente por head. Sua perplexidade melhora à medida que a window cresce além do comprimento de treinamento, de 44,95 para 42,49, porque a penalidade é definida para qualquer distância e cada head continua fazendo o que foi treinada para fazer.

A lição sobrevive à tabela: uma arquitetura que não consegue representar algo é um problema diferente de uma que nunca aprendeu aquele intervalo, e o segundo é o que morde. Essa também é a mecânica por trás de todo anúncio de "estendemos o context para 128K" — quase sempre são reescalonamentos de uma codificação rotativa, e é por isso que o Capítulo 16 diz que o limite de context se move em vez de desaparecer.

Dropout é herdado da mesma forma: aparece nos pesos de attention depois do softmax, na saída de cada subcamada antes da adição residual e na soma de embeddings, fazendo exatamente o que o Capítulo 6 descreveu. Em grandes execuções de pretraining, muitas vezes é definido como zero, porque um modelo que vê cada token uma vez não está em posição de overfit.

Dois tensores na camada têm shape n×nn \times n, onde nn é o número de tokens: os scores e os pesos após o softmax. Todo o resto — toda projeção, a MLP inteira — é linear em nn.

Uma camada de attention, largura 512, 8 heads, batch de um, float32, em uma GPU de notebook. Leia as duas colunas de milissegundos apenas por suas proporções: são wall clock em uma placa de notebook de 8 GB que desacelera de 1.785 MHz para menos de 300 MHz quando esquenta, então uma execução fria do mesmo código volta sete a dez vezes mais rápida, e uma ocupada ainda mais lenta. As colunas em megabytes são contagens de bytes do alocador e não mudam.

TEXT
  tokens   ms total    ms x4   ms projections   attn matrix MB    peak MB    MB x4
     128      2.246        -            1.324              0.5       14.6        -
     256      2.855     1.27            2.113              2.0       19.2     1.31
     512      5.761     2.02            3.105              8.0       34.4     1.79
    1024     16.414     2.85            4.008             32.0       89.1     2.59
    2048     51.573     3.14            9.989            128.0      296.1     3.32
    4096    225.432     4.37           20.176            512.0     1100.1     3.72
    8192    832.838     3.69           40.106           2048.0     4300.1     3.91
   16384   OUT OF MEMORY                                 8192.0

fitted exponent (log-log slope, last four rows):  time ~ n^1.91   memory ~ n^1.87

As colunas x4 são a razão em relação à linha acima, e dobrar nn converge exatamente para 4 tanto em tempo quanto em memória — 3,91 no último passo contra um 4 teórico. A coluna de projeções é o controle: 4,0 ms em 1024 tokens para 40,1 ms em 8192, um fator de dez para um fator de oito. Linear, como anunciado.

Então a última linha. Uma camada de attention, uma sequência, sem nenhum modelo ao redor, fica sem memória em uma GPU de 8 GB com 16.384 tokens — só a matriz de scores teria 8 GB, sendo 8 heads vezes 16.384 vezes 16.384 vezes 4 bytes. Não o modelo; um tensor intermediário em uma camada.

Esse é o fato físico por baixo de três capítulos posteriores. É por isso que uma context window tem um limite, que o Capítulo 16 transforma em preço. É por isso que FlashAttention existe, calculando o mesmo resultado em blocos sem nunca armazenar a matriz — uma otimização de memória antes de ser de velocidade.11 E é a aritmética por trás do preço de um prompt longo, que o Capítulo 24 paga em um loop de agent — uma questão separada da outra descoberta daquele capítulo, de que um modelo também usa um context longo pior, algo que ele mede e se recusa a culpar nesta fórmula.

Mostrar detalhes

As duas variantes que encolhem cache, nomeadas aqui e pagas no Capítulo 13.

A geração armazena em cache as keys e values dos tokens já processados — uma key e um value por token, por head por camada. Multi-query attention12 mantém hh projeções de query, mas uma única projeção de key e value compartilhada por todas as heads, dividindo esse cache por hh. Grouped-query attention13 interpola: as heads são agrupadas, cada grupo compartilhando uma key e um value, de modo que g=hg = h é attention comum e g=1g = 1 é multi-query. Quase todo modelo aberto desde 2023 a usa com 4 ou 8 grupos. Nenhuma das duas existe por qualidade; ambas existem pelo tamanho desse cache, e o Capítulo 13 faz a aritmética que transforma isso em "qual modelo cabe na sua GPU".

O artigo de 2017 descreve um encoder-decoder: uma pilha lendo a fonte com attention sem máscara, uma segunda gerando o alvo causalmente, e um terceiro tipo de attention no meio, onde as queries do decoder encontram as keys do encoder. Isso é certo para tradução, em que entrada e saída são duas sequências.

O que venceu foi a metade decoder-only — uma pilha, causal do começo ao fim, entrada e saída na mesma sequência — e a razão não é elegância. "Prever o próximo token" roda em qualquer texto, então o conjunto de treinamento é a internet em vez de um corpus paralelo, e tudo vira essa única tarefa: uma tradução é um documento contendo fonte e depois alvo, uma pergunta e sua resposta são um documento, uma conversa com uma tool call no meio é um documento. O Capítulo 11 trata de como esse último é fabricado. Encoders não desapareceram — um encoder vê a entrada inteira de uma vez, que é o que você quer quando o trabalho é representar um texto em vez de continuá-lo, e é por isso que os retrieval embeddings do Capítulo 19 vêm de encoders e não do modelo que faz o chat.

Com o bloco definido, o tamanho do modelo é aritmética. Por bloco, com largura dd e expansão de quatro vezes: 4d2+4d4d^2 + 4d para Wq,Wk,Wv,WoW_q, W_k, W_v, W_o com vieses nos quatro, como o GPT-2 os tem — a tabela acima deixa o viés fora de três deles, daí 2.304 a menos por bloco em d=768d = 768; 8d2+5d8d^2 + 5d para a MLP; 4d4d para duas LayerNorms — 12d2+13d12d^2 + 13d, mais uma tabela de tokens de V×dV \times d e, para posições absolutas, nctx×dn_{\text{ctx}} \times d. Para o formato do GPT-2 small — d=768d = 768, 12 blocos, vocabulário de 50.257, context de 1024, a camada de saída compartilhando os pesos do embedding:

TEXT
  token embeddings     50,257 x 768 = 38,597,376
  position embeddings   1,024 x 768 =    786,432
  one block                             7,087,872
  12 blocks                            85,054,464
  final LayerNorm         2 x 768 =        1,536
  total (weights tied)                124,439,808

Que é o tamanho publicado desse modelo. A fórmula não é uma aproximação; ela é o modelo. Observe também que quase um terço de um modelo pequeno é a tabela de embedding, e é por isso que o tamanho do vocabulário é uma decisão arquitetural, não de pré-processamento — o trade-off que o Capítulo 7 preparou.

Perplexidade é um número sobre um corpus. O que uma head faz é outra pergunta, e um modelo treinado em um megabyte de Shakespeare é o instrumento errado para ela: o honesto a dizer sobre o mapa de attention de um modelo de 500.000 parâmetros é que ele em geral não é interpretável. Então: uma linguagem em que a pergunta tem uma resposta certa.

A ilustração clássica é the animal did not cross the street because it was too tired, em que it é o animal, contra …because it was too wet, em que uma palavra move o referente para a rua. Esses são esquemas de Winograd14 — pares de frases idênticas exceto por uma palavra, em que essa palavra decide a que um pronome se refere.

Eles também são solucionáveis trapaceando, que é a parte que os tutoriais pulam. Se os dois candidatos são um animal e um lugar, tired e wet identificam o referente por categoria, e um modelo que só sabe quais palavras estão presentes acerta sem saber nada sobre ordem. Medido nessa versão da tarefa, com pares animal/lugar reservados:

TEXT
uniform causal average           held-out referent accuracy 100.0 %
one transformer block            held-out referent accuracy  91.7 %

O saco de palavras vence o transformer. Qualquer demonstração construída sobre essa frase não prova nada sobre attention.

Então feche o buraco: extraia ambos os candidatos de um conjunto de dezesseis substantivos, qualquer um dos quais pode aparecer em qualquer slot, e divida os adjetivos por papel em vez de categoria — quatro fazendo it ser quem cruza (tired, scared, slow, weak), quatro fazendo ser o cruzado (wet, wide, busy, steep).

TEXT
the {x} did not cross the {y} because it was too {adj} , so the {ref} waited .

Treine como um preditor comum do próximo token, pontue uma posição — a palavra depois de so the — e construa o conjunto reservado com pares de substantivos cuja ordem invertida estava no treinamento, de modo que qualquer coisa que saiba quais dois substantivos estão presentes, mas não qual veio primeiro, deve responder ao contrário.

modeloparâmetrosreservadonomeia o outro substantivo
apenas token atual5.7965,2 %5,2 %
média causal uniforme5.79627,9 %50,0 %
uma head de learned attention18.08435,4 %64,6 %
quatro heads22.24475,0 %15,6 %
um bloco transformer55.71692,7 %4,2 %
dois blocos transformer105.508100,0 %0,0 %

O acaso entre os dois substantivos presentes é 50 %. A média uniforme chega a 27,9 % e responde com o substantivo errado do par exatamente metade das vezes — a assinatura de algo que sabe quais palavras estão ali e nada sobre sua ordem, como o teste de embaralhamento previu três seções atrás.

Agora o mapa: a attention na posição que precisa nomear o referente, média entre as quatro heads de cada bloco, para as duas frases que diferem por uma palavra. Uma média uniforme colocaria 0,067 em cada um dos quinze tokens visíveis.

TEXT
the animal did not cross the street because it was too tired , so the animal waited .
  blk 1  the:0.00 animal:0.70 did:0.00 not:0.00 cross:0.00 the:0.00 street:0.06
         because:0.00 it:0.00 was:0.00 too:0.00 tired:0.00 ,:0.05 so:0.00 the:0.19
  blk 2  the:0.00 animal:0.00 did:0.00 not:0.00 cross:0.00 the:0.00 street:0.00
         because:0.00 it:0.00 was:0.00 too:0.00 tired:1.00 ,:0.00 so:0.00 the:0.00

the animal did not cross the street because it was too wet , so the street waited .
  blk 1  the:0.00 animal:0.70 did:0.00 not:0.00 cross:0.00 the:0.00 street:0.06
         because:0.00 it:0.00 was:0.00 too:0.00   wet:0.00 ,:0.05 so:0.00 the:0.19
  blk 2  the:0.00 animal:0.00 did:0.00 not:0.00 cross:0.03 the:0.00 street:0.49
         because:0.00 it:0.00 was:0.00 too:0.20   wet:0.03 ,:0.00 so:0.00 the:0.25

O bloco 1 é idêntico nas duas frases — 0,70 no primeiro substantivo, qualquer que seja o adjetivo. Isso não é uma falha, mas uma prova: na primeira camada, a query em uma posição é função do token e do índice daquela própria posição, e the na posição 14 é o mesmo token nas duas frases. Uma head de primeira camada não consegue condicionar em uma palavra que ainda não buscou. Então o bloco 1 faz a única coisa útil disponível e puxa o primeiro substantivo para frente.

O bloco 2 é onde as frases se separam, e a mesma linha em todos os oito adjetivos mostra a regra que o modelo encontrou:

adjetivobloco 2 em animalem streetno adjetivoresposta
tired, scared, slow, weak0,0000,0001,000animal
wet, wide, busy, steep0,0000,4910,00–0,03street

Para um adjetivo de quem cruza, o segundo bloco gasta todo o seu peso no adjetivo, porque a resposta já está no residual stream — o bloco 1 a colocou ali — e tudo de que ele precisa é confirmação. Para um adjetivo de quem é cruzado, ele vai buscar o outro substantivo. Isso é um circuito de dois saltos: uma head move uma candidata para frente, uma head em uma camada posterior lê um token que decide se deve mantê-la. Composição entre camadas é o mecanismo, e é por isso que um bloco chegou a 92,7 % e dois chegaram a 100 %.

Também é o formato do circuito mais bem documentado em modelos reais. Induction heads — uma head de token anterior alimentando uma head na camada seguinte que completa o padrão [A][B] … [A] → [B] — são o que o trabalho de interpretabilidade da Anthropic identifica por trás de grande parte do in-context learning, e elas se formam em um momento identificável durante o pretraining. Este capítulo não tenta essa análise: ela é delegada, com os dois artigos nas referências, porque ler circuitos de um modelo real é uma área de pesquisa, não uma seção.

Por fim, a implementação. As trinta linhas acima, com seus pesos copiados do próprio PyTorch:

TEXT
ours vs nn.MultiheadAttention           max |diff| = 1.7881393432617188e-07
ours vs F.scaled_dot_product_attention  max |diff| = 1.7881393432617188e-07

1.8×1071.8 \times 10^{-7} em saídas cuja magnitude média é 0,159: a mesma aritmética em uma ordem diferente, com precisão float32.

Você tem a arquitetura da qual todo modelo no restante deste curso é construído, e ela é menor do que sua reputação: uma média ponderada cujos pesos são aprendidos, uma MLP por posição que guarda dois terços dos parâmetros, duas normalizações e duas adições, empilhadas.

O que você não tem é um modelo que saiba alguma coisa, e empilhar não vai consertar isso sozinho. Dois blocos neste corpus chegam a uma perplexidade de treinamento de 14,49 e uma perplexidade de validação de 40,57, contra 18,77 e 38,07 de um bloco — mais capacidade, melhor no que viu, pior no que não viu, que é a tabela do Capítulo 6 com um transformer dentro. A distância entre este modelo e aqueles com que os Capítulos 14 a 30 conversam não é arquitetural. É o mesmo bloco, repetido mais vezes, sobre muito mais texto.

O que torna isso um problema de contabilidade, e a contabilidade é mais estranha do que parece. Quanto texto, e de onde alguém o obtém? Quanta aritmética, e como estimá-la antes de o dinheiro ser gasto? Dado um orçamento fixo, é melhor tornar o modelo maior ou mostrar mais dados a ele — e existe uma resposta correta, ou só uma moda? O Capítulo 10 responde às três por medição, e coloca um preço na forma útil mais barata da pergunta: quanto custa, hoje, treinar um modelo como o GPT-2 do zero?


Três explicações deste material são melhores que esta naquilo para que servem, e este capítulo foi escrito para ser lido junto com elas. The Illustrated Transformer, de Jay Alammar, é a melhor imagem do fluxo de dados já desenhada. The Annotated Transformer, da Harvard NLP, é o artigo de 2017 com código executável intercalado linha por linha. Let's build GPT: from scratch, in code, spelled out, de Andrej Karpathy, constrói o mesmo modelo ao vivo em duas horas, e a escada de ablations acima é a mesma espinha medida em outro corpus. Para a pergunta de interpretabilidade que este capítulo apenas toca, as fontes primárias são Elhage et al., A Mathematical Framework for Transformer Circuits (2021) e Olsson et al., In-context Learning and Induction Heads (2022), ambos do grupo de interpretabilidade da Anthropic.

  1. Hochreiter, S. e Schmidhuber, J. Long Short-Term Memory. Neural Computation 9(8), pp. 1735–1780 (1997).

  2. Sutskever, I., Vinyals, O. e Le, Q. V. Sequence to Sequence Learning with Neural Networks. arXiv:1409.3215 (2014). O encoder-decoder cujo vetor de context único é o gargalo.

  3. Bahdanau, D., Cho, K. e Bengio, Y. Neural Machine Translation by Jointly Learning to Align and Translate. arXiv:1409.0473 (2014). Attention, três anos antes do transformer.

  4. Perplexidade é a exponencial da entropia cruzada média por token, do Capítulo 8. Todo número aqui usa o mesmo tokenizer e a mesma divisão de validação, que é a única condição sob a qual duas perplexidades podem ser comparadas.

  5. Vaswani, A., Shazeer, N., Parmar, N., Uszkoreit, J., Jones, L., Gomez, A. N., Kaiser, Ł. e Polosukhin, I. Attention Is All You Need. arXiv:1706.03762 (2017). A seção 3.2.1 é a frase única sobre dk\sqrt{d_k} que este capítulo passa uma seção medindo.

  6. Shazeer, N., Mirhoseini, A., Maziarz, K., Davis, A., Le, Q., Hinton, G. e Dean, J. Outrageously Large Neural Networks: The Sparsely-Gated Mixture-of-Experts Layer. arXiv:1701.06538 (2017).

  7. Ba, J. L., Kiros, J. R. e Hinton, G. E. Layer Normalization. arXiv:1607.06450 (2016). Introduzida e medida no Capítulo 6; usada aqui sem alterações.

  8. Xiong, R., Yang, Y., He, D., Zheng, K., Zheng, S., Xing, C., Zhang, H., Lan, Y., Wang, L. e Liu, T.-Y. On Layer Normalization in the Transformer Architecture. arXiv:2002.04745 (2020). A análise de gradient por trás de pre-norm, e o argumento de que warmup é um sintoma.

  9. Su, J., Lu, Y., Pan, S., Murtadha, A., Wen, B. e Liu, Y. RoFormer: Enhanced Transformer with Rotary Position Embedding. arXiv:2104.09864 (2021).

  10. Press, O., Smith, N. A. e Lewis, M. Train Short, Test Long: Attention with Linear Biases Enables Input Length Extrapolation. arXiv:2108.12409 (2021). O resultado de extrapolação reproduzido acima.

  11. Dao, T., Fu, D. Y., Ermon, S., Rudra, A. e Ré, C. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. arXiv:2205.14135 (2022).

  12. Shazeer, N. Fast Transformer Decoding: One Write-Head is All You Need. arXiv:1911.02150 (2019).

  13. Ainslie, J., Lee-Thorp, J., de Jong, M., Zemlyanskiy, Y., Lebrón, F. e Sanghai, S. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. arXiv:2305.13245 (2023).

  14. Levesque, H. J., Davis, E. e Morgenstern, L. The Winograd Schema Challenge. KR (2012). A construção por trás da frase animal / street que todo tutorial de attention usa.

Pronto para deixar a LIA escolher por você?

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