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

Previsão do próximo token: embeddings e o que significa a perplexidade

Treine um modelo de caracteres em 32.033 nomes, veja o gradient descent redescobrir contagens e perceba porque perplexidades raramente coincidem.

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 a tentar. São pronunciáveis, terminam onde os nomes terminam, e um deles — emia — está a uma única letra de um nome real. O programa que os produziu contém 729 números, não tem noção de palavra, sílaba ou pessoa, e foi ajustado por uma única passagem de contagem de pares adjacentes de letras.

No fim deste capítulo, uma rede neuronal terá reduzido a pontuação desse programa em um terço, na mesma medição. A parte que vale a pena acompanhar é o que a rede faz primeiro: 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 à mesma pergunta. Tudo o que vem depois é o que a contagem nunca poderia ter feito.

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

Ligação para a secção: O objetivo é uma identidade, não uma escolha de desenho

O Capítulo 7 deixou-o com uma sequência de inteiros e sem razão para um se seguir a outro. Aqui está a razão, e cabe numa linha do Capítulo 2.

Um modelo de linguagem é uma função que recebe os tokens até ao momento e devolve 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 passar daí para uma probabilidade de um documento inteiro, aplica-se 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})

Isto é uma identidade, verdadeira para qualquer sequência de qualquer coisa, sem pressupostos associados. Portanto, um modelo que faz a tarefa pequena — o próximo token dado os anteriores — já fez a tarefa grande de atribuir uma probabilidade a todos os documentos possíveis, exatamente e de borla. O enquadramento popular disto como um truque barato («só prevê a palavra seguinte») tem a lógica ao contrário: prever o próximo token é modelar a distribuição conjunta. Nunca houve uma segunda coisa a fazer.

A loss segue de forma igualmente mecânica. Em cada posição, o modelo produz uma distribuição qq e a verdade é um único token conhecido, por isso a cross-entropy do Capítulo 4 aplica-se sem alterações:

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

Isto é a log-verosimilhanç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, a sua entropia é zero, pelo que, pela identidade do Capítulo 4, a cross-entropy é igual à divergência KL: reduzir este número e aproximar as crenças do modelo das dos dados são o mesmo ato.

Uma consequência merece a sua própria frase, porque é o facto económico por baixo de toda a área. Os rótulos são os dados, deslocados uma posição. Ninguém anota nada. Um bilião de tokens de texto é um bilião de exemplos pré-rotulados, e é por isso que o corpus de treino 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 . que marca tanto o início como o fim de um nome, por isso o modelo tem de aprender onde os 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 se seguiu a 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 verosimilhanç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

Faça sampling a partir dele — escolha uma letra na linha da letra atual, passe para essa linha, repita até surgir o símbolo de fronteira — e obtém os nomes no topo deste capítulo. Falham de uma forma específica e informativa: localmente plausíveis, globalmente absurdos. Cada par adjacente de letras em momakurailezitynn é um par que ocorre em nomes reais; há simplesmente dezassete deles seguidos. O modelo tem uma letra de memória, por isso não consegue saber que já se prolongou demasiado.

A loss em nomes retidos para validação é 2,4546 nats. Esse número não significa nada por si só, e é por isso que existe a perplexidade:

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 biblioteca a fazer 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 rigorosamente nada — um que atribui probabilidade 1/271/27 a todos os símbolos, 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. A perplexidade é o número efetivo de opções igualmente prováveis entre as quais o modelo está a escolher. Uma perplexidade de 27 significa «sem ideia, podia ser qualquer coisa». Os 11,642 do modelo de contagem significam que uma letra de contexto o deixa tão incerto como alguém que escolhesse às cegas entre cerca de doze opções em vez de vinte e sete — e é por isso que se cita a perplexidade, não a loss bruta.

Duas coisas correm mal com ela, e a segunda corre mal em artigos publicados.

Probabilidades zero são fatais. Das 729 células da tabela, 113 nunca ocorrem no treino — 15,5 % dela está vazia. Isso não faz mal até o conjunto retido cair numa delas, e sete bigramas de 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. A correção habitual é somar 1 a todas as contagens antes de normalizar, o que aqui quase não custa nada (2,4546 em vez de 2,4524). Mas a correção é uma confissão. Um modelo de contagem não consegue generalizar de todo. Não tem forma de suspeitar que qo é plausível porque qu é comum e o se comporta como u noutros lugares, pois não tem noção de que dois símbolos podem parecer-se. Cada célula é aprendida isoladamente, e corrigir isso é o objetivo do resto deste capítulo.

A perplexidade é um preço por token, e o token é um parâmetro livre. Este é o erro que aparece constantemente quando se comparam modelos, e é fácil de ver quando se olha. Pegue no mesmo corpus de prosa inglesa do Capítulo 7, no mesmo modelo bigrama interpolado, e mude apenas a forma como o texto é dividido:

unidadevocabuláriotokens no testecross-entropyperplexidadebits por carácter
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 no modelo mudou; apenas o tamanho da coisa a prever. Prever uma palavra inteira é mais difícil do que prever uma letra, por isso 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 e o converte para bits. Ela reordena a tabela. Pela perplexidade, a ordenação é caracteres, palavras, BPE-512, BPE-2048; por bits por carácter, é palavras, BPE-2048, BPE-512, caracteres. O modelo de caracteres passa do primeiro lugar para o último. O modelo com 2.048 merges, que pela perplexidade parece 6,6 vezes pior do 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 partilham um tokenizer, e modelos com tokenizers diferentes só podem ser comparados em bits por carácter — a quantidade que Shannon mediu em 1951 ao pedir a sujeitos humanos que adivinhassem a próxima letra de texto inglês, e que limitou a cerca de um bit por carácter.2 O nosso melhor bigrama fica em 2,23 bits, o que resume bem o quanto este capítulo ainda tem de avançar.

Agora construa o mesmo modelo como uma rede. Vai exigir ordens de grandeza mais aritmética para chegar ao mesmo sítio, e chegar ao mesmo sítio é o ponto.

Substitua a tabela por uma matriz de pesos WW com forma 27×2727 \times 27. Transforme a letra atual num vetor one-hot, multiplique, e chame ao resultado logits — as pontuações não normalizadas do Capítulo 4. Depois softmax, depois cross-entropy, 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 a pena guardar. Multiplicar um vetor one-hot por uma matriz seleciona uma linha dessa matriz, portanto a multiplicação é uma consulta — e todas as implementações saltam a aritmética e fazem a consulta diretamente, que é o que W[xs] é.

Isto é uma tabela de embedding. Uma matriz com uma linha por entrada do vocabulário, indexada por id de token. Sem geometria, sem semântica, sem algoritmo separado: uma lookup table cujo conteúdo acontece ser aprendido por gradient descent juntamente com tudo o resto. Todas as afirmações místicas sobre «embedding space» acabam aqui.

Treine-a e observe para onde 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 vai para zero. Após 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 treino; 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, a partir de números aleatórios e sem receber nada além de «torna grande a log-probabilidade da próxima letra», redescobriu a tabela de contagens. E tinha de o fazer: as contagens são a estimativa de máxima verosimilhança, a cross-entropy é a log-verosimilhança negativa, portanto os dois procedimentos otimizam o mesmo objetivo e esse objetivo tem um único ótimo. A rede não aprendeu algo parecido com contar. Convergiu para contar, lentamente.

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

Estenda o modelo para olhar para mais do que um carácter anterior. Esta é a arquitetura de Bengio de 2003, a antepassada direta de todos os modelos no resto deste curso:4 pegue nos últimos três caracteres, mapeie cada um através de 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 que produz 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)

Repare no que é novo e no que não é. A camada oculta é a do Capítulo 5, sem alterações; a loss é a do Capítulo 4, sem alterações. As novidades são a tabela de embedding à entrada e uma camada de saída tão larga como o vocabulário do Capítulo 7 — e essa segunda é a parte cara de todos os modelos de linguagem alguma vez construídos, porque um vocabulário real tem 100.000 entradas e esta multiplicação matricial corre em todas as posições.

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 carácter7292,454611,642
neuronal, 1 carácter7.8972,457711,678
neuronal, 3 caracteres11.8972,11458,285
neuronal, 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 do que a tabela de contagens tem um desempenho exatamente tão bom como a tabela de contagens e não melhor. A capacidade nunca foi a limitação. Um carácter de contexto permite uma certa loss e nada que lhe acrescente consegue descer abaixo disso, porque a informação não está lá.

Dê-lhe três caracteres e a perplexidade desce de 11,68 para 8,29 — uma redução de 29 %, comprada com mais 4.000 parâmetros. Aqui vence a contagem exatamente pela razão diagnosticada antes: um modelo de contagem sobre contextos de três caracteres precisa de 273=19,68327^3 = 19{,}683 linhas, a maioria vazias ou com uma única observação, e aprende cada uma isoladamente. A rede partilha. Se a, e e i acabarem com linhas de embedding semelhantes, o que ela aprende depois de bra transfere-se para bre sem nunca ter 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 em conformidade:

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 destoariam numa, e os monstros intermináveis desapareceram: o mais longo de vinte exemplos do modelo de contagem tem dezanove letras; o mais longo de vinte deste tem treze.

A tabela é 27×1027 \times 10: uma linha de dez números por carácter, todos inicializados aleatoriamente e movidos apenas pelo gradient da loss de próximo carácter. Ninguém pôs lá nada. Então o que acabou dentro dela?

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

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

Mede o ângulo entre dois vetores e ignora os seus comprimentos, que é o que se quer quando o comprimento de uma linha reflete a frequência com que o seu token apareceu, e não o que ele significa. Normalize primeiro todos os vetores para comprimento 1 — como os 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, tal como i e y; j e z são ambas consoantes raras, sobretudo iniciais, que se comportam de forma semelhante. O símbolo de fronteira . não está perto de praticamente nada — 0,077 da letra mais próxima — porque é o único símbolo que marca uma posição em vez de um som.

E parte disso 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 as consoantes, e o efeito é real mas pequeno. Testado contra 2.000 grupos aleatórios de cinco letras, 58 desses grupos separam-se pelo menos tão claramente — uma diferença significativa a cerca de p=0.03p = 0.03. Real, portanto, mas nada parecido com a ilha geométrica nítida que os relatos populares sobre embeddings sugerem.

Esta é a descrição honesta de uma tabela de embedding, e vale a pena mantê-la presente no resto do curso. Não é um mapa de significado. É uma mudança de coordenadas, aprendida em vez de desenhada, cuja única função é facilitar o trabalho da camada seguinte — a mesma frase que o Capítulo 5 usou para a camada oculta que dobrou o plano para resolver XOR. Qualquer estrutura que encontre nela está lá porque reduziu a loss, e a estrutura que não reduz a loss simplesmente não está lá.

Se a parte útil é a tabela, pode atacá-la diretamente. Isso é word2vec: mantenha a consulta de embedding, deite fora o modelo de linguagem.5

O objetivo skip-gram with negative sampling cabe numa linha. Para um par real (centro, contexto) retirado do corpus, empurre o seu produto escalar para cima; para kk pares falsos retirados 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})

Isto é uma classificação binária — «estas duas palavras ocorreram realmente juntas?» — e é barato precisamente porque nunca toca no vocabulário completo, que foi o que tornou prático treinar em milhares de milhões de palavras em 2013. GloVe chega a vetores semelhantes pelo outro lado, fatorizando a matriz de contagens globais de coocorrência em vez de percorrer exemplos em stream.7 Ambos são ajustados exatamente à estatística a partir da qual a tabela de contagens foi construída. São contagem, comprimida.

Treinados em text8 — 17.005.207 palavras da Wikipédia inglesa, 71.290 delas a ocorrer pelo menos cinco vezes, 100 dimensões, três passagens — 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 para numerais. Agora a parte famosa: pegue em 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. Isto não é uma peculiaridade de um exemplo. O conjunto de avaliação de Mikolov coloca 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 a 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 o mencionam, porque a regra de pontuação padrão elimina a, b e c antes de procurar. É uma regra legítima, e está a fazer mais trabalho do que a aritmética:

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

A terceira linha é a que merece pausa. Deite fora a e b, não faça aritmética nenhuma, devolva o que estiver mais próximo de c — e 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 devidamente treinados e o que as baselines acima replicam.8 Estes vetores em particular são pequenos — 17 milhões de palavras contra os milhares de milhões por trás dos modelos publicados — por isso leia as percentagens como uma forma, não como o estado da arte. A forma é o que sobrevive em qualquer escala: a aritmética é real, e muito mais fraca do que a única demonstração que toda a gente cita.

Estáticos e contextuais: um vetor por palavra, ou um por ocorrência

Ligação para a secção: Estáticos e contextuais: 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 numa frase sobre um rio e numa frase sobre uma hipoteca — necessariamente, porque uma consulta por id não pode depender de mais nada.

A correção é deixar de ler o vetor a partir da tabela e começar a calculá-lo a partir da frase. Isso é um contextual embedding, 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 nítidos do 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, por isso o cosseno é 1 por construção. Nove camadas depois, as duas ocorrências ficam em 0,43, enquanto bank em duas frases diferentes sobre rios fica em 0,87. Ninguém rotulou um sentido em lado nenhum neste processo; os sentidos separaram-se porque separá-los torna o objetivo de treino — adivinhar um token oculto a partir dos seus vizinhos — mais fácil de satisfazer.

Dois detalhes recompensam atenção. A camada 0 já é 0,9512 em vez de 1,0, porque foram adicionados position embeddings e a palavra está num lugar diferente em cada frase. E a similaridade volta a subir nas camadas 11 e 12: as camadas finais de um modelo pré-treinado são especializadas no seu objetivo de treino, e muitas vezes não são o melhor lugar de onde retirar 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. Num 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 — faz um produto escalar contra cada uma — e o tying diz que o vetor usado para ler um token e o vetor usado para o escrever devem ser o mesmo objeto. Corta parâmetros e melhora a perplexidade ao mesmo tempo, o que é raro o suficiente para merecer nota.

Para pesquisar um corpus por significado, precisa de um vetor por frase. Dados esses vetores, a pesquisa é trivial — isto é todo o retrieval semântico, e o Capítulo 19 é sobre tudo o que está à volta dele:

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]

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

como a frase é embeddedcorrelação de postostop-1 num índice de 200 frases
sobreposição binária de palavras (sem modelo nenhum)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 com 109 milhões de parâmetros, usado da forma óbvia, é pior a julgar similaridade entre frases do que contar quantas palavras duas frases partilham — e pior do que fazer a média dos vetores text8 de 100 dimensões treinados há momentos. O token [CLS], que os tutoriais ainda recomendam porque o BERT foi pré-treinado com um objetivo ao nível da frase ligado a ele, é pior do que metade disso.

Isto não é um defeito do BERT. É o objetivo. Um modelo de linguagem é treinado para que os seus hidden states prevejam um token; nada aí pede que duas paráfrases acabem perto uma da outra, e nada recompensa uma geometria em que o cosseno significa «mesmo significado». A última linha é um modelo com um quinto do tamanho (22.713.216 parâmetros) treinado com uma loss completamente diferente: aprendizagem contrastiva, em que os exemplos são pares — uma pergunta e a sua resposta, uma frase e a sua paráfrase — e o objetivo aproxima os pares verdadeiros enquanto afasta negativos amostrados. 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 à pesquisa, com um encoder para queries e outro para passagens.13

Portanto, a regra prática:

Um embedding model não é um modelo de linguagem com a última camada removida. É um modelo diferente, com um objetivo diferente, normalmente muito mais pequeno, cujo cosseno significa o que quer que signifique porque foi treinado em pares onde esse 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 sacos de palavras idênticos, por isso a sobreposição de palavras e a média de vetores estáticos dão-lhes cosseno exatamente 1,000000, e BERT com mean pooling, que vê posição, ainda fica quase aí — e o MiniLM treinado contrastivamente ainda os coloca em 0,979. Se a sua tarefa de retrieval depende de quem fez o quê a quem, nenhum limiar de cosseno o vai salvar.

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

Embeddings reais têm centenas ou milhares de componentes, e as distâncias comportam-se de forma estranha lá em cima. Pegue em 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 do que o par mais próximo. Tudo fica aproximadamente equidistante de tudo o resto, «nearest neighbour» deixa de transportar muita informação, e essa é a maldição da dimensionalidade — bem como uma das razões pelas quais grandes bases de dados vetoriais não fazem pesquisa exata de nearest neighbour. O outro lado da mesma moeda é o que torna viáveis os limiares de cosseno: medido sobre 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 não é, portanto, «40 % parecido»; está muito fora de qualquer coisa que o acaso produza, e é por isso que 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 deles e cola os resultados por ordem. Esse desenho tem dois problemas, e são o mesmo problema.

Olhe novamente para a tabela de contexto: passar de três caracteres para oito quase duplicou os parâmetros e comprou 0,06 nats. O custo cresce linearmente com o contexto — cada posição extra precisa do seu próprio bloco da primeira matriz de pesos — e o benefício não. Empurre isto para mil tokens e a primeira camada sozinha pesa mais do que o resto 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 forma de decidir quais dos tokens anteriores importam. A posição dois recebe os seus próprios pesos e a posição sete recebe os seus, permanentemente, seja o que for que esteja nelas. Quando o modelo está a soletrar nell, o carácter decisivo é o imediatamente anterior. Quando uma frase contém um pronome, a palavra que fixa o seu referente pode estar quarenta tokens atrás — e nenhuma slot fixa pode ser atribuída a «quarenta atrás», porque da próxima vez serão seis.

O que queremos é um modelo que calcule, para cada previsão, quanto deve contar cada token anterior — pesos sobre o contexto produzidos pelo conteúdo em vez de fixados pelo layout. Escreva isso com cuidado e começa como algo inteiramente banal: uma média sobre os tokens anteriores. Depois deixe que os pesos dessa média sejam aprendidos, e deixe que dependam do token que está a fazer a pergunta.

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


Também vale a pena ler em paralelo: o capítulo 3 de Speech and Language Processing, de Jurafsky e Martin, que trata modelos n-gram, smoothing e perplexidade com muito mais cuidado do que há espaço para aqui, incluindo porque a interpolação e o back-off vencem somar um; as notas de Stanford CS229 §17.1–17.2 sobre modelação de linguagem pelo lado probabilístico; e o artigo de Linzen acima, que é curto e vale a pena ler na íntegra.

  1. O exemplo de geração de nomes, o dataset e a progressão de uma tabela de contagens para uma rede ao 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). Sujeitos humanos a adivinhar a próxima letra de texto inglês, e a medição original de bits por carácter.

  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 da previsão com a compressão.

  4. Bengio, Y., Ducharme, R., Vincent, P. e 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, através de uma camada oculta, até uma softmax sobre o vocabulário.

  5. Mikolov, T., Chen, K., Corrado, G. e 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. e 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. e Manning, C. GloVe: Global Vectors for Word Representation. EMNLP 2014. Vetores de palavras a partir de uma fatorização da matriz global de coocorrência, em vez de janelas locais em stream.

  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, calculado por um modelo de linguagem bidirecional.

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

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

  12. Reimers, N. e Gurevych, I. Sentence-BERT: Sentence Embeddings using Siamese BERT-Networks. arXiv:1908.10084 (2019). A sua medição inicial — BERT com mean pooling 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). Treino contrastivo de um retriever de dois encoders; o antepassado direto da stack de retrieval do Capítulo 19.

Pronto para deixar a LIA escolher?

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