Zum Inhalt springen
8/30Kapitel 8 von 30

Next-token prediction: Embeddings und was Perplexity bedeutet

Ein Zeichenmodell auf 32.033 Namen: gradient descent entdeckt eine Zähltabelle neu – und warum Perplexities selten vergleichbar sind.

Auf dieser Seite

Hier sind zehn Namen, erzeugt von einem Programm, das noch nie ein Wort gesehen hat:

TEXT
cexze   momakurailezitynn   konimittain   llayn   ka
da      moliellavo          emia          sade    ftlsp

Keiner davon ist ein Name. Fast alle versuchen es. Sie sind aussprechbar, sie enden dort, wo Namen enden, und einer von ihnen — emia — ist nur einen Buchstaben von einem echten entfernt. Das Programm, das sie erzeugt hat, enthält 729 Zahlen, hat keine Vorstellung von einem Wort, einer Silbe oder einer Person und wurde durch einen einzigen Durchlauf trainiert, der benachbarte Buchstabenpaare gezählt hat.

Am Ende dieses Kapitels wird ein neuronales Netz den Score dieses Programms auf derselben Messgröße um ein Drittel gesenkt haben. Der Teil, für den es sich lohnt dranzubleiben, ist das, was das Netz zuerst tut: Es reproduziert die Zähltabelle auf jeder gut gefüllten Zeile bis auf drei Dezimalstellen, ohne prompt, weil beide Objekte Antworten auf dieselbe Frage sind. Alles danach ist das, was Zählen niemals hätte leisten können.

Das Ziel ist eine Identität, keine Designentscheidung

Link zum Abschnitt: Das Ziel ist eine Identität, keine Designentscheidung

Kapitel 7 ließ dich mit einer Folge von Ganzzahlen zurück und ohne Grund, warum eine auf die andere folgen sollte. Hier ist der Grund, und er ist eine Zeile aus Kapitel 2.

Ein Sprachmodell ist eine Funktion, die die bisherigen tokens nimmt und eine Verteilung darüber zurückgibt, welches token als Nächstes kommt: eine Zahl pro Vokabulareintrag, nicht negativ, zusammen eins. Sonst nichts. Um daraus eine Wahrscheinlichkeit für ein ganzes Dokument zu bekommen, wende die Kettenregel der Wahrscheinlichkeit an:

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})

Das ist eine Identität, wahr für jede Folge von irgendetwas, ohne zusätzliche Annahmen. Ein Modell, das die kleine Aufgabe erledigt — nächstes token gegeben die vorherigen — hat also bereits die große Aufgabe erledigt, jedem möglichen Dokument eine Wahrscheinlichkeit zuzuweisen, exakt und gratis. Die populäre Darstellung als billiger Trick („es sagt nur das nächste Wort voraus“) dreht die Logik um: Das nächste token vorherzusagen ist die Modellierung der gemeinsamen Verteilung. Es gab nie eine zweite Sache zu tun.

Der Loss folgt genauso mechanisch. An jeder Position erzeugt das Modell eine Verteilung qq und die Wahrheit ist ein einzelnes bekanntes token, also gilt die Kreuzentropie aus Kapitel 4 unverändert:

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

Das ist die durchschnittliche negative Log-Likelihood — das Rezept aus Kapitel 2 mit einer kategorialen Verteilung an der Stelle, an der vorher die Gauß-Verteilung saß. Und weil die wahre Verteilung one-hot ist, ist ihre Entropie null; nach der Identität aus Kapitel 4 ist die Kreuzentropie also gleich der KL-Divergenz: Diese Zahl zu senken und die Überzeugungen des Modells näher an die Daten zu ziehen, ist derselbe Akt.

Eine Konsequenz verdient einen eigenen Satz, weil sie die ökonomische Tatsache unter dem ganzen Feld ist. Die Labels sind die Daten, um eine Position verschoben. Niemand annotiert irgendetwas. Eine Billion tokens Text sind eine Billion vorlabelte Beispiele, weshalb der Trainingskorpus eines modernen Modells „das Internet“ ist und nicht „ein Datensatz, den jemand gebaut hat“.

Vor jedem Netz die Baseline: 32.033 Namen, einer pro Zeile, und die Aufgabe, mehr davon zu erzeugen, Buchstabe für Buchstabe.1

Das Vokabular besteht aus 26 Buchstaben plus einem Grenzsymbol ., das sowohl den Anfang als auch das Ende eines Namens markiert; das Modell muss also lernen, wo Namen beginnen und wo sie aufhören. Das sind 27 Symbole, und das kleinstmögliche Modell ist eine Tabelle, die zählt, wie oft jedes Symbol auf jedes andere Symbol folgte.

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   

Zwei Zeilen Arithmetik und das Modell ist angepasst — und es ist keine Heuristik: Counts durch Zeilensummen zu teilen ist der Maximum-Likelihood-Schätzer für eine kategoriale Verteilung, also das Rezept aus Kapitel 2 mit bereits erledigter Analysis.

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

Ziehe Samples daraus — wähle einen Buchstaben aus der Zeile des aktuellen Buchstabens, gehe zu dieser Zeile, wiederhole, bis das Grenzsymbol erscheint — und du bekommst die Namen vom Anfang dieses Kapitels. Sie scheitern auf eine spezifische und aufschlussreiche Weise: lokal plausibel, global Unsinn. Jedes benachbarte Buchstabenpaar in momakurailezitynn kommt in echten Namen vor; es sind nur eben siebzehn davon hintereinander. Das Modell hat ein Zeichen Gedächtnis, also kann es nicht wissen, dass es schon zu lange weiterläuft.

Der Loss auf zurückgehaltenen Namen beträgt 2,4546 nats. Diese Zahl bedeutet für sich genommen nichts, und genau deshalb gibt es Perplexity:

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}

Ausgeschrieben, ohne dass eine Library die Arbeit macht:

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())             

Exponentieren macht den Logarithmus rückgängig und bringt die Zahl zurück in Einheiten des Zählens von Dingen. Am saubersten sieht man, was sie zählt, wenn man ein Modell misst, das gar nichts weiß — eines, das jedem Symbol unabhängig vom Kontext die Wahrscheinlichkeit 1/271/27 zuweist:

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

Exakt 27,000, weil elog27=27e^{\log 27} = 27. Perplexity ist die effektive Zahl gleich wahrscheinlicher Optionen, zwischen denen das Modell wählt. Eine Perplexity von 27 bedeutet: „Keine Ahnung, könnte alles sein.“ Die 11,642 des Zählmodells bedeuten, dass ein Buchstabe Kontext es so unsicher lässt wie jemanden, der blind aus ungefähr zwölf statt siebenundzwanzig Optionen wählt — deshalb wird Perplexity zitiert und der rohe Loss nicht.

Zwei Dinge laufen dabei schief, und das zweite läuft in veröffentlichten Papers schief.

Nullwahrscheinlichkeiten sind tödlich. Von den 729 Zellen der Tabelle kommen 113 im Training nie vor — 15,5 % davon sind leer. Das ist in Ordnung, bis das zurückgehaltene Set in einer landet, und sieben Bigramme in der Validierung tun das, darunter dq, zj und qo zweimal. Wahrscheinlichkeit null bedeutet log -\infty, also unendlichen Loss und unendliche Perplexity: Ein Name in dreitausend zerstört die Metrik. Der übliche Patch ist, vor der Normalisierung zu jedem Count 1 hinzuzufügen, was hier fast nichts kostet (2,4546 statt 2,4524). Aber der Patch ist ein Geständnis. Ein Zählmodell kann überhaupt nicht generalisieren. Es hat keine Möglichkeit zu vermuten, dass qo plausibel ist, weil qu häufig ist und o sich anderswo wie u verhält, denn es hat keinen Begriff davon, dass zwei Symbole einander ähneln können. Jede Zelle wird allein gelernt, und genau das zu beheben ist der Zweck des restlichen Kapitels.

Perplexity ist ein Preis pro token, und das token ist ein freier Parameter. Das ist der Fehler, der ständig auftaucht, wenn Modelle verglichen werden, und er ist leicht zu sehen, sobald man hinschaut. Nimm denselben Korpus englischer Prosa aus Kapitel 7, dasselbe interpolierte Bigramm-Modell, und ändere nur, wie der Text zerlegt wird:

EinheitVokabulartokens im TestKreuzentropiePerplexityBits pro Zeichen
Zeichen7614.4692,521712,453,6378
BPE, 512 Merges3296.8713,854747,212,6407
BPE, 2.048 Merges1.8204.2335,7468313,202,4254
Wörter2.9916.2843,562735,262,2322

Perplexity schwankt zwischen diesen Zeilen um den Faktor 25. Am Modell hat sich nichts geändert; nur die Größe des vorhergesagten Dings. Ein ganzes Wort vorherzusagen ist schwieriger als einen Buchstaben vorherzusagen, kostet also mehr pro Vorhersage — und es gibt weniger Vorhersagen zu machen.

Lies jetzt die letzte Spalte: Sie teilt die Gesamtkosten stattdessen durch die Zahl der Zeichen und wandelt sie in Bits um. Sie sortiert die Tabelle neu. Nach Perplexity lautet das Ranking Zeichen, Wörter, BPE-512, BPE-2048; nach Bits pro Zeichen lautet es Wörter, BPE-2048, BPE-512, Zeichen. Das Zeichenmodell fällt vom ersten auf den letzten Platz. Das 2.048-Merge-Modell, das nach Perplexity 6,6-mal schlechter aussieht als das 512-Merge-Modell, ist tatsächlich das bessere der beiden: 2,4254 Bits gegenüber 2,6407.

Eine Perplexity ist also nur zwischen zwei Modellen vergleichbar, die denselben tokenizer teilen, und Modelle mit unterschiedlichen tokenizers lassen sich nur in Bits pro Zeichen vergleichen — der Größe, die Shannon 1951 gemessen hat, indem er Versuchspersonen den nächsten Buchstaben englischen Textes raten ließ, und die er auf ungefähr ein Bit pro Zeichen eingrenzte.2 Unser bestes Bigramm liegt bei 2,23 Bits, eine faire Zusammenfassung dafür, wie weit dieses Kapitel noch gehen muss.

Baue nun dasselbe Modell als Netz. Es wird um Größenordnungen mehr Arithmetik brauchen, um am selben Ort anzukommen — und am selben Ort anzukommen ist der Punkt.

Ersetze die Tabelle durch eine Gewichtsmatrix WW der Form 27×2727 \times 27. Verwandle den aktuellen Buchstaben in einen One-hot-Vektor, multipliziere, und nenne das Ergebnis logits — die unnormalisierten Scores aus Kapitel 4. Dann softmax, dann Kreuzentropie, dann 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

Die hervorgehobene Zeile enthält eine Definition, die man haben sollte. Eine One-hot-Vektor-Matrix-Multiplikation wählt eine Zeile der Matrix aus, die Multiplikation ist also ein Lookup — und jede Implementierung überspringt die Arithmetik und macht den Lookup direkt, genau das ist W[xs].

Das ist eine embedding-Tabelle. Eine Matrix mit einer Zeile pro Vokabulareintrag, indiziert durch token id. Keine Geometrie, keine Semantik, kein separater Algorithmus: eine Lookup-Tabelle, deren Inhalte zufällig zusammen mit allem anderen durch gradient descent gelernt werden. Jede mystische Behauptung über „embedding space“ endet hier.

Trainiere sie und schau, wohin sie läuft:

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

Die letzte Spalte ist die größte absolute Differenz zwischen irgendeiner Zelle von softmax(W) und der passenden Zelle der Zähltabelle, und sie geht gegen null. Nach 3.000 Schritten beträgt die größte Abweichung irgendwo in den 729 Zellen 0,004048, der Mittelwert 0,000224. Die schlechteste Zelle ist qi, im gesamten Trainingsset zwölfmal gesehen; unter den 22 Zeilen mit mehr als tausend Vorkommen beträgt die schlechteste Abweichung 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, gestartet mit Zufallszahlen und mit nichts anderem instruiert als „mach die Log-Wahrscheinlichkeit des nächsten Buchstabens groß“, hat die Zähltabelle wiederentdeckt. Und das musste es: Die Counts sind der Maximum-Likelihood-Schätzer, Kreuzentropie ist die negative Log-Likelihood, beide Verfahren optimieren also dasselbe Ziel, und dieses Ziel hat ein Optimum. Das Netz hat nicht etwas Ähnliches wie Zählen gelernt. Es konvergierte zu Zählen, langsam.

Das wirft die berechtigte Frage auf, warum sich irgendjemand damit beschäftigen sollte. Weil die Zähltabelle von hier aus nirgendwohin kann — das Netz aber schon.

Erweitere das Modell so, dass es mehr als ein vorheriges Zeichen anschaut. Das ist Bengios Architektur von 2003, der direkte Vorfahr jedes Modells im Rest dieses Kurses:4 Nimm die letzten drei Zeichen, mappe jedes davon durch eine embedding-Tabelle in eine 10-dimensionale Zeile, konkateniere die Zeilen zu 30 Zahlen, schiebe sie durch die Hidden Layer aus Kapitel 5, und beende mit einer Ausgabeschicht, die ein logit pro Vokabulareintrag erzeugt.

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)

Beachte, was neu ist und was nicht. Die Hidden Layer ist die aus Kapitel 5, unverändert; der Loss ist der aus Kapitel 4, unverändert. Neu sind die embedding-Tabelle am Anfang und eine Ausgabeschicht, die so breit ist wie das Vokabular aus Kapitel 7 — und diese zweite ist der teure Teil jedes jemals gebauten Sprachmodells, weil ein echtes Vokabular 100.000 Einträge hat und diese Matrixmultiplikation an jeder Position läuft.

Derselbe Code, identisch trainiert, nur mit geänderter Größe des context window:

contextParameterValidierungs-LossValidierungs-Perplexity
Zählen, 1 Zeichen7292,454611,642
neural, 1 Zeichen7.8972,457711,678
neural, 3 Zeichen11.8972,11458,285
neural, 8 Zeichen21.8972,05067,773

Die zweite Zeile ist die interessante. Ein Netz mit einer 200-Unit-Hidden-Layer und elfmal so vielen Parametern wie die Zähltabelle leistet genau so viel wie die Zähltabelle und nicht mehr. Kapazität war nie die Begrenzung. Ein Zeichen context erlaubt einen bestimmten Loss, und nichts, was du daraufsetzt, kann daruntergehen, weil die Information nicht da ist.

Gib ihm drei Zeichen und die Perplexity fällt von 11,68 auf 8,29 — eine Senkung um 29 %, gekauft mit 4.000 zusätzlichen Parametern. Es schlägt das Zählen hier aus genau dem zuvor diagnostizierten Grund: Ein Zählmodell über Drei-Zeichen-Kontexte braucht 273=19,68327^3 = 19{,}683 Zeilen, die meisten davon leer oder mit einer einzigen Beobachtung, und es lernt jede allein. Das Netz teilt. Wenn a, e und i ähnliche embedding-Zeilen bekommen, überträgt sich, was es nach bra lernt, auf bre, ohne dass es bre jemals gesehen hat. Dieser Transfer ist der ganze Wert der embedding-Tabelle, und er ist die Lücke zwischen Zeile zwei und drei.

Die Samples werden entsprechend besser:

TEXT
deliah   nellara   joce     kael      quintis
salayson  reety    khyrmin  mahnen    madiaryxia

Immer noch keine Liste echter Namen. Aber deliah, nellara und kael würden auf einer solchen nicht fehl am Platz wirken, und die ausufernden Monster sind verschwunden: Das längste von zwanzig Samples des Zählmodells hat neunzehn Buchstaben, das längste von zwanzig dieses Modells dreizehn.

Was tatsächlich in der embedding-Tabelle steckt

Link zum Abschnitt: Was tatsächlich in der embedding-Tabelle steckt

Die Tabelle ist 27×1027 \times 10: eine Zeile mit zehn Zahlen pro Zeichen, alle zufällig initialisiert und nur durch den Gradienten des Next-Character-Loss bewegt. Niemand hat dort etwas hineingelegt. Was ist also darin gelandet?

Das Werkzeug für diese Frage ist Kosinusähnlichkeit, also das Skalarprodukt aus Kapitel 1, nachdem die Längen herausgeteilt wurden:

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

Sie misst den Winkel zwischen zwei Vektoren und ignoriert ihre Längen — genau das, was du willst, wenn die Länge einer Zeile eher widerspiegelt, wie oft ihr token vorkam, als was es bedeutet. Normalisiere jeden Vektor zuerst auf Länge 1 — wie reale Systeme es einmalig zur Indexierungszeit tun — und Kosinusähnlichkeit ist einfach das Skalarprodukt.

Hier sind die nächsten Nachbarn einiger Zeichen in der trainierten Tabelle:

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

Ein Teil davon ist das, was die Folklore verspricht. c und k sind in Namen austauschbar, ebenso i und y; j und z sind beide seltene, meist initiale Konsonanten, die sich ähnlich verhalten. Das Grenzsymbol . ist nahe bei gar nichts — 0,077 zu seinem nächsten Buchstaben — weil es das einzige Symbol ist, das eine Position statt eines Lauts markiert.

Und ein Teil davon nicht. Der nächste Nachbar von a ist h, kein anderer Vokal. Über alle Paare gemittelt:

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

Die Vokale ähneln einander stärker als Konsonanten, und der Effekt ist real, aber klein. Gegen 2.000 zufällig gewählte Gruppen von fünf Buchstaben getestet, trennen sich 58 dieser Gruppen mindestens genauso sauber — eine Lücke, signifikant bei ungefähr p=0.03p = 0.03. Also real, aber nichts wie die scharf abgegrenzte geometrische Insel, die populäre Darstellungen von embeddings nahelegen.

Das ist die ehrliche Beschreibung einer embedding-Tabelle, und es lohnt sich, sie für den Rest des Kurses festzuhalten. Sie ist keine Karte von Bedeutung. Sie ist ein Koordinatenwechsel, gelernt statt entworfen, dessen einzige Aufgabe es ist, die Aufgabe der nächsten Schicht leicht zu machen — derselbe Satz, den Kapitel 5 für die Hidden Layer nutzte, die die Ebene faltete, um XOR zu lösen. Jede Struktur, die du darin findest, ist dort, weil sie den Loss senkte; und Struktur, die den Loss nicht senkt, ist schlicht nicht dort.

word2vec, GloVe und die Arithmetik, die alle zitieren

Link zum Abschnitt: word2vec, GloVe und die Arithmetik, die alle zitieren

Wenn der nützliche Teil die Tabelle ist, kannst du direkt auf sie zielen. Das ist word2vec: Behalte den embedding-Lookup, wirf das Sprachmodell weg.5

Das Ziel von skip-gram with negative sampling ist eine Zeile. Für ein echtes (Zentrum, Kontext)-Paar aus dem Korpus schiebe ihr Skalarprodukt nach oben; für kk Fake-Paare aus einer Rauschverteilung schiebe es nach unten: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})

Das ist binäre Klassifikation — „Sind diese zwei Wörter wirklich zusammen vorgekommen?“ — und sie ist genau deshalb billig, weil sie nie das vollständige Vokabular berührt, was 2013 Training auf Milliarden Wörtern praktikabel machte. GloVe kommt von der anderen Seite zu ähnlichen Vektoren, indem es die Matrix globaler Kookkurrenz-Counts faktorisiert, statt durch Beispiele zu streamen.7 Beide werden exakt auf die Statistik angepasst, aus der die Zähltabelle gebaut wurde. Sie sind Zählen, komprimiert.

Trainiert auf text8 — 17.005.207 Wörter aus der englischen Wikipedia, 71.290 davon mindestens fünfmal vorkommend, 100 Dimensionen, drei Durchläufe — kommen die Vektoren mit der Eigenschaft heraus, die sie berühmt gemacht hat:

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

Niemand hat eine Kategorie für Instrumente oder Numeralia geliefert. Jetzt der berühmte Teil: Nimm king, subtrahiere man, addiere woman und finde den nächsten Vektor zum Ergebnis.

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)

Der nächste Vektor zu king - man + woman ist king. Das ist keine Eigenheit eines Beispiels. Mikolovs Evaluationsset stellt Fragen der Form a : b :: c : ? — 8.869 semantische (paris : france :: rome : italy) und 10.675 syntaktische (walking : walked :: swimming : swam) — und über die 4.103 semantischen Fragen hinweg, die dieses Vokabular beantworten kann, ist der Gewinner in 99,8 % der Fälle eines der drei Eingabewörter. Die veröffentlichten Demonstrationen erwähnen das nicht, weil die Standard-Scoring-Regel a, b und c vor der Suche löscht. Das ist eine legitime Regel, und sie leistet mehr als die Arithmetik:

wie die Antwort gewählt wirdsemantischsyntaktisch
Offset, mit ausgeschlossenen Eingaben (Standard)17,0 %11,9 %
Offset, ohne Ausschlüsse0,1 %0,4 %
nächster Nachbar von c allein, Eingaben ausgeschlossen13,1 %9,3 %
nächster Nachbar von b allein, Eingaben ausgeschlossen2,3 %0,4 %

Die dritte Zeile ist die, bei der man verweilen sollte. Wirf a und b weg, mache überhaupt keine Arithmetik, gib zurück, was c am nächsten ist — und du behältst 77 % des semantischen Scores. Das meiste, was wie analogisches Schließen aussieht, ist Nähe plus eine Regel, die die offensichtlichen Antworten verbietet; genau das hat Linzen an korrekt trainierten Vektoren gemessen, und genau das replizieren die Baselines oben.8 Diese konkreten Vektoren sind klein — 17 Millionen Wörter gegenüber den Milliarden hinter den veröffentlichten Modellen — lies die Prozentsätze also als Form, nicht als State of the Art. Die Form überlebt auf jeder Skala: Die Arithmetik ist real und viel schwächer als die eine Demonstration, die alle zitieren.

Statisch und kontextuell: ein Vektor pro Wort oder einer pro Vorkommen

Link zum Abschnitt: Statisch und kontextuell: ein Vektor pro Wort oder einer pro Vorkommen

Alles bisher hat eine harte Grenze in die Datenstruktur eingebaut. Eine Tabelle hat eine Zeile pro token. Das Wort bank bekommt einen Vektor, denselben in einem Satz über einen Fluss und in einem Satz über eine Hypothek — notwendigerweise, weil ein Lookup per id von nichts anderem abhängen kann.

Die Lösung ist, den Vektor nicht mehr aus der Tabelle zu lesen, sondern ihn aus dem Satz zu berechnen. Das ist ein contextual embedding, eingeführt 2018 durch ELMo und im selben Jahr durch BERT zum Standard gemacht.910 Am echten Modell gemessen sind die Zahlen schärfer als die Erklärung:

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

Die erste Zeile ist exakt, nicht approximativ: Der statische Vektor für bank besteht in beiden Sätzen aus denselben 768 Zahlen, der Kosinus ist also per Konstruktion 1. Neun Schichten später liegen die zwei Vorkommen bei 0,43, während bank in zwei verschiedenen Fluss-Sätzen bei 0,87 bleibt. Niemand hat in diesem Prozess irgendwo eine Bedeutung gelabelt; die Bedeutungen trennten sich, weil ihre Trennung das Trainingsziel — ein verborgenes token aus seinen Nachbarn zu erraten — leichter erfüllbar macht.

Zwei Details lohnen Aufmerksamkeit. Layer 0 liegt bereits bei 0,9512 statt 1,0, weil position embeddings addiert wurden und das Wort in jedem Satz an einer anderen Stelle steht. Und die Ähnlichkeit steigt wieder in den Schichten 11 und 12: Die letzten Schichten eines vortrainierten Modells sind auf sein Trainingsziel spezialisiert und oft nicht der beste Ort, um eine Repräsentation zu entnehmen.

Details anzeigen

Optional: Weight tying.

In bert-base-uncased ist die embedding-Tabelle 30,522×76830{,}522 \times 768 — 23.440.896 Zahlen, 21,4 % der 109.482.240 Parameter des Modells. In einem kleinen Sprachmodell ist der Anteil noch größer, weshalb ein Trick fast universell ist: Die Eingabetabelle und die Ausgabeschicht, die die logits erzeugt, sind dieselbe Matrix, einmal per Zeilen-Lookup und einmal transponiert verwendet.11 Die Ausgabeschicht weist bereits jedem Vokabulareintrag einen Vektor zu — sie bildet ein Skalarprodukt gegen jeden davon — und Tying sagt, dass der Vektor, mit dem ein token gelesen wird, und der Vektor, mit dem es geschrieben wird, dasselbe Objekt sein sollten. Das reduziert Parameter und verbessert Perplexity zugleich, was selten genug ist, um es zu bemerken.

Um einen Korpus nach Bedeutung zu durchsuchen, brauchst du einen Vektor pro Satz. Sind diese vorhanden, ist die Suche trivial — das ist der Kern semantischer Retrieval-Systeme, und Kapitel 19 handelt von allem drumherum:

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]

Die einzige wirkliche Frage ist also, woher embed kommt. Der naheliegende Schritt ist, ein vortrainiertes Sprachmodell zu nehmen, jeden Satz hindurchlaufen zu lassen und die token-Vektoren zu mitteln. Hier steht diese Methode gegen vier Alternativen, auf zwei Arten bewertet: die Rangkorrelation zwischen Kosinus und menschlichen Ähnlichkeitsurteilen über die 1.379 Paare des STS-Benchmarks, und Top-1-Retrieval auf einem Index aus den 200 am stärksten paraphrasierten dieser Paare — eine Seite jedes Paars indexiert, die andere als Query verwendet.

wie der Satz eingebettet wirdRangkorrelationTop-1 auf einem 200-Satz-Index
binäre Wortüberlappung (gar kein Modell)0,550089,0 %
Mittelwert der oben trainierten statischen Vektoren0,526385,5 %
BERT, das [CLS] token0,203067,0 %
BERT, Mittelwert der token-Vektoren0,472984,0 %
MiniLM, kontrastiv trainiert0,820392,0 %

Lies die mittleren drei Zeilen gegen die ersten zwei. Ein vortrainierter transformer mit 109 Millionen Parametern, auf die naheliegende Weise verwendet, ist schlechter darin, Satzähnlichkeit zu beurteilen, als zu zählen, wie viele Wörter zwei Sätze teilen — und schlechter als der Mittelwert der vorhin trainierten 100-dimensionalen text8-Vektoren. Das [CLS] token, das Tutorials immer noch empfehlen, weil BERT mit einem daran angehängten Satz-Level-Ziel vortrainiert wurde, ist schlechter als die Hälfte davon.

Das ist kein Defekt in BERT. Es ist das Ziel. Ein Sprachmodell wird so trainiert, dass seine Hidden States ein token vorhersagen; nichts darin verlangt, dass zwei Paraphrasen nahe beieinander landen, und nichts belohnt eine Geometrie, in der Kosinus „gleiche Bedeutung“ heißt. Die letzte Zeile ist ein Modell mit einem Fünftel der Größe (22.713.216 Parameter), trainiert auf einem völlig anderen Loss: contrastive learning, bei dem die Beispiele Paare sind — eine Frage und ihre Antwort, ein Satz und seine Paraphrase — und das Ziel echte Paare zusammenzieht, während gesampelte Negative auseinander gedrückt werden. Das ist der Beitrag von Sentence-BERT und der Ursprung der ganzen embedding-model-Industrie.12 Dense Passage Retrieval wendet dasselbe Rezept direkt auf Suche an, mit einem Encoder für Queries und einem für Passagen.13

Also die praktische Regel:

Ein embedding model ist kein Sprachmodell mit entfernter letzter Schicht. Es ist ein anderes Modell mit einem anderen Ziel, meist viel kleiner, dessen Kosinus bedeutet, was du willst, weil es auf Paaren trainiert wurde, bei denen genau das das Ziel war. Die Tabelle oben zeigt die Kosten, wenn man das eine durch das andere ersetzt.

Und die Familie scheitert an Wortreihenfolge. „The dog bit the man“ und „the man bit the dog“ haben identische Bags of Words; Wortüberlappung und der Mittelwert statischer Vektoren geben ihnen also einen Kosinus von exakt 1,000000, und mean-pooled BERT, das Position durchaus sieht, landet immer noch fast dort — und selbst das kontrastiv trainierte MiniLM setzt sie noch bei 0,979. Wenn deine Retrieval-Aufgabe davon abhängt, wer wem was getan hat, wird dich kein Kosinus-Schwellenwert retten.

Kapitel 19 baut auf dieser Grundlage ein Produktions-Retrieval-System und kommt zu einem konkreten Kosinus-Cut-off. Die letzte Messung in diesem Kapitel macht so eine Zahl vertretbar statt magisch.

Der Fluch der Dimensionalität in einer Tabelle

Link zum Abschnitt: Der Fluch der Dimensionalität in einer Tabelle

Reale embeddings haben Hunderte oder Tausende Komponenten, und Abstände verhalten sich dort oben seltsam. Nimm 1.000 zufällige Punkte im Einheitswürfel von dd Dimensionen und betrachte das Verhältnis zwischen dem größten und dem kleinsten Abstand irgendeines Paars:

Dimensionennächstes Paarentferntestes PaarVerhältnis
20,00071,36121921,66
100,23612,33979,91
1003,00475,17521,72
1.00011,780914,03061,19
10.00039,615242,01251,06

In zehntausend Dimensionen ist das entfernteste Punktepaar nur 6 % weiter auseinander als das nächste Paar. Alles ist ungefähr gleich weit von allem anderen entfernt, „nächster Nachbar“ trägt nicht mehr viel Information, und das ist der Fluch der Dimensionalität — sowie ein Grund, warum große Vektordatenbanken keine exakte Nearest-Neighbour-Suche machen. Die andere Seite derselben Medaille macht Kosinus-Schwellenwerte praktikabel: Über tausend Paare zufälliger Einheitsvektoren gemessen liegt der mittlere Kosinus bei 0.0052-0.0052 in 100 Dimensionen und +0.0003+0.0003 in 768, mit Standardabweichungen von 0,0968 und 0,0357 — und in 768 Dimensionen überschreiten nur 0,2 % zufälliger Paare den Betrag 0,1. Eine gemessene Ähnlichkeit von 0,4 ist daher nicht „40 % ähnlich“; sie liegt weit außerhalb dessen, was Zufall produziert, weshalb Schwellen zwischen 0,3 und 0,7 Signal von Rauschen trennen, statt in dessen Mitte zu sitzen.

Das Modell in diesem Kapitel liest eine feste Zahl vorheriger Zeichen, schaut jedes nach und klebt die Ergebnisse der Reihe nach zusammen. Dieses Design hat zwei Probleme, und sie sind dasselbe Problem.

Schau noch einmal auf die context-Tabelle: Von drei Zeichen auf acht zu gehen hat die Parameter fast verdoppelt und 0,06 nats gebracht. Die Kosten wachsen linear mit dem context — jede zusätzliche Position braucht ihren eigenen Block der ersten Gewichtsmatrix — und der Nutzen tut es nicht. Schieb es auf tausend tokens, und allein die erste Schicht wiegt mehr als der Rest des Modells, größtenteils ausgegeben für Positionen, die für eine gegebene Vorhersage keine Rolle spielen.

Das ist auch das zweite Problem: Das Modell hat keine Möglichkeit zu entscheiden, welche der vorherigen tokens wichtig sind. Position zwei bekommt ihre eigenen Gewichte und Position sieben ihre eigenen, dauerhaft, egal was darin steht. Wenn das Modell nell buchstabiert, ist das entscheidende Zeichen das direkt davor. Wenn ein Satz ein Pronomen enthält, kann das Wort, das seinen Bezug festlegt, vierzig tokens zurückliegen — und kein fester Slot kann „vierzig zurück“ zugewiesen werden, weil es beim nächsten Mal sechs sein wird.

Was wir wollen, ist ein Modell, das für jede Vorhersage berechnet, wie stark jedes frühere token zählen sollte — Gewichte über den context, erzeugt durch den Inhalt statt festgelegt durch das Layout. Schreibe das sorgfältig hin, und es beginnt als etwas völlig Alltägliches: ein Durchschnitt über die vorherigen tokens. Dann lass die Gewichte dieses Durchschnitts lernen, und lass sie davon abhängen, welches token die Frage stellt.

Das ist attention, und sie ist Kapitel 9.


Ebenfalls lohnend parallel dazu: Kapitel 3 von Jurafsky und Martins Speech and Language Processing, das n-gram-Modelle, Smoothing und Perplexity deutlich sorgfältiger behandelt, als hier Platz ist, einschließlich der Frage, warum Interpolation und Back-off besser sind als Add-one; die Stanford-CS229-Notizen §17.1–17.2 für Sprachmodellierung von der probabilistischen Seite; und Linzens oben genanntes Paper, das kurz ist und sich komplett zu lesen lohnt.

  1. Das Beispiel zur Namensgenerierung, der Datensatz und die Entwicklung von einer Zähltabelle zu einem Netz im Bengio-Stil folgen Andrej Karpathys Reihe building makemore, deren erste zwei Teile die beste Begleitung zu diesem Kapitel sind.

  2. Shannon, C. E. Prediction and Entropy of Printed English. Bell System Technical Journal 30(1), S. 50–64 (1951). Versuchspersonen, die den nächsten Buchstaben englischen Textes raten, und die ursprüngliche Bits-pro-Zeichen-Messung.

  3. Shannon, C. E. A Mathematical Theory of Communication. Bell System Technical Journal 27 (1948). Das Source-Coding-Theorem und die Identifikation von Vorhersage mit Kompression.

  4. Bengio, Y., Ducharme, R., Vincent, P. und Jauvin, C. A Neural Probabilistic Language Model. Journal of Machine Learning Research 3, S. 1137–1155 (2003). Die oben verwendete Architektur: ein embedding pro Wort, über ein festes Fenster konkateniert, durch eine Hidden Layer, zu einem softmax über das Vokabular.

  5. Mikolov, T., Chen, K., Corrado, G. und Dean, J. Efficient Estimation of Word Representations in Vector Space. arXiv:1301.3781 (2013). CBOW und skip-gram sowie das oben verwendete Analogie-Set.

  6. Mikolov, T., Sutskever, I., Chen, K., Corrado, G. und Dean, J. Distributed Representations of Words and Phrases and their Compositionality. arXiv:1310.4546 (2013). Negative Sampling, Subsampling häufiger Wörter und die oben verwendete Rauschverteilung hoch 3/4.

  7. Pennington, J., Socher, R. und Manning, C. GloVe: Global Vectors for Word Representation. EMNLP 2014. Wortvektoren aus einer Faktorisierung der globalen Kookkurrenzmatrix statt aus gestreamten lokalen Fenstern.

  8. Linzen, T. Issues in evaluating semantic spaces using word analogies. RepEval 2016, arXiv:1606.07736. Die Quelle der offset-freien Baselines, die oben repliziert wurden.

  9. Peters, M. et al. Deep contextualized word representations. arXiv:1802.05365 (2018). ELMo: ein Vektor pro Vorkommen, berechnet von einem bidirektionalen Sprachmodell.

  10. Devlin, J., Chang, M.-W., Lee, K. und Toutanova, K. BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding. arXiv:1810.04805 (2018). Das im bank-Experiment gemessene Modell.

  11. Press, O. und Wolf, L. Using the Output Embedding to Improve Language Models. arXiv:1608.05859 (2016), und Inan, H., Khosravi, K. und Socher, R. Tying Word Vectors and Word Classifiers. arXiv:1611.01462 (2016). Zwei unabhängige Argumente für denselben Trick.

  12. Reimers, N. und Gurevych, I. Sentence-BERT: Sentence Embeddings using Siamese BERT-Networks. arXiv:1908.10084 (2019). Seine Eingangsmessung — mean-pooled BERT schneidet bei Satzähnlichkeit schlechter ab als gemittelte statische Vektoren — ist das, was die Tabelle oben reproduziert.

  13. Karpukhin, V. et al. Dense Passage Retrieval for Open-Domain Question Answering. arXiv:2004.04906 (2020). Kontrastives Training eines Two-Encoder-Retrievers; der direkte Vorfahr des Retrieval-Stacks aus Kapitel 19.

Bereit, LIA die Wahl zu überlassen?

Bau mit jedem KI-Modell an einem Ort — starte heute kostenlos.