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

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

Partimos da média como resumo mais barato do contexto, medimos onde falha e deixamos a fórmula de attention surgir da correção.

Nesta página

Chega aqui com um tokenizador do Capítulo 7, uma tabela de embedding do Capítulo 8 e o objetivo que os acompanha: 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 vem antes, e nada do que construiu produz esse vetor. O embedding do token t1t-1 não serve — isso é um modelo de bigramas, e não consegue saber que a frase começou com uma pergunta. Uma concatenação de todos os embeddings anteriores também não serve: o seu número muda a cada passo, e uma matriz de pesos fixa não consegue receber uma entrada de comprimento variável.

Portanto: um vetor de tamanho fixo, a resumir um número variável de vetores. Esse é o problema todo, e attention é o que se obtém ao resolvê-lo da forma mais preguiçosa possível e depois corrigir as duas coisas que se partem.

A resposta que a área tinha, e porque não a estamos a construir

Ligação para a secção: A resposta que a área tinha, e porque não a estamos a construir

De 1997 até cerca 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 a forma certa.

Falhou de três formas, e a arquitetura deste capítulo responde às três. Fazer backpropagation através de TT passos multiplica TT jacobianos, por isso o gradiente desaparece ou explode — a doença que o Capítulo 5 mediu dentro de um único nó tanh\tanh. A LSTM1 foi desenhada exatamente contra isso e levou o intervalo utilizável de dezenas de passos para centenas, sem alterar o facto de a informação do token 5 chegar ao token 500 apenas se sobreviver a 495 atualizações sequenciais. Toda a fonte tinha de caber num vetor: na tradução sequência-para-sequência2, um codificador comprime a entrada no seu estado final. Bahdanau, Cho e Bengio deram nome a esse estrangulamento e corrigiram-no em 2014, três anos antes do transformer, permitindo que o descodificador tomasse uma soma ponderada de todos os estados do codificador, com pesos que ele próprio calculava.3 Tudo o que vem abaixo é essa ideia, aplicada por uma sequência a si própria, com a recorrência eliminada. 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 nada pode fazer 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 pequeno filtro sobre toda a entrada, para que uma característica detetada em qualquer ponto seja detetada em todo o lado — também não é construído aqui; é quase exatamente certo para imagens e fica entregue a um curso de visão. Nem recorrência nem convolução voltam a aparecer depois desta página, e é por isso que nenhuma delas 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 devolve 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, tamanho de saída fixo, diferenciável, grátis. Uma tabela de embedding mais esta média mais uma camada linear para o vocabulário é um modelo de linguagem completo em quinze linhas. Também é péssimo, e a forma como é péssimo é toda a derivação.

O corpus abaixo é um megabyte de Shakespeare, 1.115.394 caracteres, através de um tokenizador BPE ao nível do byte do tipo construído no Capítulo 7, com um vocabulário de 1024: 459.760 tokens a 2,43 caracteres cada, dividido 90/10. Todos os modelos têm largura 128, veem 128 tokens e treinam durante 3000 passos de AdamW a 10310^{-3} com um batch de 64. A perplexidade é medida na divisão retida.4

modeloparâmetrosperplexidade de validação
apenas o token atual, sem contexto nenhum263.16859,71
mais a média uniforme de tudo o que vem antes263.168248,07
mais embeddings de posição aprendidos279.552245,93
média uniforme adicionada ao token em vez de o substituir263.16860,45

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

A média não consegue ver a ordem. A adição comuta, por isso baralhar a janela 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 vírgula flutuante numa soma reordenada: os dois resumos são o mesmo vetor. Um modelo cuja única visão do contexto é uma média não consegue distinguir o cão mordeu o homem de o homem mordeu o cão. A terceira linha prova que isto não se corrige adicionando posições às entradas — um embedding de posição aprendido em cada token antes da média comprou 2,14 pontos em 188. As posições entram na soma, e a soma esquece-as.

E a média afoga o presente. Na posição 100, o token atual é um centésimo do resumo. Isso tem uma correção barata que já conhece: manter o token e adicionar o resumo a ele — uma ligação residual, do Capítulo 6, e a quarta linha mostra o que faz. Com a diluição corrigida, a média uniforme não contribui nada: 60,45 contra uma referência de 59,71. Todos os tokens estão lá dentro, ponderados de forma igual, e ponderação igual é o mesmo que ausência de informação.

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

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

Ligação para a secção: A média é uma multiplicação de matrizes, e a máscara é uma softmax

Fazer a média sobre um prefixo crescente parece um ciclo. É uma multiplicação por uma matriz triangular inferior cujas linhas somam um — e também, exatamente, uma 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 estão agora no ecrã. 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 — a fuga que o Capítulo 6 lhe disse para auditar, só que dentro da arquitetura. A softmax é a forma como a máscara é implementada: definir entradas proibidas como -\infty envia-as exatamente para zero e normaliza o que resta, pelo que mascarar e normalizar são uma só operação. (Use -\infty, não -1e9: é o valor que a máscara significa, sobrevive a uma conversão para float16 como -\infty, e evita-lhe decidir se a constante que escolheu é suficientemente grande para o intervalo em que calhou estar — que é a caixa de vírgula flutuante do Capítulo 2 a fazer uma pergunta a que não precisa de responder.) E os scores são o parâmetro livre. A média uniforme é o que se obtém quando todos os scores permitidos são o mesmo número; ponha lá quaisquer números e a softmax transforma-os em pesos válidos.

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

Não podem ser parâmetros simples. Uma matriz T×TT \times T aprendida seria idêntica para todas as frases — poderia codificar “olha quatro tokens para trás”, mas nunca “olha para o substantivo a que este pronome se refere”. O peso que liga a posição tt à posição ii tem de depender do que está em ambas as posições, porque relevância é uma relação, não uma propriedade: a palavra isso não é intrinsecamente relevante, é relevante para alguma coisa.

A função mais barata de dois vetores que devolve um número é o produto escalar do Capítulo 1. Pontue a posição ii para a posição tt como xtxi\mathbf{x}_t \cdot \mathbf{x}_i e o mecanismo funciona — mal, de duas formas que obrigam a tudo o resto. O produto escalar de um vetor consigo próprio é a sua norma ao quadrado, por isso cada token atenderia sobretudo a si próprio. E a relação seria simétrica: se isso atende fortemente a animal, então animal atende fortemente a isso, o que é falso sobre a linguagem, onde um adjetivo precisa muito mais do seu substantivo do que o substantivo precisa do adjetivo.

Por isso, dê a cada token dois papéis, como dois mapas lineares aprendidos dele: aquilo que esta posição está à procura, qt=Wqxt\mathbf{q}_t = W_q\mathbf{x}_t, a query; e aquilo que oferece para ser encontrado por, ki=Wkxi\mathbf{k}_i = W_k\mathbf{x}_i, a key. Pontue 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.

Ainda falta uma coisa. A soma ponderada era sobre os próprios xi\mathbf{x}_i, o que obriga a que a coisa que é copiada seja a coisa que é correspondida. A correspondência quer as características que identificam um token; a cópia quer as características úteis a jusante. Portanto, 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 a causal mask, zero na diagonal e abaixo dela, e -\infty acima. Em código são trinta linhas, vinte das quais são formas:

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                                                  

Pontuar, mascarar, normalizar, misturar. Tudo o resto é uma projeção.

Quase todas as explicações de dk\sqrt{d_k} dizem “para impedir a softmax de saturar”, o que é verdade e não explica nada. O argumento são duas linhas de variância do Capítulo 2. Se as entradas de q\mathbf{q} e k\mathbf{k} forem independentes, com média zero e variância um, cada produto qjkjq_j k_j tem variância um, e as variâncias de coisas independentes somam-se:

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

Portanto, 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

Porque isto importa: a softmax é sensível à escala de uma forma que uma camada linear não é. Duplicar a entrada de uma camada linear duplica a sua saída; multiplicar scores por dez antes de uma softmax transforma uma mistura suave numa escolha rígida. 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: por quantas posições a linha está realmente a fazer média. Sem divisão, em dk=256d_k = 256, uma cabeça recém-inicializada atende exatamente a um token em 64, escolhido apenas pelo sorteio aleatório.

Isso é mau no sentido direto e pior no sentido inverso, numa forma que o Capítulo 5 já mediu numa tanh\tanh. Uma softmax comprometida com uma entrada quase não tem derivada: a diagonal do seu jacobiano é 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 treino começar, e uma cabeça que começa congelada não consegue aprender para onde olhar. Com a divisão, a quantidade fica estável em 0,96 em todas as larguras e nada satura.

Agora a parte que ninguém publica: isto muda a perplexidade final? Elimine a divisão e treine, em quatro larguras de cabeça:

largura da cabeçasem divisãodividido por dk\sqrt{d_k}dividido por dkd_k
quatro cabeças, dk=32d_k = 3237,2938,0737,89
uma cabeça, dk=128d_k = 12848,5146,1045,99
uma cabeça, dk=256d_k = 25665,3747,53
uma cabeça, dk=512d_k = 51267,0649,15
uma cabeça, 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 cabeça, sem normalização antes das projeções — com ambas as variantes nas mesmas definições.

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

dkd_kdesvio-padrão 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 cabeça sem divisão não recupera. Dispara: o desvio-padrão dos seus scores passa de 21 na inicialização para 5147, a entropia de attention cai para zero, e 99,9 % das linhas colocam mais de 0,99 do seu peso num único token. Assim que uma cabeça se torna um seletor rígido, o seu gradiente fica quase zero e nada a puxa de volta, por isso o colapso é estável. A cabeça dividida fica num desvio-padrão de score de 3,44 após o mesmo treino, que é uma mistura suave que ainda pode ser alterada.

Vaswani et al. dizem exatamente isto e nada mais — suspeitam que os produtos “crescem 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 em 256.

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

Ligação para a secção: Mais do que uma opinião, e os dois terços de que ninguém fala

Uma cabeça é uma linha de softmax por posição, por isso contém uma resposta para “o que é relevante aqui”. Prever a palavra depois de o em o animal que atravessou a rua molhada precisa ao mesmo tempo do espaço sintático, do sujeito e do token anterior, e uma distribuição de probabilidade não consegue concentrar-se em três lugares. Portanto, execute várias cabeças em paralelo, cada uma com largura dmodel/hd_{\text{model}}/h, concatene e misture com mais uma matriz WoW_o: particionou a largura, não a aumentou.

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

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

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

Pesos aprendidos batem pesos uniformes por 14 pontos de perplexidade, que é todo o argumento deste capítulo numa linha. Quatro cabeças compram mais 3 por 16.512 parâmetros extra. E a mesma cabeça vale mais 9 pontos adicionada do que a substituir: attention traz informação para dentro, não decide o que uma posição é.

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

larguracabeçasattentionfeed-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 cada bloco transformer são a rede feed-forward, em todas as escalas, porque attention tem quatro matrizes d×dd \times d e a MLP tem o equivalente a oito. Seja o que for que um modelo saiba, a maioria dos parâmetros que o guardam está na MLP por posição.

LayerNorm foi construído e medido no Capítulo 6, e este capítulo usa-o como foi deixado lá; as ligações residuais foram nomeadas e ablated lá, e são construídas aqui. As linhas “adicionada, não a substituir” acima são ligações residuais, valendo 188 pontos de perplexidade para a média e 9 para uma cabeça. LayerNorm7 normaliza cada exemplo através das suas características, e o Capítulo 6 deu as razões pelas quais ele, e não BatchNorm, sobreviveu aqui — sem dependência do batch, sem estatísticas correntes, idêntico em treino e inferência, indiferente ao comprimento da sequência — cada uma das quais se torna um requisito quando se gera um token de cada vez para um utilizador, que é onde o Capítulo 13 acaba. Custa 768 parâmetros e compra 1,8 pontos 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 fica a normalização: na entrada de cada subcamada, com o caminho residual da entrada à saída nunca normalizado. Isto é 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 através do gradiente na inicialização, que numa rede post-norm fica mal escalado com a profundidade — a razão pela qual o transformer original precisava de um warmup da learning rate para sequer 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 exatamente pre-norm. Warmup não é aqui uma boa prática geral; é um remendo para uma disposição específica da normalização, e mover a LayerNorm remove a necessidade dele. É por isso que praticamente todos os modelos desde 2019 são pre-norm, e porque o diagrama de 2017 deve ser lido como história, não como especificação.

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

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

posiçõesperplexidade em 64em 128em 256
nenhuma48,7952,6357,52
embeddings absolutos aprendidos38,63108,47181,94
sinusoides fixos42,9695,26152,25
RoPE44,1250,5284,84
ALiBi44,9543,5142,49

Embeddings absolutos aprendidos — um vetor por posição, adicionado ao token — vencem no comprimento treinado e depois caem a pique, porque a posição 100 nunca esteve num batch e o seu embedding continua a ser o vetor aleatório com que começou. Sinusoides, a escolha original, são calculados em vez de aprendidos, 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 lá. RoPE9 não adiciona nada e, em vez disso, roda query e key por um ângulo proporcional à posição, em fatias bidimensionais; como rodar igualmente os dois lados de um produto escalar o deixa inalterado, o score acaba por depender apenas de tit - i, por isso a posição torna-se relativa de graça e não há tabela que se esgote. Degrada-se, mas degrada-se. ALiBi10 é o resultado mais simples e mais estranho aqui: uma penalização linear no score proporcional à distância, com uma inclinação diferente por cabeça. A sua perplexidade melhora à medida que a janela cresce para lá do comprimento de treino, de 44,95 para 42,49, porque a penalização está definida a qualquer distância e cada cabeça continua a fazer 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 esse intervalo, e é o segundo que morde. É também a maquinaria por trás de cada anúncio “alargámos o contexto para 128K” — quase sempre são reescalamentos de uma codificação rotativa, e são a razão pela qual o Capítulo 16 diz que o limite de contexto se move em vez de desaparecer.

Dropout é herdado da mesma forma: aparece nos pesos de attention depois da softmax, na saída de cada subcamada antes da adição residual, e na soma dos embeddings, fazendo exatamente o que o Capítulo 6 descreveu. Em grandes execuções de pré-treino, é 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 forma n×nn \times n, onde nn é o número de tokens: os scores e os pesos depois da softmax. Tudo o resto — cada projeção, toda a MLP — é linear em nn.

Uma camada de attention, 512 de largura, 8 cabeças, batch de um, float32, numa GPU de portátil. Leia as duas colunas de milissegundos apenas pelos seus rácios: são tempo real numa placa de portátil de 8 GB que abranda de 1.785 MHz para menos de 300 MHz quando aquece, pelo que uma execução fria do mesmo código regressa sete a dez vezes mais depressa, e uma ocupada ainda mais devagar. As colunas de megabytes são contagens de bytes do alocador e não se mexem.

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 o rácio face à linha acima, e duplicar nn converge exatamente para 4 tanto em tempo como em memória — 3,91 no último passo contra um 4 teórico. A coluna das projeções é o controlo: 4,0 ms em 1024 tokens para 40,1 ms em 8192, um fator de dez para um fator de oito. Linear, como anunciado.

Depois a última linha. Uma camada de attention, uma sequência, sem modelo à volta, esgota a memória numa GPU de 8 GB aos 16.384 tokens — só a matriz de scores teria 8 GB, sendo 8 cabeças vezes 16.384 vezes 16.384 vezes 4 bytes. Não o modelo; um tensor intermédio numa camada.

Esse é o facto físico por baixo de três capítulos posteriores. É por isso que uma context window tem sequer 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 uma de velocidade.11 E é a aritmética por trás do preço de um prompt longo, que o Capítulo 24 paga num ciclo de agent — uma questão separada da outra conclusão desse capítulo, a de que um modelo também usa pior um contexto longo, que ele mede e se recusa a culpar nesta fórmula.

Mostrar detalhes

As duas variantes que reduzem a 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 cabeça e por camada. Multi-query attention12 mantém hh projeções de query, mas uma única projeção de key e value partilhada por todas as cabeças, dividindo essa cache por hh. Grouped-query attention13 interpola: as cabeças são agrupadas, cada grupo partilhando uma key e um value, por isso g=hg = h é attention normal e g=1g = 1 é multi-query. Quase todos os modelos abertos desde 2023 a usam com 4 ou 8 grupos. Nenhuma existe por qualidade; ambas existem pelo tamanho dessa cache, e o Capítulo 13 faz a aritmética que a transforma em “que modelo cabe na sua GPU”.

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

O que venceu foi a metade decoder-only — uma pilha, causal do princípio ao fim, entrada e saída na mesma sequência — e a razão não é elegância. “Prever o próximo token” funciona em qualquer texto, por isso o conjunto de treino é a internet em vez de um corpus paralelo, e tudo se torna essa única tarefa: uma tradução é um documento que contém a fonte e depois o alvo, uma pergunta e a 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. Os codificadores não desapareceram — um codificador vê toda a entrada de uma vez, que é o que se quer quando o trabalho é representar um texto em vez de o continuar, e é por isso que os retrieval embeddings do Capítulo 19 vêm de codificadores e não do modelo que conversa.

Com o bloco definido, o tamanho do modelo é aritmética. Por bloco, com largura dd e expansão por quatro: 4d2+4d4d^2 + 4d para Wq,Wk,Wv,WoW_q, W_k, W_v, W_o com biases em todas as quatro, como o GPT-2 os tem — a tabela acima deixa o bias fora de três delas, daí menos 2.304 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 token de V×dV \times d e, para posições absolutas, nctx×dn_{\text{ctx}} \times d. Para a forma do GPT-2 small — d=768d = 768, 12 blocos, um vocabulário de 50.257, um contexto de 1024, a camada de saída a partilhar os pesos de 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; é o modelo. Note também que quase um terço de um modelo pequeno é a tabela de embedding, razão pela qual o tamanho do vocabulário é uma decisão arquitetural e não de pré-processamento — o trade-off que o Capítulo 7 preparou.

Perplexidade é um número sobre um corpus. O que uma cabeça faz é uma pergunta diferente, e um modelo treinado num megabyte de Shakespeare é o instrumento errado para ela: a coisa honesta a dizer sobre o mapa de attention de um modelo de 500.000 parâmetros é que é, na maior parte, não interpretável. Portanto: uma linguagem em que a pergunta tem uma resposta certa.

A ilustração clássica é o animal não atravessou a rua porque estava demasiado cansado, onde estava se refere ao animal, contra …porque estava demasiado molhada, onde uma palavra desloca o referente para a rua. Estes são esquemas de Winograd14 — pares de frases idênticos exceto por uma palavra, em que essa palavra decide a que se refere um pronome.

Também são solúveis fazendo batota, que é a parte que os tutoriais saltam. Se os dois candidatos são um animal e um lugar, cansado e molhada identificam o referente por categoria, e um modelo que só sabe que palavras estão presentes acerta sem saber nada sobre ordem. Medido nessa versão da tarefa, com pares animal/lugar retidos:

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

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

Feche então a brecha: retire ambos os candidatos de um conjunto de dezasseis substantivos, qualquer um dos quais pode aparecer em qualquer posição, e divida os adjetivos por papel em vez de categoria — quatro que tornam isso o atravessador (cansado, assustado, lento, fraco), quatro que tornam isso o atravessado (molhado, largo, movimentado, íngreme).

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 por isso o — e construa o conjunto retido a partir de pares de substantivos cuja ordem inversa esteve no treino, para que qualquer coisa que saiba quais são os dois substantivos presentes, mas não qual veio primeiro, tenha de responder ao contrário.

modeloparâmetrosretidonomeia o outro substantivo
apenas o token atual5.7965,2 %5,2 %
média causal uniforme5.79627,9 %50,0 %
uma cabeça de attention aprendido18.08435,4 %64,6 %
quatro cabeças22.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 fica em 27,9 % e responde com o substantivo errado do par exatamente metade das vezes — a assinatura de algo que sabe que palavras estão lá e nada sobre a sua ordem, como o teste de baralhamento previu três secções atrás.

Agora o mapa: a attention na posição que tem de nomear o referente, média pelas quatro cabeças de cada bloco, para as duas frases que diferem numa 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, seja qual for o adjetivo. Isso não é uma falha, mas uma prova: na primeira camada, a query numa posição é uma função do token dessa própria posição e do índice, e o na posição 14 é o mesmo token nas duas frases. Uma cabeça de primeira camada não pode condicionar numa palavra que ainda não foi buscar. Portanto, o bloco 1 faz a única coisa útil disponível e arrasta o primeiro substantivo para a frente.

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

adjetivobloco 2 em animalem ruano adjetivoresposta
cansado, assustado, lento, fraco0,0000,0001,000animal
molhado, largo, movimentado, íngreme0,0000,4910,00–0,03rua

Para um adjetivo do atravessador, o segundo bloco gasta todo o seu peso no adjetivo, porque a resposta já está no fluxo residual — o bloco 1 pô-la lá — e tudo de que precisa é confirmação. Para um adjetivo do atravessado, vai buscar o outro substantivo. Isto é um circuito de dois saltos: uma cabeça move um candidato para a frente, uma cabeça numa camada posterior lê um token que decide se o mantém. A composição entre camadas é o mecanismo, e é por isso que um bloco chegou a 92,7 % e dois chegaram a 100 %.

É também a forma do circuito mais bem documentado em modelos reais. Induction heads — uma cabeça de token anterior a alimentar uma cabeça 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 uma grande parte da in-context learning, e formam-se num momento identificável durante o pré-treino. Este capítulo não tenta essa análise: delega-a, com ambos os artigos nas referências, porque ler circuitos de um modelo real é um campo de investigação e não uma secção.

Por fim, a implementação. As trinta linhas acima, com os seus pesos copiados dos próprios pesos do 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 por outra ordem, em precisão float32.

Tem agora a arquitetura de que todos os modelos no resto deste curso são feitos, e ela é menor do que a 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 não tem é um modelo que saiba alguma coisa, e empilhar não vai resolver isso por si só. Dois blocos neste corpus chegam a uma perplexidade de treino de 14,49 e a 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 faz disto um problema de contabilidade, e a contabilidade é mais estranha do que parece. Quanto texto, e onde é que alguém o arranja? Quanta aritmética, e como a estima antes de o dinheiro ser gasto? Dado um orçamento fixo, é melhor tornar o modelo maior ou mostrar-lhe mais dados — e há uma resposta correta, ou apenas uma moda? O Capítulo 10 responde às três por medição, e põe 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 do que esta naquilo a que se destinam, e este capítulo foi escrito para ser lido ao lado delas. The Illustrated Transformer, de Jay Alammar, é a melhor imagem do fluxo de dados alguma vez desenhada. The Annotated Transformer, da Harvard NLP, é o artigo de 2017 com código executável intercalado linha a 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 dorsal medida noutro 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. and Schmidhuber, J. Long Short-Term Memory. Neural Computation 9(8), pp. 1735–1780 (1997).

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

  3. Bahdanau, D., Cho, K. and 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. Todos os números aqui usam o mesmo tokenizador 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, Ł. and Polosukhin, I. Attention Is All You Need. arXiv:1706.03762 (2017). A secção 3.2.1 é a frase única sobre dk\sqrt{d_k} que este capítulo passa uma secção a medir.

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

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

  8. Xiong, R., Yang, Y., He, D., Zheng, K., Zheng, S., Xing, C., Zhang, H., Lan, Y., Wang, L. and Liu, T.-Y. On Layer Normalization in the Transformer Architecture. arXiv:2002.04745 (2020). A análise de gradiente 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. and Liu, Y. RoFormer: Enhanced Transformer with Rotary Position Embedding. arXiv:2104.09864 (2021).

  10. Press, O., Smith, N. A. and 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. and 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. and Sanghai, S. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. arXiv:2305.13245 (2023).

  14. Levesque, H. J., Davis, E. and Morgenstern, L. The Winograd Schema Challenge. KR (2012). A construção por trás da frase animal / rua que todos os tutoriais de attention usam.

Pronto para deixar a LIA escolher?

Construa com todos os modelos de IA num só sítio — comece grátis hoje.