<?xml version="1.0" encoding="utf-8" standalone="yes"?><rss version="2.0" xmlns:atom="http://www.w3.org/2005/Atom"><channel><title>Coding Practice | YuyaoGe's Website</title><link>https://geyuyao.com/category/coding-practice/</link><atom:link href="https://geyuyao.com/category/coding-practice/index.xml" rel="self" type="application/rss+xml"/><description>Coding Practice</description><generator>Wowchemy (https://wowchemy.com)</generator><language>en-us</language><lastBuildDate>Wed, 03 Jul 2024 00:00:00 +0000</lastBuildDate><image><url>https://geyuyao.com/media/icon_hucac340dfc176d8b4c8a8aa7a23204f12_18561_512x512_fill_lanczos_center_3.png</url><title>Coding Practice</title><link>https://geyuyao.com/category/coding-practice/</link></image><item><title>Coding Practice | Implementing a Tokenizer: Word-based and Byte Pair Encoding</title><link>https://geyuyao.com/post/tokenizer-en/</link><pubDate>Wed, 03 Jul 2024 00:00:00 +0000</pubDate><guid>https://geyuyao.com/post/tokenizer-en/</guid><description>
&lt;div class="travel-langswitch" role="group" aria-label="Language">
&lt;span class="travel-langswitch__btn is-active" aria-current="true">English&lt;/span>
&lt;a class="travel-langswitch__btn" href="https://geyuyao.com/post/tokenizer/">中文&lt;/a>
&lt;/div>
&lt;h1 id="introduction">Introduction&lt;/h1>
&lt;p>A &lt;strong>tokenizer&lt;/strong> is a key tool in natural language processing. Its job is to split a text string into lexical units, or tokens. Tokenizers play a crucial role during text preprocessing, providing the foundation for downstream tasks such as text analysis, information retrieval, and machine translation.&lt;/p>
&lt;h2 id="origins-and-development">Origins and Development&lt;/h2>
&lt;p>The history of tokenizers goes back to the early days of computational linguistics and information retrieval. As computer science matured, researchers realized that processing natural language text requires splitting a continuous sequence of characters into basic units that carry meaning (called morphemes in linguistics). Early tokenization methods were mainly rule- and dictionary-based, relying on predefined rules and word lists to identify word boundaries. With the rise of statistical NLP and machine learning, statistical models and data-driven approaches gradually replaced rule-based ones.&lt;/p>
&lt;p>In recent years, advances in deep learning have pushed tokenization further. In particular, the introduction of pretrained language models such as BERT (Bidirectional Encoder Representations from Transformers) brought subword-level tokenization methods such as Byte Pair Encoding (BPE), WordPiece, and SentencePiece, which handle out-of-vocabulary words and linguistic diversity far better.&lt;/p>
&lt;h2 id="motivation">Motivation&lt;/h2>
&lt;p>To understand tokenizers, we first need to understand why we use them at all.&lt;/p>
&lt;p>In NLP tasks, the data we work with is usually raw text.&lt;/p>
&lt;p>Take the following sentence as an example:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-fallback" data-lang="fallback">&lt;span class="line">&lt;span class="cl">Jim Henson was a puppeteer
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Models, however, can only process numbers, so we need a way to turn raw text into numbers. That is exactly what a tokenizer does.&lt;/p>
&lt;p>In short, the goal of a tokenizer is to convert text that humans understand into numbers that machines understand.&lt;/p>
&lt;p>Next, we introduce the simplest and most intuitive tokenization method — the &lt;strong>Word-based Tokenizer&lt;/strong>.&lt;/p>
&lt;h1 id="the-simplest-and-most-intuitive-approach--word-based-tokenizer">The Simplest and Most Intuitive Approach — Word-based Tokenizer&lt;/h1>
&lt;p>&lt;em>Part of this section draws on the Hugging Face NLP Course&lt;/em> &lt;sup id="fnref:1">&lt;a href="#fn:1" class="footnote-ref" role="doc-noteref">1&lt;/a>&lt;/sup>&lt;/p>
&lt;p>Picking up where we left off: how do we turn the text &lt;code>Jim Henson was a puppeteer&lt;/code> into numbers?&lt;/p>
&lt;p>One intuitive approach is to split the string on whitespace and assign each word a unique number as its index.&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-fallback" data-lang="fallback">&lt;span class="line">&lt;span class="cl">tokenized_text = &amp;#34;Jim Henson was a puppeteer&amp;#34;.split()
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">print(tokenized_text)
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-fallback" data-lang="fallback">&lt;span class="line">&lt;span class="cl">[&amp;#39;Jim&amp;#39;, &amp;#39;Henson&amp;#39;, &amp;#39;was&amp;#39;, &amp;#39;a&amp;#39;, &amp;#39;puppeteer&amp;#39;]
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>This way each word gets an ID, starting from 0 and running up to the size of the vocabulary. The model uses these IDs to identify each word.&lt;/p>
&lt;p>If we wanted a word-based tokenizer to cover an entire language, we would have to assign a unique integer index to every word in that language, which produces an enormous number of indices. English, for example, has more than 500,000 words, so building a mapping from every word to an index would require a dictionary with 500,000 entries. And that is not the only drawback of word-based tokenization.&lt;/p>
&lt;p>Furthermore, a word like &amp;ldquo;dog&amp;rdquo; is represented differently from a word like &amp;ldquo;dogs&amp;rdquo;, and the model has no way of knowing that &amp;ldquo;dog&amp;rdquo; and &amp;ldquo;dogs&amp;rdquo; are related: it treats them as two unrelated words. The same applies to other similar pairs, such as &amp;ldquo;run&amp;rdquo; and &amp;ldquo;running&amp;rdquo; — the model will not consider them similar.&lt;/p>
&lt;p>We also need a special token to represent words that are not in our vocabulary. This is the &amp;ldquo;unknown&amp;rdquo; token, usually written as &amp;ldquo;[UNK]&amp;rdquo; or &amp;ldquo;&amp;lt;unk&amp;gt;&amp;rdquo;. If you see a tokenizer producing many of these tokens, it is generally a bad sign: it means the tokenizer could not retrieve an index for a word, so information is lost during tokenization — a fatal shortcoming.&lt;/p>
&lt;p>To summarize, the word-based tokenizer has several major drawbacks:&lt;/p>
&lt;ol>
&lt;li>&lt;strong>Huge vocabulary&lt;/strong>: because the vocabulary of a language or application domain is very large, a word-based tokenizer needs a very large vocabulary. This not only increases storage and computation costs, but also degrades training and inference efficiency. In practice, many words are low-frequency, and including large numbers of rare words in the vocabulary leads to data sparsity problems.&lt;/li>
&lt;li>&lt;strong>Difficulty handling morphological variation&lt;/strong>: many languages have rich morphology (plurals and tenses in English, grammatical gender in French, and so on). A word-based tokenizer has to include every one of these surface forms, inflating the vocabulary even further. Worse, different forms of the same word are treated as distinct tokens, so the model cannot share semantic information across them.&lt;/li>
&lt;li>&lt;strong>Limited coverage, unable to handle out-of-vocabulary words&lt;/strong>: a word-based tokenizer requires a fixed vocabulary. New words or misspellings that are not in the vocabulary simply cannot be handled, which hurts generalization.&lt;/li>
&lt;/ol>
&lt;h1 id="a-finer-grained-approach--byte-pair-encoding-bpe">A Finer-grained Approach — Byte Pair Encoding (BPE)&lt;/h1>
&lt;p>As discussed above, word-based tokenizers suffer from large vocabularies, difficulty with morphological variation, and limited coverage. Is there a tokenization method that solves all three?&lt;/p>
&lt;p>Notice that English words reuse a great many letter combinations. Take &amp;ldquo;happy&amp;rdquo;: many words are built on it as a root, such as &amp;ldquo;happily&amp;rdquo;, &amp;ldquo;happiness&amp;rdquo;, and &amp;ldquo;unhappy&amp;rdquo;. These derived words can be decomposed into reusable letter combinations — the sequence &amp;ldquo;happ&amp;rdquo;, for instance, appears in all three. Beyond that, combinations like &amp;ldquo;ily&amp;rdquo; and &amp;ldquo;un&amp;rdquo; occur frequently in other English words, and pairing them with other fragments can express new words.&lt;/p>
&lt;p>Naturally, this suggests a question: could we design a tokenizer around the idea of finding frequently occurring letter combinations, covering every English word by reusing them?&lt;/p>
&lt;p>If so, such a tokenizer would resolve all three drawbacks of the word-based tokenizer, because:&lt;/p>
&lt;ol>
&lt;li>&lt;strong>High reuse means a small vocabulary&lt;/strong>: since words are built from letter combinations that are heavily reused, only a small number of combinations are needed to cover most English words.&lt;/li>
&lt;li>&lt;strong>Easy handling of morphological variation&lt;/strong>: in the example above, our vocabulary contains tokens like &amp;ldquo;ily&amp;rdquo; and &amp;ldquo;un&amp;rdquo;, which also appear frequently in other inflected words, making the method well suited to morphological variation.&lt;/li>
&lt;li>&lt;strong>Broad word coverage&lt;/strong>: the method starts from individual letters and searches for letter combinations, so it can cover every English word.&lt;/li>
&lt;/ol>
&lt;p>This method is in fact &lt;strong>Byte Pair Encoding (BPE)&lt;/strong>.&lt;/p>
&lt;p>&lt;strong>Byte Pair Encoding (BPE)&lt;/strong> is a subword segmentation technique that originated in data compression and is now widely used in natural language processing. &lt;strong>The core idea of BPE is to iteratively merge the most frequent byte pair, building up subword units step by step and thereby segmenting the text.&lt;/strong> The method was first proposed by Philip Gage in 1994&lt;sup id="fnref:2">&lt;a href="#fn:2" class="footnote-ref" role="doc-noteref">2&lt;/a>&lt;/sup> for file compression, and was later introduced into machine translation by Sennrich et al. in 2015&lt;sup id="fnref:3">&lt;a href="#fn:3" class="footnote-ref" role="doc-noteref">3&lt;/a>&lt;/sup> to address oversized vocabularies and sparsity.&lt;/p>
&lt;p>Many well-known LLMs use BPE for tokenization, including the GPT series, BERT, RoBERTa, and T5. It is fair to say that BPE is now one of the classic algorithms of natural language processing.&lt;/p>
&lt;p>In the next section we will build a BPE tokenizer step by step, following Andrej Karpathy&amp;rsquo;s project &lt;a href="https://github.com/karpathy/minbpe" target="_blank" rel="noopener">minbpe&lt;/a>.&lt;/p>
&lt;blockquote>
&lt;p>It is worth mentioning that &lt;a href="https://github.com/karpathy" target="_blank" rel="noopener">Karpathy&lt;/a> has many well-known LLM projects, such as &lt;a href="https://github.com/karpathy/llm.c" target="_blank" rel="noopener">llm.c&lt;/a> and &lt;a href="https://github.com/karpathy/llama2.c" target="_blank" rel="noopener">llama2.c&lt;/a>.&lt;/p>
&lt;/blockquote>
&lt;h1 id="code-practice-implementing-bpe-from-scratch">Code Practice: Implementing BPE from Scratch&lt;/h1>
&lt;p>In this section we build a BPE tokenizer step by step, following Andrej Karpathy&amp;rsquo;s project &lt;a href="https://github.com/karpathy/minbpe" target="_blank" rel="noopener">minbpe&lt;/a>&lt;sup id="fnref:4">&lt;a href="#fn:4" class="footnote-ref" role="doc-noteref">4&lt;/a>&lt;/sup>.&lt;/p>
&lt;h2 id="the-base-class-tokenizer">The Base Class Tokenizer&lt;/h2>
&lt;p>We start with the base form of BPE, namely the definition of the &lt;code>Tokenizer&lt;/code> class:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-python" data-lang="python">&lt;span class="line">&lt;span class="cl">&lt;span class="k">class&lt;/span> &lt;span class="nc">Tokenizer&lt;/span>&lt;span class="p">:&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="s2">&amp;#34;&amp;#34;&amp;#34;Base class for Tokenizers&amp;#34;&amp;#34;&amp;#34;&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">def&lt;/span> &lt;span class="fm">__init__&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="bp">self&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># default: vocab size of 256 (all bytes), no merges, no patterns&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">merges&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="p">{}&lt;/span> &lt;span class="c1"># (int, int) -&amp;gt; int&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">vocab&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">_build_vocab&lt;/span>&lt;span class="p">()&lt;/span> &lt;span class="c1"># int -&amp;gt; bytes&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">def&lt;/span> &lt;span class="nf">train&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="bp">self&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">text&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">vocab_size&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">verbose&lt;/span>&lt;span class="o">=&lt;/span>&lt;span class="kc">False&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># Tokenizer can train a vocabulary of size vocab_size from text&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">raise&lt;/span> &lt;span class="ne">NotImplementedError&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">def&lt;/span> &lt;span class="nf">encode&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="bp">self&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">text&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># Tokenizer can encode a string into a list of integers&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">raise&lt;/span> &lt;span class="ne">NotImplementedError&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">def&lt;/span> &lt;span class="nf">decode&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="bp">self&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">ids&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># Tokenizer can decode a list of integers into a string&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">raise&lt;/span> &lt;span class="ne">NotImplementedError&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">def&lt;/span> &lt;span class="nf">_build_vocab&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="bp">self&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># vocab is simply and deterministically derived from merges&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">vocab&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="p">{&lt;/span>&lt;span class="n">idx&lt;/span>&lt;span class="p">:&lt;/span> &lt;span class="nb">bytes&lt;/span>&lt;span class="p">([&lt;/span>&lt;span class="n">idx&lt;/span>&lt;span class="p">])&lt;/span> &lt;span class="k">for&lt;/span> &lt;span class="n">idx&lt;/span> &lt;span class="ow">in&lt;/span> &lt;span class="nb">range&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="mi">256&lt;/span>&lt;span class="p">)}&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">return&lt;/span> &lt;span class="n">vocab&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>The &lt;code>Tokenizer&lt;/code> class has four basic functions — &lt;code>init&lt;/code>, &lt;code>train&lt;/code>, &lt;code>encode&lt;/code>, and &lt;code>decode&lt;/code> — corresponding to initializing, training, encoding, and decoding. Among them, &lt;code>train&lt;/code>, &lt;code>encode&lt;/code>, and &lt;code>decode&lt;/code> are virtual functions that will be overridden by the subclass.&lt;/p>
&lt;p>The &lt;a href="https://github.com/karpathy/minbpe" target="_blank" rel="noopener">minbpe&lt;/a> project also supports saving and loading tokenizer models. Since this is only loosely related to our topic, and for reasons of space, we will not go into it here; interested readers can consult the source&lt;sup id="fnref:5">&lt;a href="#fn:5" class="footnote-ref" role="doc-noteref">5&lt;/a>&lt;/sup>.&lt;/p>
&lt;p>In the &lt;code>init&lt;/code> function we define two variables, &lt;code>self.merges = {}&lt;/code> and &lt;code>self.vocab&lt;/code>. The former records, during tokenization, the mapping from the indices of two adjacent original tokens to the index of the new token they merge into, i.e. &lt;code>(int, int) -&amp;gt; int&lt;/code>. The latter represents the vocabulary, recording the mapping from an index to a token, i.e. &lt;code>int -&amp;gt; bytes&lt;/code>.&lt;/p>
&lt;p>&lt;code>self.vocab&lt;/code> is initialized inside &lt;code>init&lt;/code>, where we map 256 integers used as indices (&lt;code>idx&lt;/code>) to the 256 tokens in hexadecimal representation (&lt;code>bytes([idx])&lt;/code>).&lt;/p>
&lt;p>If we print &lt;code>vocab&lt;/code> inside the &lt;code>_build_vocab&lt;/code> function, we get:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-txt" data-lang="txt">&lt;span class="line">&lt;span class="cl">{0: b&amp;#39;\x00&amp;#39;, 1: b&amp;#39;\x01&amp;#39;, 2: b&amp;#39;\x02&amp;#39;, 3: b&amp;#39;\x03&amp;#39;, ..., 254: b&amp;#39;\xfe&amp;#39;, 255: b&amp;#39;\xff&amp;#39;}
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;h2 id="bpe">BPE&lt;/h2>
&lt;p>Next, here is the definition of the BPE class:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-python" data-lang="python">&lt;span class="line">&lt;span class="cl">&lt;span class="k">class&lt;/span> &lt;span class="nc">BPE&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">Tokenizer&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">def&lt;/span> &lt;span class="fm">__init__&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="bp">self&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="nb">super&lt;/span>&lt;span class="p">()&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="fm">__init__&lt;/span>&lt;span class="p">()&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">def&lt;/span> &lt;span class="nf">train&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="bp">self&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">text&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">vocab_size&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">verbose&lt;/span>&lt;span class="o">=&lt;/span>&lt;span class="kc">False&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">def&lt;/span> &lt;span class="nf">get_stats&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">ids&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">counts&lt;/span>&lt;span class="o">=&lt;/span>&lt;span class="kc">None&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="s2">&amp;#34;&amp;#34;&amp;#34;
&lt;/span>&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="s2"> Given a list of integers, return a dictionary of counts of consecutive pairs
&lt;/span>&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="s2"> Example: [1, 2, 3, 1, 2] -&amp;gt; {(1, 2): 2, (2, 3): 1, (3, 1): 1}
&lt;/span>&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="s2"> Optionally allows to update an existing dictionary of counts
&lt;/span>&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="s2"> &amp;#34;&amp;#34;&amp;#34;&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">counts&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="p">{}&lt;/span> &lt;span class="k">if&lt;/span> &lt;span class="n">counts&lt;/span> &lt;span class="ow">is&lt;/span> &lt;span class="kc">None&lt;/span> &lt;span class="k">else&lt;/span> &lt;span class="n">counts&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">for&lt;/span> &lt;span class="n">pair&lt;/span> &lt;span class="ow">in&lt;/span> &lt;span class="nb">zip&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">ids&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">ids&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="mi">1&lt;/span>&lt;span class="p">:]):&lt;/span> &lt;span class="c1"># iterate consecutive elements&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">counts&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="n">pair&lt;/span>&lt;span class="p">]&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">counts&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">get&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">pair&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="mi">0&lt;/span>&lt;span class="p">)&lt;/span> &lt;span class="o">+&lt;/span> &lt;span class="mi">1&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">return&lt;/span> &lt;span class="n">counts&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">def&lt;/span> &lt;span class="nf">merge&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">ids&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">pair&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">idx&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="s2">&amp;#34;&amp;#34;&amp;#34;
&lt;/span>&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="s2"> In the list of integers (ids), replace all consecutive occurrences
&lt;/span>&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="s2"> of pair with the new integer token idx
&lt;/span>&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="s2"> Example: ids=[1, 2, 3, 1, 2], pair=(1, 2), idx=4 -&amp;gt; [4, 3, 4]
&lt;/span>&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="s2"> &amp;#34;&amp;#34;&amp;#34;&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">newids&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="p">[]&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">i&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="mi">0&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">while&lt;/span> &lt;span class="n">i&lt;/span> &lt;span class="o">&amp;lt;&lt;/span> &lt;span class="nb">len&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">ids&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># if not at the very last position AND the pair matches, replace it&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">if&lt;/span> &lt;span class="n">ids&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="n">i&lt;/span>&lt;span class="p">]&lt;/span> &lt;span class="o">==&lt;/span> &lt;span class="n">pair&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="mi">0&lt;/span>&lt;span class="p">]&lt;/span> &lt;span class="ow">and&lt;/span> &lt;span class="n">i&lt;/span> &lt;span class="o">&amp;lt;&lt;/span> &lt;span class="nb">len&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">ids&lt;/span>&lt;span class="p">)&lt;/span> &lt;span class="o">-&lt;/span> &lt;span class="mi">1&lt;/span> &lt;span class="ow">and&lt;/span> &lt;span class="n">ids&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="n">i&lt;/span>&lt;span class="o">+&lt;/span>&lt;span class="mi">1&lt;/span>&lt;span class="p">]&lt;/span> &lt;span class="o">==&lt;/span> &lt;span class="n">pair&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="mi">1&lt;/span>&lt;span class="p">]:&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">newids&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">append&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">idx&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">i&lt;/span> &lt;span class="o">+=&lt;/span> &lt;span class="mi">2&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">else&lt;/span>&lt;span class="p">:&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">newids&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">append&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">ids&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="n">i&lt;/span>&lt;span class="p">])&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">i&lt;/span> &lt;span class="o">+=&lt;/span> &lt;span class="mi">1&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">return&lt;/span> &lt;span class="n">newids&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">assert&lt;/span> &lt;span class="n">vocab_size&lt;/span> &lt;span class="o">&amp;gt;=&lt;/span> &lt;span class="mi">256&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">num_merges&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">vocab_size&lt;/span> &lt;span class="o">-&lt;/span> &lt;span class="mi">256&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># input text preprocessing&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">text_bytes&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">text&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">encode&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="s2">&amp;#34;utf-8&amp;#34;&lt;/span>&lt;span class="p">)&lt;/span> &lt;span class="c1"># raw bytes&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">ids&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="nb">list&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">text_bytes&lt;/span>&lt;span class="p">)&lt;/span> &lt;span class="c1"># list of integers in range 0..255&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># iteratively merge the most common pairs to create new tokens&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">merges&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="p">{}&lt;/span> &lt;span class="c1"># (int, int) -&amp;gt; int&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">vocab&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="p">{&lt;/span>&lt;span class="n">idx&lt;/span>&lt;span class="p">:&lt;/span> &lt;span class="nb">bytes&lt;/span>&lt;span class="p">([&lt;/span>&lt;span class="n">idx&lt;/span>&lt;span class="p">])&lt;/span> &lt;span class="k">for&lt;/span> &lt;span class="n">idx&lt;/span> &lt;span class="ow">in&lt;/span> &lt;span class="nb">range&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="mi">256&lt;/span>&lt;span class="p">)}&lt;/span> &lt;span class="c1"># int -&amp;gt; bytes&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">for&lt;/span> &lt;span class="n">i&lt;/span> &lt;span class="ow">in&lt;/span> &lt;span class="nb">range&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">num_merges&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># count up the number of times every consecutive pair appears&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">stats&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">get_stats&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">ids&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># find the pair with the highest count&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">pair&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="nb">max&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">stats&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">key&lt;/span>&lt;span class="o">=&lt;/span>&lt;span class="n">stats&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">get&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">pair&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># mint a new token: assign it the next available id&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">idx&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="mi">256&lt;/span> &lt;span class="o">+&lt;/span> &lt;span class="n">i&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># replace all occurrences of pair in ids with idx&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">ids&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">merge&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">ids&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">pair&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">idx&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># save the merge&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">merges&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="n">pair&lt;/span>&lt;span class="p">]&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">idx&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">vocab&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="n">idx&lt;/span>&lt;span class="p">]&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">vocab&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="n">pair&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="mi">0&lt;/span>&lt;span class="p">]]&lt;/span> &lt;span class="o">+&lt;/span> &lt;span class="n">vocab&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="n">pair&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="mi">1&lt;/span>&lt;span class="p">]]&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">if&lt;/span> &lt;span class="n">verbose&lt;/span>&lt;span class="p">:&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="sa">f&lt;/span>&lt;span class="s2">&amp;#34;merge &lt;/span>&lt;span class="si">{&lt;/span>&lt;span class="n">i&lt;/span>&lt;span class="o">+&lt;/span>&lt;span class="mi">1&lt;/span>&lt;span class="si">}&lt;/span>&lt;span class="s2">/&lt;/span>&lt;span class="si">{&lt;/span>&lt;span class="n">num_merges&lt;/span>&lt;span class="si">}&lt;/span>&lt;span class="s2">: &lt;/span>&lt;span class="si">{&lt;/span>&lt;span class="n">pair&lt;/span>&lt;span class="si">}&lt;/span>&lt;span class="s2"> -&amp;gt; &lt;/span>&lt;span class="si">{&lt;/span>&lt;span class="n">idx&lt;/span>&lt;span class="si">}&lt;/span>&lt;span class="s2"> (&lt;/span>&lt;span class="si">{&lt;/span>&lt;span class="n">vocab&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="n">idx&lt;/span>&lt;span class="p">]&lt;/span>&lt;span class="si">}&lt;/span>&lt;span class="s2">) had &lt;/span>&lt;span class="si">{&lt;/span>&lt;span class="n">stats&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="n">pair&lt;/span>&lt;span class="p">]&lt;/span>&lt;span class="si">}&lt;/span>&lt;span class="s2"> occurrences&amp;#34;&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># save class variables&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">merges&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">merges&lt;/span> &lt;span class="c1"># used in encode()&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">vocab&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">vocab&lt;/span> &lt;span class="c1"># used in decode()&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">def&lt;/span> &lt;span class="nf">decode&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="bp">self&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">ids&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># given ids (list of integers), return Python string&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">text_bytes&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="sa">b&lt;/span>&lt;span class="s2">&amp;#34;&amp;#34;&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">join&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">vocab&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="n">idx&lt;/span>&lt;span class="p">]&lt;/span> &lt;span class="k">for&lt;/span> &lt;span class="n">idx&lt;/span> &lt;span class="ow">in&lt;/span> &lt;span class="n">ids&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">text&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">text_bytes&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">decode&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="s2">&amp;#34;utf-8&amp;#34;&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">errors&lt;/span>&lt;span class="o">=&lt;/span>&lt;span class="s2">&amp;#34;replace&amp;#34;&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">return&lt;/span> &lt;span class="n">text&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">def&lt;/span> &lt;span class="nf">encode&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="bp">self&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">text&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># given a string text, return the token ids&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">text_bytes&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">text&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">encode&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="s2">&amp;#34;utf-8&amp;#34;&lt;/span>&lt;span class="p">)&lt;/span> &lt;span class="c1"># raw bytes&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">ids&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="nb">list&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">text_bytes&lt;/span>&lt;span class="p">)&lt;/span> &lt;span class="c1"># list of integers in range 0..255&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">while&lt;/span> &lt;span class="nb">len&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">ids&lt;/span>&lt;span class="p">)&lt;/span> &lt;span class="o">&amp;gt;=&lt;/span> &lt;span class="mi">2&lt;/span>&lt;span class="p">:&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># find the pair with the lowest merge index&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">stats&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">get_stats&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">ids&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">pair&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="nb">min&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">stats&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">key&lt;/span>&lt;span class="o">=&lt;/span>&lt;span class="k">lambda&lt;/span> &lt;span class="n">p&lt;/span>&lt;span class="p">:&lt;/span> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">merges&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">get&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">p&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="nb">float&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="s2">&amp;#34;inf&amp;#34;&lt;/span>&lt;span class="p">)))&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># subtle: if there are no more merges available, the key will&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># result in an inf for every single pair, and the min will be&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># just the first pair in the list, arbitrarily&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># we can detect this terminating case by a membership check&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">if&lt;/span> &lt;span class="n">pair&lt;/span> &lt;span class="ow">not&lt;/span> &lt;span class="ow">in&lt;/span> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">merges&lt;/span>&lt;span class="p">:&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">break&lt;/span> &lt;span class="c1"># nothing else can be merged anymore&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="c1"># otherwise let&amp;#39;s merge the best pair (lowest merge index)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">idx&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">merges&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="n">pair&lt;/span>&lt;span class="p">]&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">ids&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">merge&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">ids&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">pair&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">idx&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">return&lt;/span> &lt;span class="n">ids&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>The BPE class overrides &lt;code>train&lt;/code>, &lt;code>encode&lt;/code>, and &lt;code>decode&lt;/code>, while &lt;code>init&lt;/code> is inherited from the &lt;code>Tokenizer&lt;/code> class.&lt;/p>
&lt;p>Let us walk through the functions in the order &lt;code>train&lt;/code>, &lt;code>encode&lt;/code>, &lt;code>decode&lt;/code>.&lt;/p>
&lt;h3 id="bpetrain">BPE.train&lt;/h3>
&lt;p>Inside &lt;code>train&lt;/code> there are two helper functions, &lt;code>get_stats&lt;/code> and &lt;code>merge&lt;/code>.&lt;/p>
&lt;p>&lt;code>def get_stats(ids)&lt;/code>: counts how often each pair of adjacent tokens occurs. For example, when &lt;code>ids = [1, 2, 3, 1, 2]&lt;/code>, the return value is &lt;code>{(1, 2): 2, (2, 3): 1, (3, 1): 1}&lt;/code>&lt;sup id="fnref:6">&lt;a href="#fn:6" class="footnote-ref" role="doc-noteref">6&lt;/a>&lt;/sup>.&lt;/p>
&lt;p>&lt;code>def merge(ids, pair, idx)&lt;/code>: merges every occurrence of &lt;code>pair&lt;/code> in &lt;code>ids&lt;/code> into &lt;code>idx&lt;/code> and returns the new list of ids. For example, with &lt;code>ids=[1, 2, 3, 1, 2], pair=(1, 2), idx=4 -&amp;gt; [4, 3, 4]&lt;/code>, we replace the adjacent &lt;code>1, 2&lt;/code> in &lt;code>ids&lt;/code> with &lt;code>4&lt;/code>, yielding the new ids &lt;code>[4, 3, 4]&lt;/code>.&lt;/p>
&lt;p>In &lt;code>train&lt;/code> we first need to convert the characters of the string into indices, which we achieve with &lt;code>text_bytes = text.encode(&amp;quot;utf-8&amp;quot;) &lt;/code> and &lt;code>ids = list(text_bytes) &lt;/code>. For instance, when &lt;code>text = 'aaabbc'&lt;/code>, we get &lt;code>ids = [97, 97, 97, 98, 98, 99]&lt;/code>.&lt;/p>
&lt;p>Next, inside the loop, &lt;code>get_stats&lt;/code> first &lt;strong>counts the frequency of each adjacent token pair in&lt;/strong> &lt;code>ids&lt;/code>. The astute reader will already have guessed the purpose: it is to &lt;strong>merge these high-frequency pairs into new tokens&lt;/strong>, and that is exactly what &lt;code>merge&lt;/code> does. At this point the length of &lt;code>ids&lt;/code> has been reduced. Finally, we record the mapping from token pair to index in &lt;code>merges&lt;/code> and the mapping from index to new token in the vocabulary &lt;code>vocab&lt;/code>.&lt;/p>
&lt;p>This process repeats until the loop finishes.&lt;/p>
&lt;p>So what is &lt;code>train&lt;/code> actually doing?&lt;/p>
&lt;p>In essence, &lt;code>train&lt;/code> merges frequently co-occurring adjacent token pairs into a single new token. By adding only a small number of mappings to the vocabulary, we can substantially shorten the training text. Note that at the start the vocabulary contains only single letters. As training proceeds, frequent letter pairs are gradually added — first pairs of two letters, then longer letter sequences, and so on.&lt;/p>
&lt;p>Taking the training text &lt;code>text = &amp;quot;happily happiness unhappy&amp;quot;&lt;/code> as an example, the log produced by each of the three rounds is:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-txt" data-lang="txt">&lt;span class="line">&lt;span class="cl">merge 1/3: (104, 97) -&amp;gt; 256 (b&amp;#39;ha&amp;#39;) had 3 occurrences
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">merge 2/3: (256, 112) -&amp;gt; 257 (b&amp;#39;hap&amp;#39;) had 3 occurrences
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">merge 3/3: (257, 112) -&amp;gt; 258 (b&amp;#39;happ&amp;#39;) had 3 occurrences
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Notice that after three iterations, &lt;code>happ&lt;/code> has appeared in the vocabulary. This means the frequent letter combination &lt;code>happ&lt;/code> can now be used in subsequent decoding and encoding.&lt;/p>
&lt;h3 id="bpeencode">BPE.encode&lt;/h3>
&lt;p>In &lt;code>encode&lt;/code> we again start by finding the frequently occurring token pairs in the text. We then look for the token pair with the smallest index in &lt;code>merges&lt;/code>. The reason is that during training, a smaller index means a higher frequency in the training set, which is statistically sensible&lt;sup id="fnref:7">&lt;a href="#fn:7" class="footnote-ref" role="doc-noteref">7&lt;/a>&lt;/sup>. We then map that token pair to its new token according to the mapping in &lt;code>merges&lt;/code>, producing a new token sequence. This process repeats until no token pair can be merged (reduced) any further.&lt;/p>
&lt;h3 id="bpedecode">BPE.decode&lt;/h3>
&lt;p>The &lt;code>decode&lt;/code> function takes a list of integer indices &lt;code>ids&lt;/code>, converts it into the corresponding byte sequence, and decodes that byte sequence into a UTF-8 string, returning the decoded version of the original text. This function is the inverse of &lt;code>encode&lt;/code>.&lt;/p>
&lt;h1 id="conclusion">Conclusion&lt;/h1>
&lt;p>To summarize, training a BPE amounts to finding the token combinations that occur frequently in the training set; encoding replaces mergeable tokens in the text with new tokens to produce a new token sequence; and decoding is the inverse of encoding.&lt;/p>
&lt;p>We can now explain the three advantages of BPE mentioned earlier in this article.&lt;/p>
&lt;ol>
&lt;li>&lt;strong>High reuse means a small vocabulary&lt;/strong>: unlike a word-based tokenizer, BPE does not need a mapping for every single word — it only needs mappings for frequently occurring token sequences.&lt;/li>
&lt;li>&lt;strong>Easy handling of morphological variation&lt;/strong>: however a word varies, frequently occurring token sequences such as the past-tense &lt;code>ed&lt;/code> or the adverbial &lt;code>ly&lt;/code> are already in the vocabulary and can be used efficiently during encoding.&lt;/li>
&lt;li>&lt;strong>Broad word coverage&lt;/strong>: even if the current word contains a token sequence that is absent from the vocabulary, we can still represent it with the individual letters that are in the vocabulary. BPE can therefore cover every English word.&lt;/li>
&lt;/ol>
&lt;div class="footnotes" role="doc-endnotes">
&lt;hr>
&lt;ol>
&lt;li id="fn:1">
&lt;p>&lt;a href="https://huggingface.co/learn/nlp-course/zh-CN/chapter2/4" target="_blank" rel="noopener">https://huggingface.co/learn/nlp-course/zh-CN/chapter2/4&lt;/a>&amp;#160;&lt;a href="#fnref:1" class="footnote-backref" role="doc-backlink">&amp;#x21a9;&amp;#xfe0e;&lt;/a>&lt;/p>
&lt;/li>
&lt;li id="fn:2">
&lt;p>Gage P. A new algorithm for data compression[J]. The C Users Journal, 1994, 12(2): 23-38.&amp;#160;&lt;a href="#fnref:2" class="footnote-backref" role="doc-backlink">&amp;#x21a9;&amp;#xfe0e;&lt;/a>&lt;/p>
&lt;/li>
&lt;li id="fn:3">
&lt;p>Sennrich R, Haddow B, Birch A. Neural machine translation of rare words with subword units[J]. arXiv preprint arXiv:1508.07909, 2015.&amp;#160;&lt;a href="#fnref:3" class="footnote-backref" role="doc-backlink">&amp;#x21a9;&amp;#xfe0e;&lt;/a>&lt;/p>
&lt;/li>
&lt;li id="fn:4">
&lt;p>&lt;a href="https://github.com/karpathy/minbpe" target="_blank" rel="noopener">https://github.com/karpathy/minbpe&lt;/a>&amp;#160;&lt;a href="#fnref:4" class="footnote-backref" role="doc-backlink">&amp;#x21a9;&amp;#xfe0e;&lt;/a>&lt;/p>
&lt;/li>
&lt;li id="fn:5">
&lt;p>&lt;a href="https://github.com/karpathy/minbpe/blob/master/minbpe/base.py" target="_blank" rel="noopener">https://github.com/karpathy/minbpe/blob/master/minbpe/base.py&lt;/a>&amp;#160;&lt;a href="#fnref:5" class="footnote-backref" role="doc-backlink">&amp;#x21a9;&amp;#xfe0e;&lt;/a>&lt;/p>
&lt;/li>
&lt;li id="fn:6">
&lt;p>Of course, in a real string the letter &lt;code>a&lt;/code> corresponds to index &lt;code>97&lt;/code> and so on; &lt;code>[1, 2, 3, 1, 2]&lt;/code> is used purely for ease of understanding.&amp;#160;&lt;a href="#fnref:6" class="footnote-backref" role="doc-backlink">&amp;#x21a9;&amp;#xfe0e;&lt;/a>&lt;/p>
&lt;/li>
&lt;li id="fn:7">
&lt;p>Note that BPE is a heuristic tokenization method and does not aim for an optimal solution, so the token sequence produced by encoding is not the shortest possible, only a relatively short one.&amp;#160;&lt;a href="#fnref:7" class="footnote-backref" role="doc-backlink">&amp;#x21a9;&amp;#xfe0e;&lt;/a>&lt;/p>
&lt;/li>
&lt;/ol>
&lt;/div></description></item><item><title>Coding Practice | Implementing Self-Attention from Scratch</title><link>https://geyuyao.com/post/transformer-en/</link><pubDate>Sun, 10 Mar 2024 00:00:00 +0000</pubDate><guid>https://geyuyao.com/post/transformer-en/</guid><description>
&lt;div class="travel-langswitch" role="group" aria-label="Language">
&lt;span class="travel-langswitch__btn is-active" aria-current="true">English&lt;/span>
&lt;a class="travel-langswitch__btn" href="https://geyuyao.com/post/transformer/">中文&lt;/a>
&lt;/div>
&lt;h1 id="introduction">Introduction&lt;/h1>
&lt;p>In contemporary natural language processing (NLP) and deep learning, the Transformer and its core component — the self-attention mechanism — have revolutionized sequence modeling. Since Vaswani et al. first introduced it in the 2017 paper &lt;a href="https://arxiv.org/abs/1706.03762" target="_blank" rel="noopener">Attention Is All You Need&lt;/a>, the Transformer has become the foundation for a wide range of complex tasks, including machine translation, text generation, speech recognition, and image processing.&lt;/p>
&lt;p>Within the Transformer, self-attention has drawn wide interest from both academia and industry thanks to its outstanding performance. Attention lets the model access every element of the sequence at each time step. The key idea is selectivity: determining which words matter most in a particular context. It enriches the input embeddings by incorporating information about the surrounding input context. In other words, self-attention lets a model weigh the importance of different elements in the input sequence and dynamically adjust how much each contributes to the output.&lt;/p>
&lt;p>In this article we implement self-attention by hand, following &lt;a href="https://magazine.sebastianraschka.com/p/understanding-and-coding-self-attention" target="_blank" rel="noopener">Understanding and Coding Self-Attention, Multi-Head Attention, Cross-Attention, and Causal-Attention in LLMs&lt;/a>. A Chinese translation of that article was published by Synced as &lt;a href="https://mp.weixin.qq.com/s/XZDG8ZaB4QqD5cgIDJT6vQ" target="_blank" rel="noopener">&amp;ldquo;Still Don&amp;rsquo;t Understand Self-Attention in the Age of LLMs? This Article Walks You Through Implementing It from Scratch&amp;rdquo;&lt;/a>. Relative to those two pieces, this article adds my own reflections and experience.&lt;/p>
&lt;h1 id="embedding">Embedding&lt;/h1>
&lt;p>In this section we use embeddings to turn the discrete symbols of an input sequence (words or characters) into continuous, high-dimensional vector representations. Put simply, this is necessary because deep learning models cannot understand raw text directly; they must learn the semantics and syntax of text from the information carried by these vectors. For a deeper treatment, see the Zhihu article &lt;a href="https://zhuanlan.zhihu.com/p/164502624" target="_blank" rel="noopener">&amp;ldquo;Understanding Embeddings and How They Relate to Deep Learning&amp;rdquo;&lt;/a> (in Chinese).&lt;/p>
&lt;p>Given a sentence, we want to turn it into a continuous vector representation via embedding. We will use the sentence &amp;ldquo;Life is short, eat dessert first&amp;rdquo; as our running example.&lt;/p>
&lt;p>In the preprocessing stage, we deduplicate the words in the sentence and map each word to an integer index. This is easy to express in Python.&lt;/p>
&lt;p>Input:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-python" data-lang="python">&lt;span class="line">&lt;span class="cl">&lt;span class="n">sentence&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="s1">&amp;#39;Life is short, eat dessert first&amp;#39;&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">dc&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="p">{&lt;/span>&lt;span class="n">s&lt;/span>&lt;span class="p">:&lt;/span>&lt;span class="n">i&lt;/span> &lt;span class="k">for&lt;/span> &lt;span class="n">i&lt;/span>&lt;span class="p">,&lt;/span>&lt;span class="n">s&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="ow">in&lt;/span> &lt;span class="nb">enumerate&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="nb">sorted&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">sentence&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">replace&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="s1">&amp;#39;,&amp;#39;&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="s1">&amp;#39;&amp;#39;&lt;/span>&lt;span class="p">)&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">split&lt;/span>&lt;span class="p">()))}&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">dc&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Output:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-txt" data-lang="txt">&lt;span class="line">&lt;span class="cl">{&amp;#39;Life&amp;#39;: 0, &amp;#39;dessert&amp;#39;: 1, &amp;#39;eat&amp;#39;: 2, &amp;#39;first&amp;#39;: 3, &amp;#39;is&amp;#39;: 4, &amp;#39;short&amp;#39;: 5}
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>With this word-to-index mapping, the sentence can be represented as a sequence of indices.&lt;/p>
&lt;p>Input:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-python" data-lang="python">&lt;span class="line">&lt;span class="cl">&lt;span class="kn">import&lt;/span> &lt;span class="nn">torch&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">sentence_int&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">tensor&lt;/span>&lt;span class="p">(&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="p">[&lt;/span>&lt;span class="n">dc&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="n">s&lt;/span>&lt;span class="p">]&lt;/span> &lt;span class="k">for&lt;/span> &lt;span class="n">s&lt;/span> &lt;span class="ow">in&lt;/span> &lt;span class="n">sentence&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">replace&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="s1">&amp;#39;,&amp;#39;&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="s1">&amp;#39;&amp;#39;&lt;/span>&lt;span class="p">)&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">split&lt;/span>&lt;span class="p">()]&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">sentence_int&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Output:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-txt" data-lang="txt">&lt;span class="line">&lt;span class="cl">tensor([0, 4, 5, 2, 1, 3])
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Now we are ready for the key step, the embedding itself.&lt;/p>
&lt;p>In an embedding, each word is represented by a multi-dimensional vector whose dimensionality is determined by the size of the vocabulary. Llama 2, for example, uses an embedding size of 4096. To keep things compact, we use three dimensions in this article.&lt;/p>
&lt;p>Input:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-python" data-lang="python">&lt;span class="line">&lt;span class="cl">&lt;span class="n">vocab_size&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="mi">50_000&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">manual_seed&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="mi">123&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">embed&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">nn&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">Embedding&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">vocab_size&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="mi">3&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">embedded_sentence&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">embed&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">sentence_int&lt;/span>&lt;span class="p">)&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">detach&lt;/span>&lt;span class="p">()&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">embedded_sentence&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">embedded_sentence&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">shape&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Output:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-txt" data-lang="txt">&lt;span class="line">&lt;span class="cl">tensor([[ 0.3374, -0.1778, -0.3035],
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> [ 0.1794, 1.8951, 0.4954],
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> [ 0.2692, -0.0770, -1.0205],
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> [-0.2196, -0.3792, 0.7671],
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> [-0.5880, 0.3486, 0.6603],
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> [-1.1925, 0.6984, -1.4097]])
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">torch.Size([6, 3])
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>As the output shows, the six words of our sentence are represented as six row vectors, each of dimension 3. That is, every word in the sentence is represented by a vector of three numbers.&lt;/p>
&lt;h1 id="defining-the-weight-matrices">Defining the Weight Matrices&lt;/h1>
&lt;p>Starting with this section, we introduce the famous &amp;ldquo;QKV&amp;rdquo; mechanism of the Transformer.&lt;/p>
&lt;p>Self-attention uses three weight matrices, denoted $W_q$, $W_k$, and $W_v$; they are model parameters and are adjusted throughout training. Their role is to project the input into the query, key, and value components of the sequence.&lt;/p>
&lt;p>The corresponding query, key, and value sequences are obtained by matrix multiplication between the weight matrix $W$ and the embedded input $x$:&lt;/p>
&lt;ul>
&lt;li>Query sequence: for $i$ in the sequence $1\ldots T$, $q^{(i)}=x^{(i)}W_q$&lt;/li>
&lt;li>Key sequence: for $i$ in the sequence $1\ldots T$, $k^{(i)}=x^{(i)}W_k$&lt;/li>
&lt;li>Value sequence: for $i$ in the sequence $1\ldots T$, $v^{(i)}=x^{(i)}W_v$&lt;/li>
&lt;li>The index $i$ refers to the $token$ index position in the input sequence, whose length is $T$.&lt;/li>
&lt;/ul>
&lt;center> &lt;img style="border-radius: 0.3125em; box-shadow: 0 2px 4px 0 rgba(34,36,38,.12),0 2px 10px 0 rgba(34,36,38,.08);" src="https://geyuyao.com/post/transformer/640.png"> &lt;br> &lt;div style="color:orange; border-bottom: 1px solid #d9d9d9; display: inline-block; color: #999; padding: 1px;">&lt;/div>&lt;/center>
&lt;p>Here, both $q^{(i)}$ and $k^{(i)}$ are vectors of dimension $d_k$. The projection matrices $W_q$ and $W_k$ have shape $d × d_k$, while $W_v$ has shape $d × d_v$. Here $d$ denotes the number of dimensions of each word vector $x$, which is 3 in this article.&lt;/p>
&lt;p>Because we need to compute the dot product of the query and key vectors, these two vectors must have the same number of elements ($d_q=d_k$). Many LLMs also use value vectors of the same size, that is, $d_q=d_k=d_v$. However, the number of elements in the value vector $v^{(i)}$ can be arbitrary; it determines the size of the resulting context vector.&lt;/p>
&lt;p>In the code that follows, we set $d_q=d_k=2$ and $d_v=4$. The projection matrices are initialized as follows:&lt;/p>
&lt;p>Input:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-python" data-lang="python">&lt;span class="line">&lt;span class="cl">&lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">manual_seed&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="mi">123&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">d&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">embedded_sentence&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">shape&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="mi">1&lt;/span>&lt;span class="p">]&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">d_q&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_k&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_v&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="mi">2&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="mi">2&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="mi">4&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">W_query&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">nn&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">Parameter&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">rand&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">d&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_q&lt;/span>&lt;span class="p">))&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">W_key&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">nn&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">Parameter&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">rand&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">d&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_k&lt;/span>&lt;span class="p">))&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">W_value&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">nn&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">Parameter&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">rand&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">d&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_v&lt;/span>&lt;span class="p">))&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>In the original paper &lt;a href="https://arxiv.org/abs/1706.03762" target="_blank" rel="noopener">Attention Is All You Need&lt;/a>, $d_q$, $d_k$, and $d_v$ are typically set to &lt;code>64&lt;/code>, while the total model dimension is &lt;code>512&lt;/code>.&lt;/p>
&lt;h1 id="computing-unnormalized-attention-weights">Computing Unnormalized Attention Weights&lt;/h1>
&lt;p>In this section we walk through the computation using the second token as our example:&lt;/p>
&lt;center> &lt;img style="border-radius: 0.3125em; box-shadow: 0 2px 4px 0 rgba(34,36,38,.12),0 2px 10px 0 rgba(34,36,38,.08);" src="https://geyuyao.com/post/transformer/f9774001-ea9d-48bf-9857-3c911b0a279d_588x962.jpg"> &lt;br> &lt;div style="color:orange; border-bottom: 1px solid #d9d9d9; display: inline-block; color: #999; padding: 1px;">&lt;/div>&lt;/center>
&lt;p>As the figure shows, we multiply the input $x$ by $W_q$, $W_k$, and $W_v$ respectively.&lt;/p>
&lt;p>The code is as follows:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-python" data-lang="python">&lt;span class="line">&lt;span class="cl">&lt;span class="n">x_2&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">embedded_sentence&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="mi">1&lt;/span>&lt;span class="p">]&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">query_2&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">x_2&lt;/span> &lt;span class="o">@&lt;/span> &lt;span class="n">W_query&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">key_2&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">x_2&lt;/span> &lt;span class="o">@&lt;/span> &lt;span class="n">W_key&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">value_2&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">x_2&lt;/span> &lt;span class="o">@&lt;/span> &lt;span class="n">W_value&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">query_2&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">shape&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">key_2&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">shape&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">value_2&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">shape&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Output:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-txt" data-lang="txt">&lt;span class="line">&lt;span class="cl">torch.Size([2])
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">torch.Size([2])
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">torch.Size([4])
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Generalizing, we can multiply &lt;code>embedded_sentence&lt;/code> by $W_k$ and $W_v$ respectively; these results will be used in the following steps.&lt;/p>
&lt;p>Input:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-python" data-lang="python">&lt;span class="line">&lt;span class="cl">&lt;span class="n">keys&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">embedded_sentence&lt;/span> &lt;span class="o">@&lt;/span> &lt;span class="n">W_key&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">values&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">embedded_sentence&lt;/span> &lt;span class="o">@&lt;/span> &lt;span class="n">W_value&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="s2">&amp;#34;keys.shape:&amp;#34;&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">keys&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">shape&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="s2">&amp;#34;values.shape:&amp;#34;&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">values&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">shape&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Output:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-txt" data-lang="txt">&lt;span class="line">&lt;span class="cl">keys.shape: torch.Size([6, 2])
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">values.shape: torch.Size([6, 4])
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Now that we have $query^{(2)}$ along with all the &lt;code>keys&lt;/code> and &lt;code>values&lt;/code>, let us compute the unnormalized attention weights $\omega$, as shown below:&lt;/p>
&lt;center> &lt;img style="border-radius: 0.3125em; box-shadow: 0 2px 4px 0 rgba(34,36,38,.12),0 2px 10px 0 rgba(34,36,38,.08);" src="https://geyuyao.com/post/transformer/baf9e308-223b-429e-8527-a7b868003e8c_814x912.jpg"> &lt;br> &lt;div style="color:orange; border-bottom: 1px solid #d9d9d9; display: inline-block; color: #999; padding: 1px;">&lt;/div>&lt;/center>
&lt;p>As the figure shows, $\omega (i,j)$ is the dot product between the query and key sequences: $\omega (i,j) = q^{(i)}k^{(j)}$.&lt;/p>
&lt;p>For example, we can compute the unnormalized attention between the 2nd token&amp;rsquo;s query and the 5th token as follows:&lt;/p>
&lt;p>Input:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-python" data-lang="python">&lt;span class="line">&lt;span class="cl">&lt;span class="n">omega_24&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">query_2&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">dot&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">keys&lt;/span>&lt;span class="p">[&lt;/span>&lt;span class="mi">4&lt;/span>&lt;span class="p">])&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">omega_24&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Output:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-txt" data-lang="txt">&lt;span class="line">&lt;span class="cl">tensor(1.2903)
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Generalizing, we can multiply &lt;code>query_2&lt;/code> by &lt;code>keys&lt;/code> to obtain the unnormalized attention scores between the 2nd token and every other token.&lt;/p>
&lt;p>Input:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-python" data-lang="python">&lt;span class="line">&lt;span class="cl">&lt;span class="n">omega_2&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">query_2&lt;/span> &lt;span class="o">@&lt;/span> &lt;span class="n">keys&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">T&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">omega_2&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Output:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-txt" data-lang="txt">&lt;span class="line">&lt;span class="cl">tensor([-0.6004, 3.4707, -1.5023, 0.4991, 1.2903, -1.3374])
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;h1 id="computing-the-attention-weights">Computing the Attention Weights&lt;/h1>
&lt;p>In the previous section we computed the unnormalized attention scores between the 2nd token and every other token. In practice, we also need to normalize these scores.&lt;/p>
&lt;p>The purpose of normalization is to put the attention weight at each sequence position between 0 and 1, with all positions summing to 1. This probabilistic interpretation lets the model express, in probabilistic terms, how much attention it pays to different parts of the sequence: a higher weight means the model attends more to that position.&lt;/p>
&lt;center> &lt;img style="border-radius: 0.3125em; box-shadow: 0 2px 4px 0 rgba(34,36,38,.12),0 2px 10px 0 rgba(34,36,38,.08);" src="https://geyuyao.com/post/transformer/2222.png"> &lt;br> &lt;div style="color:orange; border-bottom: 1px solid #d9d9d9; display: inline-block; color: #999; padding: 1px;">&lt;/div>&lt;/center>
&lt;p>As the figure shows, self-attention first scales $\omega$ by $1/√{d_k} $ and then normalizes it with the softmax function.&lt;/p>
&lt;p>Scaling by $d_k$ ensures that the Euclidean lengths of the weight vectors are all roughly on the same scale. This helps prevent the attention weights from becoming too small or too large — which could cause numerical instability or hurt the model&amp;rsquo;s ability to converge during training.&lt;/p>
&lt;p>We can implement the attention weight computation in code like this:&lt;/p>
&lt;p>Input:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-python" data-lang="python">&lt;span class="line">&lt;span class="cl">&lt;span class="kn">import&lt;/span> &lt;span class="nn">torch.nn.functional&lt;/span> &lt;span class="k">as&lt;/span> &lt;span class="nn">F&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">attention_weights_2&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">F&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">softmax&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">omega_2&lt;/span> &lt;span class="o">/&lt;/span> &lt;span class="n">d_k&lt;/span>&lt;span class="o">**&lt;/span>&lt;span class="mf">0.5&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">dim&lt;/span>&lt;span class="o">=&lt;/span>&lt;span class="mi">0&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">attention_weights_2&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Output:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-txt" data-lang="txt">&lt;span class="line">&lt;span class="cl">tensor([0.0386, 0.6870, 0.0204, 0.0840, 0.1470, 0.0229])
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>The final step is to compute the context vector $z^{(2)}$ — an attention-weighted version of the original query input $x^{(2)}$ that incorporates all the other input elements as context through the attention weights:&lt;/p>
&lt;center> &lt;img style="border-radius: 0.3125em; box-shadow: 0 2px 4px 0 rgba(34,36,38,.12),0 2px 10px 0 rgba(34,36,38,.08);" src="https://geyuyao.com/post/transformer/9999999999999999999999.png"> &lt;br> &lt;div style="color:orange; border-bottom: 1px solid #d9d9d9; display: inline-block; color: #999; padding: 1px;">&lt;/div>&lt;/center>
&lt;p>Input:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-python" data-lang="python">&lt;span class="line">&lt;span class="cl">&lt;span class="n">context_vector_2&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">attention_weights_2&lt;/span> &lt;span class="o">@&lt;/span> &lt;span class="n">values&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">context_vector_2&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">shape&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">context_vector_2&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Output:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-txt" data-lang="txt">&lt;span class="line">&lt;span class="cl">torch.Size([4])
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">tensor([0.5313, 1.3607, 0.7891, 1.3110])
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Note that this output vector has more dimensions ($d_v=4$) than the input vector ($d=3$), because we set $d_v &amp;gt; d$ earlier. The embedding size $d_v$, however, can be chosen arbitrarily.&lt;/p>
&lt;h1 id="self-attention">Self-Attention&lt;/h1>
&lt;p>Now let us pull together the code for the self-attention mechanism from the previous sections.&lt;/p>
&lt;p>We can condense everything above into a compact Self-Attention class:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-python" data-lang="python">&lt;span class="line">&lt;span class="cl">&lt;span class="kn">import&lt;/span> &lt;span class="nn">torch.nn&lt;/span> &lt;span class="k">as&lt;/span> &lt;span class="nn">nn&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="k">class&lt;/span> &lt;span class="nc">SelfAttention&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">nn&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">Module&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">def&lt;/span> &lt;span class="fm">__init__&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="bp">self&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_in&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_out_kq&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_out_v&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="nb">super&lt;/span>&lt;span class="p">()&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="fm">__init__&lt;/span>&lt;span class="p">()&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">d_out_kq&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">d_out_kq&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">W_query&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">nn&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">Parameter&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">rand&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">d_in&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_out_kq&lt;/span>&lt;span class="p">))&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">W_key&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">nn&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">Parameter&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">rand&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">d_in&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_out_kq&lt;/span>&lt;span class="p">))&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">W_value&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">nn&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">Parameter&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">rand&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">d_in&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_out_v&lt;/span>&lt;span class="p">))&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">def&lt;/span> &lt;span class="nf">forward&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="bp">self&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">x&lt;/span>&lt;span class="p">):&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">keys&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">x&lt;/span> &lt;span class="o">@&lt;/span> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">W_key&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">queries&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">x&lt;/span> &lt;span class="o">@&lt;/span> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">W_query&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">values&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">x&lt;/span> &lt;span class="o">@&lt;/span> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">W_value&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">attn_scores&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">queries&lt;/span> &lt;span class="o">@&lt;/span> &lt;span class="n">keys&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">T&lt;/span> &lt;span class="c1"># unnormalized attention weights &lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">attn_weights&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">softmax&lt;/span>&lt;span class="p">(&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">attn_scores&lt;/span> &lt;span class="o">/&lt;/span> &lt;span class="bp">self&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">d_out_kq&lt;/span>&lt;span class="o">**&lt;/span>&lt;span class="mf">0.5&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">dim&lt;/span>&lt;span class="o">=-&lt;/span>&lt;span class="mi">1&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="n">context_vec&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">attn_weights&lt;/span> &lt;span class="o">@&lt;/span> &lt;span class="n">values&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> &lt;span class="k">return&lt;/span> &lt;span class="n">context_vec&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Input:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-python" data-lang="python">&lt;span class="line">&lt;span class="cl">&lt;span class="n">torch&lt;/span>&lt;span class="o">.&lt;/span>&lt;span class="n">manual_seed&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="mi">123&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="c1"># reduce d_out_v from 4 to 1, because we have 4 heads&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">d_in&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_out_kq&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_out_v&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="mi">3&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="mi">2&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="mi">4&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="n">sa&lt;/span> &lt;span class="o">=&lt;/span> &lt;span class="n">SelfAttention&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">d_in&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_out_kq&lt;/span>&lt;span class="p">,&lt;/span> &lt;span class="n">d_out_v&lt;/span>&lt;span class="p">)&lt;/span>
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl">&lt;span class="nb">print&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">sa&lt;/span>&lt;span class="p">(&lt;/span>&lt;span class="n">embedded_sentence&lt;/span>&lt;span class="p">))&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>Output:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" class="chroma">&lt;code class="language-txt" data-lang="txt">&lt;span class="line">&lt;span class="cl">tensor([[-0.1564, 0.1028, -0.0763, -0.0764],
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> [ 0.5313, 1.3607, 0.7891, 1.3110],
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> [-0.3542, -0.1234, -0.2627, -0.3706],
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> [ 0.0071, 0.3345, 0.0969, 0.1998],
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> [ 0.1008, 0.4780, 0.2021, 0.3674],
&lt;/span>&lt;/span>&lt;span class="line">&lt;span class="cl"> [-0.5296, -0.2799, -0.4107, -0.6006]], grad_fn=&amp;lt;MmBackward0&amp;gt;)
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>The second row is exactly the value of &lt;code>context_vector_2&lt;/code> from the previous section: &lt;code>tensor([0.5313, 1.3607, 0.7891, 1.3110])&lt;/code>.&lt;/p>
&lt;h1 id="conclusion">Conclusion&lt;/h1>
&lt;p>In this article we implemented self-attention by hand, with code and explanations, following &lt;a href="https://magazine.sebastianraschka.com/p/understanding-and-coding-self-attention" target="_blank" rel="noopener">Understanding and Coding Self-Attention, Multi-Head Attention, Cross-Attention, and Causal-Attention in LLMs&lt;/a>.&lt;/p></description></item></channel></rss>