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

Predição do próximo token: embeddings e o que perplexidade significa

Treine um modelo de caracteres em 32.033 nomes, veja gradient descent redescobrir contagens e entenda por que perplexidades raramente batem.

Nesta página

Aqui estão dez nomes produzidos por um programa que nunca viu uma palavra:

TEXT
cexze   momakurailezitynn   konimittain   llayn   ka
da      moliellavo          emia          sade    ftlsp

Nenhum deles é um nome. Quase todos estão tentando. Eles são pronunciáveis, terminam onde nomes terminam, e um deles — emia — está a uma única letra de um nome real. O programa que os produziu guarda 729 números, não tem noção de palavra, sílaba ou pessoa, e foi ajustado por uma única passada de contagem de pares adjacentes de letras.

Até o fim deste capítulo, uma rede neural terá reduzido a pontuação desse programa em um terço na mesma medição. A parte que vale acompanhar é o que a rede faz primeiro: ela reproduz a tabela de contagens com três casas decimais em todas as linhas bem povoadas, sem prompt, porque os dois objetos são respostas para a mesma pergunta. Tudo depois disso é o que a contagem jamais poderia ter feito.

O objetivo é uma identidade, não uma escolha de design

Link para a seção: O objetivo é uma identidade, não uma escolha de design

O Capítulo 7 deixou você com uma sequência de inteiros e nenhum motivo para um seguir o outro. Aqui está o motivo, e ele é uma linha do Capítulo 2.

Um modelo de linguagem é uma função que recebe os tokens até agora e retorna uma distribuição sobre qual token vem a seguir: um número por entrada do vocabulário, não negativo, somando um. Nada mais. Para sair disso e chegar a uma probabilidade para um documento inteiro, aplique a regra da cadeia da probabilidade:

P(x1,x2,,xT)=t=1TP(xtx1,,xt1)P(x_1, x_2, \ldots, x_T) = \prod_{t=1}^{T} P(x_t \mid x_1, \ldots, x_{t-1})

Isso é uma identidade, verdadeira para qualquer sequência de qualquer coisa, sem pressupostos anexados. Portanto, um modelo que faz o trabalho pequeno — próximo token dados os anteriores — já fez o trabalho grande de atribuir uma probabilidade a todo documento possível, exatamente e de graça. O enquadramento popular disso como um truque barato («ele só prevê a próxima palavra») inverte a lógica: prever o próximo token é modelar a distribuição conjunta. Nunca houve uma segunda coisa a fazer.

A loss vem de modo igualmente mecânico. Em cada posição, o modelo produz uma distribuição qq e a verdade é um único token conhecido, então a entropia cruzada do Capítulo 4 se aplica sem alteração:

L=1Tt=1Tlogqθ(xtx<t)L = -\frac{1}{T}\sum_{t=1}^{T} \log q_\theta(x_t \mid x_{<t})

Essa é a log-verossimilhança negativa média — a receita do Capítulo 2 com uma distribuição categórica no lugar onde estava a gaussiana. E, como a distribuição verdadeira é one-hot, sua entropia é zero; portanto, pela identidade do Capítulo 4, a entropia cruzada é igual à divergência KL: baixar esse número e puxar as crenças do modelo na direção dos dados são o mesmo ato.

Uma consequência merece sua própria frase, porque é o fato econômico por baixo de todo o campo. Os rótulos são os dados, deslocados em uma posição. Ninguém anota nada. Um trilhão de tokens de texto é um trilhão de exemplos pré-rotulados, motivo pelo qual o corpus de treinamento de um modelo moderno é «a internet» e não «um dataset que alguém construiu».

Antes de qualquer rede, a baseline: 32.033 nomes, um por linha, e a tarefa de produzir mais nomes, uma letra de cada vez.1

O vocabulário tem 26 letras mais um símbolo de fronteira . marcando tanto o começo quanto o fim de um nome, então o modelo precisa aprender onde nomes começam e onde param. São 27 símbolos, e o menor modelo possível é uma tabela de quantas vezes cada símbolo veio depois de cada outro símbolo.

bigram.pyPYTHON
N = torch.zeros((27, 27), dtype=torch.int32)
for w in words:
    cs = ["."] + list(w) + ["."]
    for a, b in zip(cs, cs[1:]):
        N[stoi[a], stoi[b]] += 1

P = N.float()
P = P / P.sum(1, keepdim=True)            # one distribution per row   

Duas linhas de aritmética e o modelo está ajustado — e não é uma heurística: dividir as contagens pelos totais das linhas é a estimativa de máxima verossimilhança para uma distribuição categórica, que é a receita do Capítulo 2 com o cálculo já feito.

TEXT
names: 32033        train/val/test: 25626 / 3203 / 3204
training bigrams: 182583

the six most likely letters after 'a':
    a -> '.'  0.1944   a -> 'n'  0.1600   a -> 'r'  0.0967
    a -> 'l'  0.0749   a -> 'h'  0.0690   a -> 'y'  0.0606

Amostre dela — escolha uma letra da linha da letra atual, mova para essa linha, repita até o símbolo de fronteira aparecer — e você obtém os nomes do início deste capítulo. Eles falham de um jeito específico e informativo: localmente plausíveis, globalmente sem sentido. Todo par adjacente de letras em momakurailezitynn é um par que ocorre em nomes reais; acontece que há dezessete deles em sequência. O modelo tem uma letra de memória, então não consegue saber que já está indo longe demais.

A loss em nomes reservados é 2,4546 nats. Esse número não significa nada sozinho, e é por isso que a perplexidade existe:

PPL=exp ⁣(1Ttlogq(xtx<t))=eL\mathrm{PPL} = \exp\!\left(-\frac{1}{T}\sum_t \log q(x_t \mid x_{<t})\right) = e^{L}

Escrito por extenso, sem nenhuma biblioteca fazendo o trabalho:

perplexity.pyPYTHON
@torch.no_grad()
def perplexity(logits, Y):
    logp = F.log_softmax(logits, dim=1)          # log q for every symbol
    chosen = logp[torch.arange(len(Y)), Y]       # log q of the one that came next   
    return torch.exp(-chosen.mean())             

Exponenciar desfaz o logaritmo e devolve o número às unidades de contar coisas. A forma limpa de ver o que ele conta é medir um modelo que não sabe absolutamente nada — um que atribui probabilidade 1/271/27 a todo símbolo, independentemente do contexto:

TEXT
uniform over 27 symbols            loss 3.2958 nats   ppl  27.000
bigram counts, add-one smoothed    loss 2.4546 nats   ppl  11.642

Exatamente 27,000, porque elog27=27e^{\log 27} = 27. Perplexidade é o número efetivo de opções igualmente prováveis entre as quais o modelo está escolhendo. Uma perplexidade de 27 significa «não faço ideia, pode ser qualquer coisa». Os 11,642 do modelo de contagem significam que uma letra de contexto o deixa tão incerto quanto alguém escolhendo às cegas entre cerca de doze opções em vez de vinte e sete — por isso a perplexidade é citada e a loss bruta não.

Duas coisas dão errado com ela, e a segunda dá errado em artigos publicados.

Probabilidades zero são fatais. Das 729 células da tabela, 113 nunca ocorrem no treinamento — 15,5 % dela está vazia. Isso é aceitável até o conjunto reservado cair em uma delas, e sete bigramas na validação caem, entre eles dq, zj e qo duas vezes. Probabilidade zero significa log -\infty, o que significa loss infinita e perplexidade infinita: um nome em três mil destrói a métrica. O remendo usual é somar 1 a toda contagem antes de normalizar, o que custa quase nada aqui (2,4546 em vez de 2,4524). Mas o remendo é uma confissão. Um modelo de contagem não consegue generalizar de jeito nenhum. Ele não tem como suspeitar que qo é plausível porque qu é comum e o se comporta como u em outros lugares, já que não tem noção de que dois símbolos podem se parecer. Cada célula é aprendida isoladamente, e corrigir isso é o motivo do restante deste capítulo.

Perplexidade é um preço por token, e o token é um parâmetro livre. Esse é o erro que aparece constantemente quando modelos são comparados, e fica fácil de enxergar quando você olha. Pegue o mesmo corpus de prosa em inglês do Capítulo 7, o mesmo modelo de bigramas interpolado, e mude apenas como o texto é recortado:

unidadevocabuláriotokens no testeentropia cruzadaperplexidadebits por caractere
caracteres7614.4692,521712,453,6378
BPE, 512 merges3296.8713,854747,212,6407
BPE, 2.048 merges1.8204.2335,7468313,202,4254
palavras2.9916.2843,562735,262,2322

A perplexidade varia por um fator de 25 entre essas linhas. Nada sobre o modelo mudou; só o tamanho da coisa sendo prevista. Prever uma palavra inteira é mais difícil do que prever uma letra, então custa mais por previsão — e há menos previsões a fazer.

Agora leia a última coluna, que divide o custo total pelo número de caracteres em vez disso e o converte para bits. Ela reordena a tabela. Pela perplexidade, o ranking é caracteres, palavras, BPE-512, BPE-2048; por bits por caractere, é palavras, BPE-2048, BPE-512, caracteres. O modelo de caracteres vai do primeiro ao último lugar. O modelo de 2.048 merges, que pela perplexidade parece 6,6 vezes pior que o de 512 merges, é na verdade o melhor dos dois, com 2,4254 bits contra 2,6407.

Portanto, uma perplexidade só é comparável entre dois modelos que compartilham um tokenizer, e modelos com tokenizers diferentes só podem ser comparados em bits por caractere — a quantidade que Shannon mediu em 1951 ao pedir que pessoas adivinhassem a próxima letra de um texto em inglês, e que ele limitou a aproximadamente um bit por caractere.2 Nosso melhor bigrama está em 2,23 bits, um bom resumo de quanto este capítulo ainda precisa avançar.

Agora construa o mesmo modelo como uma rede. Ele vai exigir ordens de grandeza mais aritmética para chegar ao mesmo lugar, e chegar ao mesmo lugar é o ponto.

Substitua a tabela por uma matriz de pesos WW com forma 27×2727 \times 27. Transforme a letra atual em um vetor one-hot, multiplique e chame o resultado de logits — as pontuações não normalizadas do Capítulo 4. Depois softmax, depois entropia cruzada, depois gradient descent.

neural_bigram.pyPYTHON
W = torch.randn((27, 27), requires_grad=True)

for step in range(3000):
    logits = W[xs]                            
    loss = F.cross_entropy(logits, ys)
    W.grad = None
    loss.backward()
    W.data -= 50.0 * W.grad

A linha destacada contém uma definição que vale guardar. Multiplicar um vetor one-hot por uma matriz seleciona uma linha dela, então a multiplicação é uma consulta — e toda implementação pula a aritmética e faz a consulta diretamente, que é o que W[xs] é.

Isso é uma tabela de embedding. Uma matriz com uma linha por entrada do vocabulário, indexada pelo id do token. Sem geometria, sem semântica, sem algoritmo separado: uma tabela de consulta cujo conteúdo por acaso é aprendido por gradient descent junto com todo o resto. Toda afirmação mística sobre «espaço de embedding» termina aqui.

Treine e observe para onde ela vai:

TEXT
  step     1   train 3.7550   val 3.3882   max gap to the count table 0.757269
  step   100   train 2.4732   val 2.4726   max gap to the count table 0.388354
  step  1000   train 2.4557   val 2.4549   max gap to the count table 0.041862
  step  3000   train 2.4547   val 2.4544   max gap to the count table 0.004048

A última coluna é a maior diferença absoluta entre qualquer célula de softmax(W) e a célula correspondente da tabela de contagens, e ela vai a zero. Depois de 3.000 passos, a maior discordância em qualquer uma das 729 células é 0,004048 e a média é 0,000224. A pior célula é qi, vista doze vezes em todo o conjunto de treinamento; entre as 22 linhas com mais de mil ocorrências, a pior discordância é 0,000562.

TEXT
                 count table   network
    a -> '.'        0.1945     0.1945
    a -> 'n'        0.1601     0.1601
    a -> 'r'        0.0967     0.0967

Gradient descent, começando de números aleatórios e sem receber nada além de «aumente a log-probabilidade da próxima letra», redescobriu a tabela de contagens. E tinha que fazê-lo: as contagens são a estimativa de máxima verossimilhança, a entropia cruzada é a log-verossimilhança negativa, então os dois procedimentos otimizam o mesmo objetivo e esse objetivo tem um único ótimo. A rede não aprendeu algo parecido com contagem. Ela convergiu para a contagem, lentamente.

Isso levanta a pergunta justa: por que alguém se daria ao trabalho? Porque a tabela de contagens não tem para onde ir a partir daqui, e a rede tem.

Estenda o modelo para olhar mais de um caractere anterior. Esta é a arquitetura de Bengio de 2003, a ancestral direta de todos os modelos no restante deste curso:4 pegue os três últimos caracteres, mapeie cada um por uma tabela de embedding para uma linha de 10 dimensões, concatene as linhas em 30 números, passe-os pela camada oculta do Capítulo 5 e termine com uma camada de saída produzindo um logit por entrada do vocabulário.

mlp.pyPYTHON
C  = torch.randn((27, 10))          # the embedding table
W1 = torch.randn((3 * 10, 200))     # the hidden layer from Chapter 5
W2 = torch.randn((200, 27))         # one output per vocabulary entry

emb = C[X].view(-1, 30)             # three lookups, concatenated   
h = torch.tanh(emb @ W1 + b1)
logits = h @ W2 + b2                
loss = F.cross_entropy(logits, Y)

Observe o que é novo e o que não é. A camada oculta é a do Capítulo 5, sem mudanças; a loss é a do Capítulo 4, sem mudanças. As novidades são a tabela de embedding na frente e uma camada de saída tão larga quanto o vocabulário do Capítulo 7 — e essa segunda é a parte cara de todo modelo de linguagem já construído, porque uma vocabulário real tem 100.000 entradas e essa multiplicação de matrizes roda em toda posição.

O mesmo código, treinado de forma idêntica, mudando apenas o tamanho da context window:

contextoparâmetrosloss de validaçãoperplexidade de validação
contagem, 1 caractere7292,454611,642
neural, 1 caractere7.8972,457711,678
neural, 3 caracteres11.8972,11458,285
neural, 8 caracteres21.8972,05067,773

A segunda linha é a interessante. Uma rede com uma camada oculta de 200 unidades e onze vezes mais parâmetros que a tabela de contagens se sai exatamente tão bem quanto a tabela de contagens e nada melhor. A capacidade nunca foi a limitação. Um caractere de contexto permite uma certa loss e nada que você acople consegue ficar abaixo dela, porque a informação não está ali.

Dê a ela três caracteres e a perplexidade cai de 11,68 para 8,29 — uma redução de 29 %, comprada com 4.000 parâmetros extras. Ela vence a contagem aqui exatamente pelo motivo diagnosticado antes: um modelo de contagem sobre contextos de três caracteres precisa de 273=19,68327^3 = 19{,}683 linhas, a maioria vazia ou contendo uma única observação, e aprende cada uma isoladamente. A rede compartilha. Se a, e e i acabam com linhas de embedding semelhantes, o que ela aprende depois de bra se transfere para bre sem que ela jamais tenha visto bre. Essa transferência é todo o valor da tabela de embedding, e é a diferença entre as linhas dois e três.

As amostras melhoram na mesma proporção:

TEXT
deliah   nellara   joce     kael      quintis
salayson  reety    khyrmin  mahnen    madiaryxia

Ainda não é uma lista de nomes reais. Mas deliah, nellara e kael não pareceriam fora de lugar em uma, e os monstros sem fim desapareceram: o mais longo de vinte exemplos do modelo de contagem tem dezenove letras, o mais longo de vinte deste tem treze.

O que há de fato dentro da tabela de embedding

Link para a seção: O que há de fato dentro da tabela de embedding

A tabela é 27×1027 \times 10: uma linha de dez números por caractere, todos inicializados aleatoriamente e movidos apenas pelo gradiente da loss do próximo caractere. Ninguém colocou nada ali. Então o que acabou nela?

A ferramenta para perguntar é similaridade cosseno, que é o produto escalar do Capítulo 1 com os comprimentos divididos para fora:

cos(a,b)=abab\cos(\mathbf{a}, \mathbf{b}) = \frac{\mathbf{a} \cdot \mathbf{b}}{\lVert \mathbf{a} \rVert \, \lVert \mathbf{b} \rVert}

Ela mede o ângulo entre dois vetores e ignora seus comprimentos, que é o que você quer quando o comprimento de uma linha reflete a frequência com que seu token apareceu, e não o que ele significa. Normalize primeiro todo vetor para comprimento 1 — como sistemas reais fazem, uma vez, no momento da indexação — e a similaridade cosseno é simplesmente o produto escalar.

Aqui estão os vizinhos mais próximos de alguns caracteres na tabela treinada:

TEXT
  'c' -> 'k':+0.598      'j' -> 'z':+0.650      'i' -> 'y':+0.541
  'u' -> 'e':+0.482      'a' -> 'h':+0.367      '.' -> 'q':+0.077

Parte disso é o que o folclore promete. c e k são intercambiáveis em nomes, assim como i e y; j e z são ambas consoantes raras, majoritariamente iniciais, que se comportam de forma parecida. O símbolo de fronteira . não fica perto de quase nada — 0,077 até a letra mais próxima — porque é o único símbolo que marca uma posição em vez de um som.

E parte não é. O vizinho mais próximo de a é h, não outra vogal. Em média, sobre todos os pares:

TEXT
mean cosine, vowel to vowel         : +0.1889
mean cosine, consonant to consonant : +0.0765
mean cosine, vowel to consonant     : -0.0042

As vogais são mais parecidas entre si do que com consoantes, e o efeito é real, mas pequeno. Testado contra 2.000 grupos de cinco letras escolhidos aleatoriamente, 58 desses grupos se separam pelo menos com a mesma nitidez — uma diferença significativa em cerca de p=0.03p = 0.03. Real, portanto, e nada parecida com a ilha geométrica nítida que relatos populares sobre embeddings sugerem.

Essa é a descrição honesta de uma tabela de embedding, e vale guardá-la para o restante do curso. Ela não é um mapa de significado. É uma mudança de coordenadas, aprendida em vez de projetada, cuja única função é facilitar o trabalho da próxima camada — a mesma frase que o Capítulo 5 usou para a camada oculta que dobrou o plano para resolver XOR. Qualquer estrutura que você encontrar nela está ali porque reduziu a loss, e estrutura que não reduz a loss simplesmente não está ali.

word2vec, GloVe e a aritmética que todo mundo cita

Link para a seção: word2vec, GloVe e a aritmética que todo mundo cita

Se a parte útil é a tabela, você pode ir atrás dela diretamente. Isso é word2vec: mantenha a consulta de embedding, descarte o modelo de linguagem.5

O objetivo skip-gram com negative sampling é uma linha. Para um par real (centro, contexto) tirado do corpus, empurre o produto escalar para cima; para kk pares falsos tirados de uma distribuição de ruído, empurre-o para baixo:6

logσ(vcvo)+i=1klogσ(vcvni)\log \sigma(\mathbf{v}_c \cdot \mathbf{v}_o) + \sum_{i=1}^{k} \log \sigma(-\mathbf{v}_c \cdot \mathbf{v}_{n_i})

Isso é uma classificação binária — «estas duas palavras realmente ocorreram juntas?» — e é barato justamente porque nunca toca o vocabulário inteiro, o que tornou prático treinar em bilhões de palavras em 2013. O GloVe chega a vetores semelhantes pela outra direção, fatorando a matriz de contagens de coocorrência global em vez de percorrer exemplos em fluxo.7 Ambos são ajustados exatamente à estatística da qual a tabela de contagens foi construída. Eles são contagem, comprimida.

Treinados em text8 — 17.005.207 palavras da Wikipédia em inglês, 71.290 delas ocorrendo pelo menos cinco vezes, 100 dimensões, três passadas — os vetores saem com a propriedade que os tornou famosos:

TEXT
king     -> charles 0.700, son 0.693, queen 0.686, henry 0.669, throne 0.667
physics  -> chemistry 0.672, electromagnetism 0.661, quantum 0.654, theoretical 0.624
guitar   -> bass 0.733, vocals 0.732, acoustic 0.728, guitars 0.703, drums 0.685
three    -> seven 0.892, two 0.877, one 0.875, five 0.871, four 0.870

Ninguém forneceu uma categoria para instrumentos ou numerais. Agora a parte famosa: pegue king, subtraia man, some woman e encontre o vetor mais próximo do resultado.

TEXT
king - man + woman
   nothing excluded : king 0.693, elizabeth 0.657, wife 0.629, woman 0.607
   a, b, c excluded : elizabeth 0.657, wife 0.629, mary 0.607   (queen is 4th, 0.604)

O vetor mais próximo de king - man + woman é king. Isso não é uma peculiaridade de um exemplo. O conjunto de avaliação de Mikolov propõe perguntas da forma a : b :: c : ? — 8.869 semânticas (paris : france :: rome : italy) e 10.675 sintáticas (walking : walked :: swimming : swam) — e, entre as 4.103 perguntas semânticas que este vocabulário consegue responder, o vencedor é uma das três palavras de entrada 99,8 % das vezes. As demonstrações publicadas não mencionam isso, porque a regra de pontuação padrão remove a, b e c antes de procurar. É uma regra legítima, e ela está fazendo mais trabalho que a aritmética:

como a resposta é escolhidasemânticasintática
offset, com as entradas excluídas (padrão)17,0 %11,9 %
offset, sem nada excluído0,1 %0,4 %
vizinho mais próximo de c sozinho, entradas excluídas13,1 %9,3 %
vizinho mais próximo de b sozinho, entradas excluídas2,3 %0,4 %

A terceira linha é a que merece atenção. Jogue fora a e b, não faça aritmética nenhuma, retorne o que estiver mais perto de c — e você mantém 77 % da pontuação semântica. A maior parte do que parece raciocínio analógico é proximidade mais uma regra que proíbe as respostas óbvias, que é o que Linzen mediu em vetores treinados corretamente e o que as baselines acima replicam.8 Estes vetores específicos são pequenos — 17 milhões de palavras contra os bilhões por trás dos modelos publicados — então leia as porcentagens como uma forma, não como estado da arte. A forma é o que sobrevive em toda escala: a aritmética é real, e muito mais fraca que a demonstração que todo mundo cita.

Estático e contextual: um vetor por palavra, ou um por ocorrência

Link para a seção: Estático e contextual: um vetor por palavra, ou um por ocorrência

Tudo até aqui tem um limite rígido embutido na estrutura de dados. Uma tabela tem uma linha por token. A palavra bank recebe um vetor, o mesmo em uma frase sobre um rio e em uma frase sobre uma hipoteca — necessariamente, porque uma consulta por id não pode depender de mais nada.

A correção é parar de ler o vetor da tabela e começar a computá-lo a partir da frase. Isso é um embedding contextual, introduzido pelo ELMo em 2018 e tornado padrão pelo BERT no mesmo ano.910 Medidos no modelo real, os números são mais claros que a explicação:

TEXT
sentence A: "He sat on the bank of the river and watched the water go by."
sentence B: "She deposited the cheque at the bank on the corner of the street."

static vector for 'bank' (a row of the input embedding table)
    cosine A vs B ........................ 1.000000

contextual vector for 'bank', layer by layer
    layer  |  A vs B  |  A vs another river sentence  |  B vs another money sentence
        0  |  0.9512  |            0.9512             |            0.9359
        4  |  0.5647  |            0.8987             |            0.7716
        9  |  0.4284  |            0.8699             |            0.7568
       12  |  0.5278  |            0.8702             |            0.7335

A primeira linha é exata, não aproximada: o vetor estático de bank é os mesmos 768 números nas duas frases, então o cosseno é 1 por construção. Nove camadas depois, as duas ocorrências ficam em 0,43, enquanto bank em duas frases diferentes sobre rio permanece em 0,87. Ninguém rotulou um sentido em momento algum desse processo; os sentidos se separaram porque separá-los torna o objetivo de treinamento — adivinhar um token oculto a partir de seus vizinhos — mais fácil de satisfazer.

Dois detalhes recompensam a atenção. A camada 0 já é 0,9512 em vez de 1,0 porque embeddings de posição foram adicionados e a palavra está em um lugar diferente em cada frase. E a similaridade sobe de novo nas camadas 11 e 12: as camadas finais de um modelo pré-treinado são especializadas em seu objetivo de treinamento e frequentemente não são o melhor lugar de onde tirar uma representação.

Mostrar detalhes

Opcional: weight tying.

Em bert-base-uncased, a tabela de embedding é 30,522×76830{,}522 \times 768 — 23.440.896 números, 21,4 % dos 109.482.240 parâmetros do modelo. Em um modelo de linguagem pequeno, a fração é ainda maior, e é por isso que um truque é quase universal: a tabela de entrada e a camada de saída que produz os logits são a mesma matriz, usada uma vez por consulta de linhas e uma vez transposta.11 A camada de saída já atribui um vetor a cada entrada do vocabulário — ela faz um produto escalar contra cada uma — e o tying diz que o vetor usado para ler um token e o vetor usado para escrevê-lo deve ser o mesmo objeto. Ele reduz parâmetros e melhora a perplexidade de uma vez, o que é raro o bastante para merecer nota.

Um embedding model não é um modelo de linguagem

Link para a seção: Um embedding model não é um modelo de linguagem

Para buscar em um corpus por significado, você precisa de um vetor por frase. Com eles, a busca é trivial — isso é toda a recuperação semântica, e o Capítulo 19 trata de tudo ao redor dela:

search.pyPYTHON
E = normalise(embed(sentences))       # (200, d), every row of length 1
q = normalise(embed([query]))         # (1, d)
scores = q @ E.T                      # one matrix multiply   
top5 = scores[0].argsort()[::-1][:5]

Então a única pergunta real é de onde vem embed. O movimento óbvio é pegar um modelo de linguagem pré-treinado, passar cada frase por ele e tirar a média dos vetores de tokens. Aqui está esse método contra quatro alternativas, pontuado de duas formas: a correlação de ranking entre cosseno e julgamentos humanos de similaridade sobre os 1.379 pares do benchmark STS, e retrieval top-1 em um índice construído a partir dos 200 pares mais fortemente parafraseados desse conjunto — um lado de cada par indexado, o outro usado como consulta.

como a frase é embutidacorrelação de rankingtop-1 em um índice de 200 frases
sobreposição binária de palavras (sem modelo algum)0,550089,0 %
média dos vetores estáticos treinados acima0,526385,5 %
BERT, o token [CLS]0,203067,0 %
BERT, média dos vetores de tokens0,472984,0 %
MiniLM, treinado contrastivamente0,820392,0 %

Leia as três linhas do meio contra as duas primeiras. Um transformer pré-treinado de 109 milhões de parâmetros, usado do jeito óbvio, é pior em julgar similaridade entre frases do que contar quantas palavras duas frases compartilham — e pior do que tirar a média dos vetores text8 de 100 dimensões treinados há pouco. O token [CLS], que tutoriais ainda recomendam porque o BERT foi pré-treinado com um objetivo em nível de frase anexado a ele, é pior que metade disso.

Isso não é um defeito do BERT. É o objetivo. Um modelo de linguagem é treinado para que seus estados ocultos prevejam um token; nada ali pede que duas paráfrases acabem próximas uma da outra, e nada recompensa uma geometria em que cosseno significa «mesmo significado». A última linha é um modelo com um quinto do tamanho (22.713.216 parâmetros) treinado em uma loss completamente diferente: aprendizado contrastivo, em que os exemplos são pares — uma pergunta e sua resposta, uma frase e sua paráfrase — e o objetivo puxa pares verdadeiros para perto enquanto empurra negativos amostrados para longe. Essa é a contribuição do Sentence-BERT e a origem de toda a indústria de embedding models.12 Dense Passage Retrieval aplica a mesma receita diretamente à busca, com um encoder para consultas e outro para passagens.13

Então, a regra prática:

Um embedding model não é um modelo de linguagem com a última camada removida. É um modelo diferente em um objetivo diferente, geralmente muito menor, cujo cosseno significa o que você quer que ele signifique porque foi treinado em pares em que isso era o alvo. A tabela acima é o custo de substituir um pelo outro.

E a família falha na ordem das palavras. «The dog bit the man» e «the man bit the dog» têm bags of words idênticos, então a sobreposição de palavras e a média de vetores estáticos lhes dão cosseno exatamente 1,000000, e o BERT com mean pooling, que vê posição, ainda cai quase nisso — e o MiniLM treinado contrastivamente ainda os coloca em 0,979. Se sua tarefa de retrieval depende de quem fez o quê com quem, nenhum limiar de cosseno vai salvar você.

O Capítulo 19 constrói um sistema de retrieval em produção sobre esse fundamento e chega a um corte de cosseno concreto. A última medição deste capítulo é o que torna um número assim defensável em vez de mágico.

A maldição da dimensionalidade, em uma tabela

Link para a seção: A maldição da dimensionalidade, em uma tabela

Embeddings reais têm centenas ou milhares de componentes, e distâncias se comportam de forma estranha lá em cima. Pegue 1.000 pontos aleatórios no cubo unitário de dd dimensões e observe a razão entre a maior e a menor distância entre quaisquer dois deles:

dimensõespar mais próximopar mais distanterazão
20,00071,36121921,66
100,23612,33979,91
1003,00475,17521,72
1.00011,780914,03061,19
10.00039,615242,01251,06

Em dez mil dimensões, o par de pontos mais distante está apenas 6 % mais afastado que o par mais próximo. Tudo está aproximadamente equidistante de tudo, «vizinho mais próximo» deixa de carregar muita informação, e essa é a maldição da dimensionalidade — além de um motivo pelo qual grandes bancos de dados vetoriais não fazem busca exata de vizinho mais próximo. O outro lado da mesma moeda é o que torna limiares de cosseno viáveis: medido em mil pares de vetores unitários aleatórios, o cosseno médio fica em 0.0052-0.0052 em 100 dimensões e +0.0003+0.0003 em 768, com desvios-padrão de 0,0968 e 0,0357 — e, em 768 dimensões, apenas 0,2 % dos pares aleatórios excedem 0,1 em valor absoluto. Uma similaridade medida de 0,4, portanto, não é «40 % parecido»; ela está muito fora de qualquer coisa que o acaso produz, por isso limiares entre 0,3 e 0,7 separam sinal de ruído em vez de ficarem no meio dele.

O modelo deste capítulo lê um número fixo de caracteres anteriores, consulta cada um e cola os resultados em ordem. Esse design tem dois problemas, e eles são o mesmo problema.

Olhe de novo a tabela de contexto: ir de três caracteres para oito quase dobrou os parâmetros e comprou 0,06 nats. O custo cresce linearmente com o contexto — cada posição extra precisa de seu próprio bloco da primeira matriz de pesos — e o benefício não. Leve isso a mil tokens e só a primeira camada já pesa mais que o restante do modelo, com a maior parte gasta em posições que não importam para uma dada previsão.

Esse é o segundo problema: o modelo não tem como decidir quais dos tokens anteriores importam. A posição dois recebe seus próprios pesos e a posição sete recebe os seus, permanentemente, seja lá o que estiver nelas. Quando o modelo está soletrando nell, o caractere decisivo é o imediatamente anterior. Quando uma frase contém um pronome, a palavra que fixa seu referente pode estar quarenta tokens atrás — e nenhum slot fixo pode ser atribuído a «quarenta atrás», porque da próxima vez serão seis.

O que queremos é um modelo que compute, para cada previsão, quanto cada token anterior deve contar — pesos sobre o contexto produzidos pelo conteúdo, e não fixados pelo layout. Escreva isso com cuidado e ele começa como algo inteiramente mundano: uma média sobre os tokens anteriores. Depois deixe os pesos dessa média serem aprendidos, e deixe-os depender de qual token está fazendo a pergunta.

Isso é attention, e é o Capítulo 9.


Também vale ler junto: o capítulo 3 de Speech and Language Processing, de Jurafsky e Martin, que trata de modelos n-gram, smoothing e perplexidade com muito mais cuidado do que cabe aqui, incluindo por que interpolation e back-off vencem somar um; as notas de Stanford CS229 §17.1–17.2 para modelagem de linguagem pelo lado probabilístico; e o artigo de Linzen acima, que é curto e vale ser lido na íntegra.

  1. O exemplo de geração de nomes, o dataset e a progressão de uma tabela de contagens para uma rede no estilo de Bengio seguem a série building makemore de Andrej Karpathy, cujas duas primeiras partes são o melhor acompanhamento para este capítulo.

  2. Shannon, C. E. Prediction and Entropy of Printed English. Bell System Technical Journal 30(1), pp. 50–64 (1951). Pessoas tentando adivinhar a próxima letra do inglês, e a medição original de bits por caractere.

  3. Shannon, C. E. A Mathematical Theory of Communication. Bell System Technical Journal 27 (1948). O teorema de codificação de fonte e a identificação de predição com compressão.

  4. Bengio, Y., Ducharme, R., Vincent, P. and Jauvin, C. A Neural Probabilistic Language Model. Journal of Machine Learning Research 3, pp. 1137–1155 (2003). A arquitetura usada acima: um embedding por palavra, concatenado sobre uma janela fixa, passando por uma camada oculta até um softmax sobre o vocabulário.

  5. Mikolov, T., Chen, K., Corrado, G. and Dean, J. Efficient Estimation of Word Representations in Vector Space. arXiv:1301.3781 (2013). CBOW e skip-gram, e o conjunto de analogias usado acima.

  6. Mikolov, T., Sutskever, I., Chen, K., Corrado, G. and Dean, J. Distributed Representations of Words and Phrases and their Compositionality. arXiv:1310.4546 (2013). Negative sampling, subsampling de palavras frequentes e a distribuição de ruído elevada à potência 3/4 usada acima.

  7. Pennington, J., Socher, R. and Manning, C. GloVe: Global Vectors for Word Representation. EMNLP 2014. Vetores de palavras a partir de uma fatoração da matriz de coocorrência global, em vez de janelas locais em fluxo.

  8. Linzen, T. Issues in evaluating semantic spaces using word analogies. RepEval 2016, arXiv:1606.07736. A fonte das baselines sem offset replicadas acima.

  9. Peters, M. et al. Deep contextualized word representations. arXiv:1802.05365 (2018). ELMo: um vetor por ocorrência, computado por um modelo de linguagem bidirecional.

  10. Devlin, J., Chang, M.-W., Lee, K. and Toutanova, K. BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding. arXiv:1810.04805 (2018). O modelo medido no experimento com bank.

  11. Press, O. and Wolf, L. Using the Output Embedding to Improve Language Models. arXiv:1608.05859 (2016), and Inan, H., Khosravi, K. and Socher, R. Tying Word Vectors and Word Classifiers. arXiv:1611.01462 (2016). Dois argumentos independentes para o mesmo truque.

  12. Reimers, N. and Gurevych, I. Sentence-BERT: Sentence Embeddings using Siamese BERT-Networks. arXiv:1908.10084 (2019). Sua medição inicial — BERT com mean pooling ficando abaixo de vetores estáticos médios em similaridade de frases — é o que a tabela acima reproduz.

  13. Karpukhin, V. et al. Dense Passage Retrieval for Open-Domain Question Answering. arXiv:2004.04906 (2020). Treinamento contrastivo de um retriever de dois encoders; o ancestral direto da stack de retrieval do Capítulo 19.

Pronto para deixar a LIA escolher por você?

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