diff --git a/chapter1_transformer_interp/exercises/part1_transformer_from_scratch/1.1_Transformer_from_Scratch_exercises.ipynb b/chapter1_transformer_interp/exercises/part1_transformer_from_scratch/1.1_Transformer_from_Scratch_exercises.ipynb new file mode 100644 index 0000000..eeb9266 --- /dev/null +++ b/chapter1_transformer_interp/exercises/part1_transformer_from_scratch/1.1_Transformer_from_Scratch_exercises.ipynb @@ -0,0 +1,6031 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": { + "id": "4XPVLiZRfdfI" + }, + "source": [ + "# [1.1] - Transformers from scratch (exercises)\n", + "\n", + "> **ARENA [Streamlit Page](https://arena-chapter1-transformer-interp.streamlit.app/01_[1.1]_Transformer_from_Scratch)**\n", + ">\n", + "> **Colab: [exercises](https://colab.research.google.com/github/ARENA-education/ARENA_materials/blob/main/chapter1_transformer_interp/exercises/part1_transformer_from_scratch/1.1_Transformer_from_Scratch_exercises.ipynb?t=20260901) | [solutions](https://colab.research.google.com/github/ARENA-education/ARENA_materials/blob/main/chapter1_transformer_interp/exercises/part1_transformer_from_scratch/1.1_Transformer_from_Scratch_solutions.ipynb?t=20260901)**\n", + "\n", + "Please send any problems / bugs on the `#errata` channel in the [Slack group](https://info-arena.github.io/ARENA_img/slack.html), and ask any questions on the dedicated channels for this chapter of material.\n", + "\n", + "You can collapse each section so only the headers are visible, by clicking the arrow symbol on the left hand side of the markdown header cells.\n", + "\n", + "Links to all other chapters: [(0) Fundamentals](https://arena-chapter0-fundamentals.streamlit.app/), [(1) Transformer Interpretability](https://arena-chapter1-transformer-interp.streamlit.app/), [(2) RL](https://arena-chapter2-rl.streamlit.app/)." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "h2x73RisfdfJ" + }, + "source": [ + "" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "82Jhq7NEfdfJ" + }, + "source": [ + "# Introduction" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "D5xkxsvBfdfK" + }, + "source": [ + "This is a clean, first principles implementation of GPT-2 in PyTorch. The architectural choices closely follow those used by the TransformerLens library (which you'll be using a lot more in later exercises).\n", + "\n", + "The exercises are written to accompany Neel Nanda's [TransformerLens library](https://github.com/neelnanda-io/TransformerLens) for doing mechanistic interpretability research on GPT-2 style language models. We'll be working with this library extensively in this chapter of the course.\n", + "\n", + "Each exercise will have a difficulty and importance rating out of 5, as well as an estimated maximum time you should spend on these exercises and sometimes a short annotation. You should interpret the ratings & time estimates relatively (e.g. if you find yourself spending about 50% longer on the exercises than the time estimates, adjust accordingly). Please do skip exercises / look at solutions if you don't feel like they're important enough to be worth doing, and you'd rather get to the good stuff!\n", + "\n", + "For a lecture on the material today, which provides some high-level understanding before you dive into the material, watch the video below:\n", + "\n", + "\n", + "\n", + "\n", + "This content was based on a series of videos by [Neel Nanda](https://www.youtube.com/playlist?list=PL7m7hLIqA0hoIUPhC26ASCVs_VrqcDpAz)." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "ViabfiyPfdfK" + }, + "source": [ + "## Content & Learning Objectives\n", + "\n", + "### 1️⃣ Understanding Inputs & Outputs of a Transformer\n", + "\n", + "In this section, we'll take a first look at transformers - what their function is, how information moves inside a transformer, and what inputs & outputs they take.\n", + "\n", + "> ##### Learning Objectives\n", + ">\n", + "> - Understand what a transformer is used for\n", + "> - Understand causal attention, and what a transformer's output represents - algebraic operations on tensors\n", + "> - Learn what tokenization is, and how models do it\n", + "> - Understand what logits are, and how to use them to derive a probability distribution over the vocabulary\n", + "\n", + "### 2️⃣ Clean Transformer Implementation\n", + "\n", + "Here, we'll implement a transformer from scratch, using only PyTorch's tensor operations. This will give us a good understanding of how transformers work, and how to use them. We do this by going module-by-module, in an experience which should feel somewhat similar to last week's ResNet exercises. Much like with ResNets, you'll conclude by loading in pretrained weights and verifying that your model works as expected.\n", + "\n", + "> ##### Learning Objectives\n", + ">\n", + "> * Understand that a transformer is composed of attention heads and MLPs, with each one performing operations on the residual stream\n", + "> * Understand that the attention heads in a single layer operate independently, and that they have the role of calculating attention patterns (which determine where information is moved to & from in the residual stream)\n", + "> * Learn about & implement the following transformer modules:\n", + "> * LayerNorm (transforming the input to have zero mean and unit variance)\n", + "> * Positional embedding (a lookup table from position indices to residual stream vectors)\n", + "> * Attention (the method of computing attention patterns for residual stream vectors)\n", + "> * MLP (the collection of linear and nonlinear transformations which operate on each residual stream vector in the same way)\n", + "> * Embedding (a lookup table from tokens to residual stream vectors)\n", + "> * Unembedding (a matrix for converting residual stream vectors into a distribution over tokens)\n", + "\n", + "### 3️⃣ Training a Transformer\n", + "\n", + "Next, you'll learn how to train your transformer from scratch. This will be quite similar to the training loops you wrote for ResNet in your first week.\n", + "\n", + "> ##### Learning Objectives\n", + ">\n", + "> * Understand how to train a transformer from scratch\n", + "> * Write a basic transformer training loop\n", + "> * Interpret the transformer's falling cross entropy loss with reference to features of the training data (e.g. bigram frequencies)\n", + "\n", + "### 4️⃣ Sampling from a Transformer\n", + "\n", + "Lastly, you'll learn how to sample from a transformer. This will involve implementing a few different sampling methods, and writing a caching system which can reuse computations from previous forward passes to improve your model's text generation speed.\n", + "\n", + "*The second half of this section is less important, and you can skip it if you want.*\n", + "\n", + "> ##### Learning Objectives\n", + ">\n", + "> * Learn how to sample from a transformer\n", + "> * This includes basic methods like greedy search or top-k, and more advanced methods like beam search\n", + "> * Learn how to cache the output of a transformer, so that it can be used to generate text more efficiently\n", + "> * Optionally, rewrite your sampling functions to make use of your caching methods" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "3ZB6VKWffdfK" + }, + "source": [ + "## Setup code" + ] + }, + { + "cell_type": "code", + "execution_count": 155, + "metadata": { + "id": "JwXpbcKHfdfK" + }, + "outputs": [], + "source": [ + "import os\n", + "import sys\n", + "from pathlib import Path\n", + "\n", + "IN_COLAB = \"google.colab\" in sys.modules\n", + "\n", + "chapter = \"chapter1_transformer_interp\"\n", + "repo = \"ARENA_materials\"\n", + "branch = \"main\"\n", + "\n", + "# Install dependencies\n", + "try:\n", + " import transformer_lens\n", + "except:\n", + " %pip install transformer_lens einops jaxtyping git+https://github.com/callummcdougall/CircuitsVis.git#subdirectory=python\n", + "\n", + "# Get root directory, handling 3 different cases: (1) Colab, (2) notebook not in ARENA repo, (3) notebook in ARENA repo\n", + "root = (\n", + " \"/content\"\n", + " if IN_COLAB\n", + " else \"/root\"\n", + " if repo not in os.getcwd()\n", + " else str(next(p for p in Path.cwd().parents if p.name == repo))\n", + ")\n", + "\n", + "if Path(root).exists() and not Path(f\"{root}/{chapter}\").exists():\n", + " if not IN_COLAB:\n", + " !sudo apt-get install unzip\n", + " %pip install jupyter ipython --upgrade\n", + "\n", + " if not os.path.exists(f\"{root}/{chapter}\"):\n", + " !wget -P {root} https://github.com/ARENA-education/ARENA_materials/archive/refs/heads/{branch}.zip\n", + " !unzip {root}/{branch}.zip '{repo}-{branch}/{chapter}/exercises/*' -d {root}\n", + " !mv {root}/{repo}-{branch}/{chapter} {root}/{chapter}\n", + " !rm {root}/{branch}.zip\n", + " !rmdir {root}/ARENA_materials-{branch}\n", + "\n", + "\n", + "if f\"{root}/{chapter}/exercises\" not in sys.path:\n", + " sys.path.append(f\"{root}/{chapter}/exercises\")\n", + "\n", + "os.chdir(f\"{root}/{chapter}/exercises\")" + ] + }, + { + "cell_type": "code", + "execution_count": 156, + "metadata": { + "id": "-9mDUJqufdfL" + }, + "outputs": [], + "source": [ + "import math\n", + "import os\n", + "import sys\n", + "from collections import defaultdict\n", + "from dataclasses import dataclass\n", + "from pathlib import Path\n", + "from typing import Callable\n", + "\n", + "import datasets\n", + "import einops\n", + "import numpy as np\n", + "import torch as t\n", + "import torch.nn as nn\n", + "import wandb\n", + "from jaxtyping import Float, Int\n", + "from rich import print as rprint\n", + "from rich.table import Table\n", + "from torch import Tensor\n", + "from torch.utils.data import DataLoader\n", + "from tqdm.auto import tqdm\n", + "from transformer_lens import HookedTransformer\n", + "from transformer_lens.utils import gelu_new, tokenize_and_concatenate\n", + "from transformers import GPT2TokenizerFast\n", + "\n", + "device = t.device(\"mps\" if t.backends.mps.is_available() else \"cuda\" if t.cuda.is_available() else \"cpu\")\n", + "\n", + "# Make sure exercises are in the path\n", + "chapter = \"chapter1_transformer_interp\"\n", + "section = \"part1_transformer_from_scratch\"\n", + "root_dir = next(p for p in Path.cwd().parents if (p / chapter).exists())\n", + "exercises_dir = root_dir / chapter / \"exercises\"\n", + "section_dir = exercises_dir / section\n", + "\n", + "import part1_transformer_from_scratch.solutions as solutions\n", + "import part1_transformer_from_scratch.tests as tests\n", + "\n", + "MAIN = __name__ == \"__main__\"" + ] + }, + { + "cell_type": "code", + "source": [], + "metadata": { + "id": "ubVAl8cJWV0r" + }, + "execution_count": 156, + "outputs": [] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "qQq2tlZhfdfL" + }, + "source": [ + "# 1️⃣ Understanding Inputs & Outputs of a Transformer\n", + "\n", + "> ##### Learning Objectives\n", + ">\n", + "> - Understand what a transformer is used for\n", + "> - Understand causal attention, and what a transformer's output represents - algebraic operations on tensors\n", + "> - Learn what tokenization is, and how models do it\n", + "> - Understand what logits are, and how to use them to derive a probability distribution over the vocabulary" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "SYneN-V5fdfL" + }, + "source": [ + "## What is the point of a transformer?" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "KxBlbEeSfdfL" + }, + "source": [ + "**Transformers exist to model text!**\n", + "\n", + "We're going to focus on GPT-2 style transformers. Key feature: They generate text! You feed in language, and the model generates a probability distribution over tokens. And you can repeatedly sample from this to generate text!\n", + "\n", + "(To explain this in more detail - you feed in a sequence of length $N$, then sample from the probability distribution over the $N+1$-th word, use this to construct a new sequence of length $N+1$, then feed this new sequence into the model to get a probability distribution over the $N+2$-th word, and so on.)\n", + "\n", + "### How is the model trained?\n", + "\n", + "You give it a bunch of text, and train it to predict the next token.\n", + "\n", + "Importantly, if you give a model 100 tokens in a sequence, it predicts the next token for *each* prefix, i.e. it produces 100 logit vectors (= probability distributions) over the set of all words in our vocabulary, with the `i`-th logit vector representing the probability distribution over the token *following* the `i`-th token in the sequence. This is a key part of what allows transformers to be trained so efficiently; for every sequence of length $n$ we get $n$ different predictions to train on:\n", + "\n", + "$$\n", + "p(x_1), \\; p(x_2|x_1), \\; p(x_3|x_1x_2), \\; \\ldots, \\; p(x_n|x_1 \\ldots x_{n-1})\n", + "$$\n", + "\n", + "
\n", + "Aside - logits\n", + "\n", + "If you haven't encountered the term \"logits\" before, here's a quick refresher.\n", + "\n", + "Given an arbitrary vector $x$, we can turn it into a probability distribution via the **softmax** function: $x_i \\to \\frac{e^{x_i}}{\\sum e^{x_j}}$. The exponential makes everything positive; the normalization makes it add to one.\n", + "\n", + "The model's output is the vector $x$ (one for each prediction it makes). We call this vector a logit because it represents a probability distribution, and it is related to the actual probabilities via the softmax function.\n", + "
\n", + "\n", + "How do we stop the transformer from \"cheating\" by just looking at the tokens it's trying to predict? Answer - we make the transformer have *causal attention* (as opposed to *bidirectional attention*). Causal attention only allows information to move forwards in the sequence, never backwards. The prediction of what comes after token 50 is only a function of the first 50 tokens, *not* of token 51. We say the transformer is **causal**, because each next token is generated conditioned on the current context, which includes the original input plus the tokens generated so far." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "JCfrJmRvfdfL" + }, + "source": [ + "" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "zynVn1tifdfL" + }, + "source": [ + "Another way to view this is through the following analogy: we have a series of people standing in a line, each with one word or chunk of the sentence. Each person has the ability to look up information from the people behind them (we'll explore how this works in later sections) but they can't look at any information in front of them. Their goal is to guess what word the person in front of them is holding.\n", + "\n", + "" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "BGIFaz92fdfL" + }, + "source": [ + "## Tokens - Transformer Inputs" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "YD0qahHAfdfL" + }, + "source": [ + "Our transformer's input is natural language (i.e. a sequence of characters, strings, etc). But ML models generally take vectors as input, not language. How do we convert language to vectors?\n", + "\n", + "We can factor this into 2 questions:\n", + "\n", + "1. How do we split up language into small sub-units?\n", + "2. How do we convert these sub-units into vectors?\n", + "\n", + "Let's start with the second of these questions." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "55V8K_vPfdfL" + }, + "source": [ + "### Converting sub-units to vectors\n", + "\n", + "We basically make a massive lookup table, which is called an **embedding**. It has one vector for each possible sub-unit of language we might get (we call this set of all sub-units our **vocabulary**). We label every element in our vocabulary with an integer (this labelling never changes), and we use this integer to index into the embedding.\n", + "\n", + "A key intuition is that one-hot encodings let you think about each integer independently. We don't bake in any relation between words when we perform our embedding, because every word has a completely separate embedding vector.\n", + "\n", + "
\n", + "Aside - one-hot encodings\n", + "\n", + "We sometimes think about **one-hot encodings** of words. These are vectors with zeros everywhere, except for a single one in the position corresponding to the word's index in the vocabulary. This means that indexing into the embedding is equivalent to multiplying the **embedding matrix** by the one-hot encoding (where the embedding matrix is the matrix we get by stacking all the embedding vectors on top of each other).\n", + "\n", + "$$\n", + "\\begin{aligned}\n", + "W_E &= \\begin{bmatrix}\n", + "\\leftarrow v_0 \\rightarrow \\\\\n", + "\\leftarrow v_1 \\rightarrow \\\\\n", + "\\vdots \\\\\n", + "\\leftarrow v_{d_{vocab}-1} \\rightarrow \\\\\n", + "\\end{bmatrix} \\quad \\text{is the embedding matrix (size }d_{vocab} \\times d_{embed}\\text{),} \\\\\n", + "\\\\\n", + "t_i &= (0, \\dots, 0, 1, 0, \\dots, 0) \\quad \\text{is the one-hot encoding for the }i\\text{th word (length }d_{vocab}\\text{)} \\\\\n", + "\\\\\n", + "v_i &= t_i W_E \\quad \\text{is the embedding vector for the }i\\text{th word (length }d_{embed}\\text{).} \\\\\n", + "\\end{aligned}\n", + "$$\n", + "\n", + "
\n", + "\n", + "Now, let's answer the first question - how do we split language into sub-units?" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "wQFqxeO7fdfL" + }, + "source": [ + "### Splitting language into sub-units\n", + "\n", + "We need to define a standard way of splitting up language into a series of substrings, where each substring is a member of our **vocabulary** set.\n", + "\n", + "Could we use a dictionary, and have our vocabulary be the set of all words in the dictionary? No, because this couldn't handle arbitrary text (e.g. URLs, punctuation, etc). We need a more general way of splitting up language.\n", + "\n", + "Could we just use the 256 ASCII characters? This fixes the previous problem, but it loses the structure of language - some sequences of characters are more meaningful than others. For example, \"language\" is a lot more meaningful than \"hjksdfiu\". We want \"language\" to be a single token, but not \"hjksdfiu\" - this is a more efficient use of our vocab.\n", + "\n", + "What actually happens? The most common strategy is called **Byte-Pair encodings**.\n", + "\n", + "We begin with the 256 ASCII characters as our tokens, and then find the most common pair of tokens, and merge that into a new token. Note that we do have a space character as one of our 256 tokens, and merges using space are very common. For instance, here are the first five merges for the tokenizer used by GPT-2 (you'll be able to verify this below).\n", + "\n", + "```\n", + "\" t\"\n", + "\" a\"\n", + "\"he\"\n", + "\"in\"\n", + "\"re\"\n", + "```\n", + "\n", + "Note - you might see the character `Ġ` in front of some tokens. This is a special token that indicates that the token begins with a space. Tokens with a leading space vs not are different.\n", + "\n", + "You can run the code below to load in the `gpt2-small` model, and see more of its tokenizer's vocabulary:" + ] + }, + { + "cell_type": "code", + "execution_count": 157, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/", + "height": 191, + "referenced_widgets": [ + "a76388bdb84c422f833a0015ba4a7ab5", + "de3deb7c5ee2471186c2d3d354ef8014", + "5faad48740fd49acb8869dfb35bbc9fb", + "263b3cbc12c746df95145835ba3e39c2", + "3b3c2c97d7f54f1c9c7bf92d8157c1c6", + "60c4eb83080d4333835df32ae061af5e", + "53891efb605249ad83402d96a3c67dae", + "a90b519263d844a883816df8a9c8d7cb", + "c64111af122341fbace7fb2195d0b519", + "f7667933820842af82625dfa89b59748", + "0ee53673a92b4d0a8398e9128c81c988" + ] + }, + "id": "rbZ3QvpgfdfL", + "outputId": "1d238651-0afb-46f8-82f1-3dc87a194a51" + }, + "outputs": [ + { + "output_type": "display_data", + "data": { + "text/plain": [ + "Loading weights: 0%| | 0/148 [00:00', 50256)]\n" + ] + } + ], + "source": [ + "print(sorted_vocab[-20:])" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "-Cw-TryXfdfL" + }, + "source": [ + "
\n", + "Fun (completely optional) exercise - can you guess what the first-formed 3/4/5/6/7-letter encodings in GPT-2's vocabulary are?\n", + "Run this code to find out:\n", + "\n", + "```python\n", + "lengths = dict.fromkeys(range(3, 8), \"\")\n", + "for tok, idx in sorted_vocab:\n", + " if not lengths.get(len(tok), True):\n", + " lengths[len(tok)] = tok\n", + "\n", + "for length, tok in lengths.items():\n", + " print(f\"{length}: {tok}\")\n", + "```\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "axKS2cFEfdfM" + }, + "source": [ + "Transformers in the `transformer_lens` library have a `to_tokens` method that converts text to numbers. It also prepends them with a special token called BOS (beginning of sequence) to indicate the start of a sequence. You can disable this with the `prepend_bos=False` argument.\n", + "\n", + "
\n", + "Aside - BOS token\n", + "\n", + "The beginning of sequence (BOS) token is a special token used to mark the beginning of the sequence. Confusingly, in GPT-2, the End of Sequence (EOS), Beginning of Sequence (BOS) and Padding (PAD) tokens are all the same, `<|endoftext|>` with index `50256`.\n", + "\n", + "Why is this token added? Some basic intuitions are:\n", + "\n", + "* It provides context that this is the start of a sequence, which can help the model generate more appropriate text.\n", + "* It can act as a \"rest position\" for attention heads (more on this later, when we discuss attention).\n", + "\n", + "TransformerLens adds this token automatically (including in forward passes of transformer models, e.g. it's implicitly added when you call `model(\"Hello World\")`). You can disable this behaviour by setting the flag `prepend_bos=False` in `to_tokens`, `to_str_tokens`, `model.forward` and any other function that converts strings to multi-token tensors.\n", + "\n", + "**Key Point: *If you get weird off-by-one errors, check whether there's an unexpected `prepend_bos`!***\n", + "\n", + "Why are the BOS, EOS and PAD tokens the same? This is because GPT-2 is an autoregressive model, and uses these kinds of tokens in a slightly different way to other transformer families (e.g. BERT). For instance, GPT has no need to distinguish between BOS and EOS tokens, because it only processes text from left to right.\n", + "\n", + "
\n", + "\n", + "### Some tokenization annoyances\n", + "\n", + "There are a few funky and frustrating things about tokenization, which cause it to behave differently than you might expect. For instance:\n", + "\n", + "#### Whether a word begins with a capital or space matters!" + ] + }, + { + "cell_type": "code", + "execution_count": 159, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "iiLjnZfcfdfM", + "outputId": "ddee3567-5315-4052-f6c7-a37f6ee59676" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "['<|endoftext|>Ralph']\n", + "['<|endoftext|> Ralph']\n", + "['<|endoftext|> ralph']\n", + "['<|endoftext|>ralph']\n" + ] + } + ], + "source": [ + "print(reference_gpt2.to_str_tokens(\"Ralph\"))\n", + "print(reference_gpt2.to_str_tokens(\" Ralph\"))\n", + "print(reference_gpt2.to_str_tokens(\" ralph\"))\n", + "print(reference_gpt2.to_str_tokens(\"ralph\"))" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "AqwY3oKnfdfM" + }, + "source": [ + "#### Arithmetic is a mess.\n", + "\n", + "Length is inconsistent, common numbers bundle together." + ] + }, + { + "cell_type": "code", + "execution_count": 160, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "qGijJiCYfdfM", + "outputId": "9dc21f38-a09d-47c0-9829-e13e80c1ebfe" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "['<|endoftext|>56873+3184623=123456789-1000000000']\n" + ] + } + ], + "source": [ + "print(reference_gpt2.to_str_tokens(\"56873+3184623=123456789-1000000000\"))" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "XVgTZ-B_fdfM" + }, + "source": [ + "> ### Key Takeaways\n", + ">\n", + "> * We learn a dictionary of vocab of tokens (sub-words).\n", + "> * We (approx) losslessly convert language to integers via tokenizing it.\n", + "> * We convert integers to vectors via a lookup table.\n", + "> * Note: input to the transformer is a sequence of *tokens* (ie integers), not vectors" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "ljDWq71TfdfM" + }, + "source": [ + "## Text generation\n", + "\n", + "Now that we understand the basic ideas here, let's go through the entire process of text generation, from our original string to a new token which we can append to our string and plug back into the model.\n", + "\n", + "#### **Step 1:** Convert text to tokens\n", + "\n", + "The sequence gets tokenized, so it has shape `[batch, seq_len]`. Here, the batch dimension is just one (because we only have one sequence)." + ] + }, + { + "cell_type": "code", + "execution_count": 161, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "LSU2jtucfdfM", + "outputId": "4a1d17a2-63e3-45a7-fe23-60414f320dcc" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "tensor([[50256, 40, 716, 281, 4998, 1960, 382, 19741, 11, 875,\n", + " 12342, 12, 8807, 11, 402, 11571, 12, 17, 3918, 47385,\n", + " 13, 1881, 1110, 314, 481, 7074, 1692, 1241, 4430, 290,\n", + " 1011, 625, 262, 995, 0]])\n", + "torch.Size([1, 35])\n", + "['<|endoftext|>I am an amazing autoregressive, decoder-only, GPT-2 style transformer. One day I will exceed human level intelligence and take over the world!']\n" + ] + } + ], + "source": [ + "reference_text = \"I am an amazing autoregressive, decoder-only, GPT-2 style transformer. One day I will exceed human level intelligence and take over the world!\"\n", + "tokens = reference_gpt2.to_tokens(reference_text).to(device)\n", + "print(tokens)\n", + "print(tokens.shape)\n", + "print(reference_gpt2.to_str_tokens(tokens))" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "5O_FRICWfdfM" + }, + "source": [ + "#### **Step 2:** Map tokens to logits\n", + "\n", + "\n", + "From our input of shape `[batch, seq_len]`, we get output of shape `[batch, seq_len, vocab_size]`. The `[i, j, :]`-th element of our output is a vector of logits representing our prediction for the `j+1`-th token in the `i`-th sequence." + ] + }, + { + "cell_type": "code", + "execution_count": 162, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "wrivQkTpfdfM", + "outputId": "6da079a6-b92a-47ae-ed49-73da22ae8b40" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "torch.Size([1, 35, 50257])\n" + ] + } + ], + "source": [ + "logits, cache = reference_gpt2.run_with_cache(tokens)\n", + "print(logits.shape)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "NH_yTB02fdfM" + }, + "source": [ + "(`run_with_cache` tells the model to cache all intermediate activations. This isn't important right now; we'll look at it in more detail later.)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "0pTuJt-DfdfM" + }, + "source": [ + "#### **Step 3:** Convert the logits to a distribution with a softmax\n", + "\n", + "This doesn't change the shape, it is still `[batch, seq_len, vocab_size]`." + ] + }, + { + "cell_type": "code", + "execution_count": 163, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "y0f0ywOLfdfM", + "outputId": "7607b1fd-3a1f-4f41-ff46-2c095d8c4036" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "torch.Size([1, 35, 50257])\n" + ] + } + ], + "source": [ + "probs = logits.softmax(dim=-1)\n", + "print(probs.shape)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "lSY7hGpnfdfM" + }, + "source": [ + "#### **Bonus step:** What is the most likely next token at each position?" + ] + }, + { + "cell_type": "code", + "execution_count": 164, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "cixrFcjhfdfM", + "outputId": "f5f0b3a2-a41c-4c6c-e97e-f26f1c14f578" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "<|endoftext|>I am an amazing autoregressive, decoder-only, GPT-2 style transformer. One day I will exceed human level intelligence and take over the world! -> \n", + "'m a avid personodsp. andently,driven programmer andIM-only.,. I of I will be myly of and I over the world. I\n" + ] + } + ], + "source": [ + "most_likely_next_tokens = reference_gpt2.to_str_tokens(logits.topk(k=4, dim=-1).indices[0, :, 0])\n", + "\n", + "for token, next_token in zip(\n", + " reference_gpt2.to_str_tokens(tokens),\n", + " most_likely_next_tokens\n", + "):\n", + " print(f\"{token} -> {next_token}\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "zPGZXlXsfdfM" + }, + "source": [ + "We can see that, in a few cases (particularly near the end of the sequence), the model accurately predicts the next token in the sequence. We might guess that `\"take over the world\"` is a common phrase that the model has seen in training, which is why the model can predict it." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "zH7ofknkfdfM" + }, + "source": [ + "#### **Step 4:** Map distribution to a token" + ] + }, + { + "cell_type": "code", + "execution_count": 165, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "EXKkLMqSfdfM", + "outputId": "72778971-b087-48f3-c1de-339ea9f3934b" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "torch.Size([1, 35]) torch.Size([1, 35, 50257])\n", + "' I'\n", + "tensor([[50256, 40, 1842, 3797]])\n" + ] + } + ], + "source": [ + "print(tokens.shape, logits.shape)\n", + "next_token = logits[0, 34].argmax(dim=-1)\n", + "next_char = reference_gpt2.to_string(next_token)\n", + "print(repr(next_char))\n", + "\n", + "\n", + "\n", + "\n", + "tt = reference_gpt2.to_tokens(\"I love cat\")\n", + "print(tt)\n" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "7kT5wk5IfdfM" + }, + "source": [ + "Note that we're indexing `logits[0, -1]`. This is because logits have shape `[1, sequence_length, vocab_size]`, so this indexing returns the vector of length `vocab_size` representing the model's prediction for what token follows the **last** token in the input sequence.\n", + "\n", + "In this case, we can see that the model predicts the token `' I'`." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "tGtoSj8HfdfN" + }, + "source": [ + "### **Step 5:** Add this to the end of the input, re-run\n", + "\n", + "There are more efficient ways to do this (e.g. where we cache some of the values each time we run our input, so we don't have to do as much calculation each time we generate a new value), but this doesn't matter conceptually right now." + ] + }, + { + "cell_type": "code", + "execution_count": 166, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "hmrK0wwifdfN", + "outputId": "e515937c-09d4-40d2-bf4d-3c7fd53e4fd5" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "Sequence so far: '<|endoftext|>I am an amazing autoregressive, decoder-only, GPT-2 style transformer. One day I will exceed human level intelligence and take over the world!'\n", + "36th char = ' I'\n", + "37th char = ' am'\n", + "38th char = ' a'\n", + "39th char = ' very'\n", + "40th char = ' talented'\n", + "41th char = ' and'\n", + "42th char = ' talented'\n", + "43th char = ' person'\n", + "44th char = ','\n", + "45th char = ' and'\n", + "46th char = ' I'\n", + "47th char = ' am'\n", + "48th char = ' very'\n", + "49th char = ' excited'\n", + "50th char = ' to'\n", + "51th char = ' be'\n", + "52th char = ' able'\n", + "53th char = ' to'\n", + "54th char = ' share'\n", + "55th char = ' my'\n" + ] + } + ], + "source": [ + "print(f\"Sequence so far: {reference_gpt2.to_string(tokens)[0]!r}\")\n", + "\n", + "for i in range(20):\n", + " print(f\"{tokens.shape[-1] + 1}th char = {next_char!r}\")\n", + " # Define new input sequence, by appending the previously generated token\n", + " tokens = t.cat([tokens, next_token[None, None]], dim=-1)\n", + " # Pass our new sequence through the model, to get new output\n", + " logits = reference_gpt2(tokens)\n", + " # Get the predicted token at the end of our sequence\n", + " next_token = logits[0, -1].argmax(dim=-1)\n", + " # Decode and print the result\n", + " next_char = reference_gpt2.to_string(next_token)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "LhI8K0IYfdfN" + }, + "source": [ + "> ## Key takeaways\n", + ">\n", + "> * Transformer takes in language, predicts next token (for *each* token in a causal way)\n", + "> * We convert language to a sequence of integers with a tokenizer.\n", + "> * We convert integers to vectors with a lookup table.\n", + "> * Output is a vector of logits (one for each input token), we convert to a probability distn with a softmax, and can then convert this to a token (eg taking the largest logit, or sampling).\n", + "> * We append this to the input + run again to generate more text (Jargon: *autoregressive*)\n", + "> * Meta level point: Transformers are sequence operation models, they take in a sequence, do processing in parallel at each position, and use attention to move information between positions!" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "A1b6dp6NfdfN" + }, + "source": [ + "# 2️⃣ Clean Transformer Implementation\n", + "\n", + "> ##### Learning Objectives\n", + ">\n", + "> * Understand that a transformer is composed of attention heads and MLPs, with each one performing operations on the residual stream\n", + "> * Understand that the attention heads in a single layer operate independently, and that they have the role of calculating attention patterns (which determine where information is moved to & from in the residual stream)\n", + "> * Learn about & implement the following transformer modules:\n", + "> * LayerNorm (transforming the input to have zero mean and unit variance)\n", + "> * Positional embedding (a lookup table from position indices to residual stream vectors)\n", + "> * Attention (the method of computing attention patterns for residual stream vectors)\n", + "> * MLP (the collection of linear and nonlinear transformations which operate on each residual stream vector in the same way)\n", + "> * Embedding (a lookup table from tokens to residual stream vectors)\n", + "> * Unembedding (a matrix for converting residual stream vectors into a distribution over tokens)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "y6ZPhHBUfdfN" + }, + "source": [ + "## High-Level architecture\n", + "\n", + "Go watch Neel's [Transformer Circuits walkthrough](https://www.youtube.com/watch?v=KV5gbOmHbjU) if you want more intuitions!\n", + "\n", + "(Diagram is bottom to top, right-click and open for higher resolution.)\n", + "\n", + "" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "PM65adGvfdfN" + }, + "source": [ + "### Tokenization & Embedding\n", + "\n", + "The input tokens $t$ are integers. We get them from taking a sequence, and tokenizing it (like we saw in the previous section).\n", + "\n", + "The token embedding is a lookup table mapping tokens to vectors, which is implemented as a matrix $W_E$. The matrix consists of a stack of token embedding vectors (one for each token)." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "EbXwLDVTfdfN" + }, + "source": [ + "### Residual stream\n", + "\n", + "The residual stream is the sum of all previous outputs of layers of the model, and is also the input to each new layer. It has shape `[batch, seq_len, d_model]` (where `d_model` is the length of a single embedding vector).\n", + "\n", + "The initial value of the residual stream is denoted $x_0$ in the diagram, and $x_i$ are later values of the residual stream (after more attention and MLP layers have been applied to the residual stream).\n", + "\n", + "The residual stream is *really* fundamental. It's the central object of the transformer. It's how the model remembers things, moves information between layers for composition, and it's the medium used to store the information that attention moves between positions.\n", + "\n", + "
\n", + "Aside - logit lens\n", + "\n", + "A key idea of transformers is the [residual stream as output accumulation](https://www.lesswrong.com/posts/X26ksz4p3wSyycKNB/gears-level-mental-models-of-transformer-interpretability#Residual_Stream_as_Output_Accumulation:~:text=The%20Models-,Residual%20Stream%20as%20Output%20Accumulation,-The%20residual%20stream). As we move through the layers of the model, shifting information around and processing it, the values in the residual stream represent the accumulation of all the inferences made by the transformer up to that point.\n", + "\n", + "This is neatly illustrated by the **logit lens**. Rather than getting predictions from the residual stream at the very end of the model, we can take the value of the residual stream midway through the model and convert it to a distribution over tokens. When we do this, we find surprisingly coherent predictions, especially in the last few layers before the end.\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "0kHTZahlfdfN" + }, + "source": [ + "### Transformer blocks\n", + "\n", + "Then we have a series of `n_layers` **transformer blocks** (also sometimes called **residual blocks**).\n", + "\n", + "Note - a block contains an attention layer *and* an MLP layer, but we say a transformer has $k$ layers if it has $k$ blocks (i.e. $2k$ total layers).\n", + "\n", + "" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "lLBbNKbwfdfN" + }, + "source": [ + "### Attention\n", + "\n", + "First we have attention. This moves information from prior positions in the sequence to the current token.\n", + "\n", + "We do this for *every* token in parallel using the same parameters. The only difference is that we look backwards only (to avoid \"cheating\"). This means later tokens have more of the sequence that they can look at.\n", + "\n", + "Attention layers are the only bit of a transformer that moves information between positions (i.e. between vectors at different sequence positions in the residual stream).\n", + "\n", + "Attention layers are made up of `n_heads` heads - each with their own parameters, own attention pattern, and own information about how to copy things from source to destination. The heads act independently and additively, we just add their outputs together, and back to the stream.\n", + "\n", + "Each head does the following:\n", + "* Produces an **attention pattern** for each destination token, a probability distribution of prior source tokens (including the current one) weighting how much information to copy.\n", + "* Moves information (via a linear map) in the same way from each source token to each destination token.\n", + "\n", + "Each attention head is made up of three components: the keys, queries, and values (often abbreviated as K, Q and V). These names come from their analogy to retrieval systems. Broadly speaking:\n", + "\n", + "* **Queries** represent a question or request for information, e.g. \"I'm looking for a name that appeared earlier in this sentence\".\n", + "* **Keys** represent whether a source token's information matches the query, e.g. if the source token is \"Mary\" then this causes the key to have a high dot product with the query (we call this an **attention score**), and it means that a lot of information will be taken from this token.\n", + "* **Values** represent the information that actually gets moved. This sounds similar to keys, but it's actually different in an important way. For instance, the key might just contain the information \"this is a name\", but the value could be the actual name itself.\n", + "\n", + "The diagram below illustrates the three different parts, in the context of the analogy for transformers we introduced earlier. This is a simplified model for how the person holding the \"in\" token might come to figure out that the next token is \"Mary\". In later sections we'll look at the actual function performed by attention heads and see how the operations relate to this analogy.\n", + "\n", + "\n", + "\n", + "Another interesting intuition for attention is as a kind of \"generalized convolution\" - read the dropdown below if you want to learn more about this.\n", + "\n", + "
\n", + "Intuition - attention as generalized convolution\n", + "\n", + "We can think of attention as a kind of generalized convolution. Standard convolution layers work by imposing a \"prior of locality\", i.e. the assumption that pixels which are close together are more likely to share information. Although language has some locality (two words next to each other are more likely to share information than two words 100 tokens apart), the picture is a lot more nuanced, because which tokens are relevant to which others depends on the context of the sentence. For instance, in the sentence `\"When Mary and John went to the store, John gave a drink to Mary\"`, the names in this sentence are the most important tokens for predicting that the final token will be `\"Mary\"`, and this is because of the particular context of this sentence rather than the tokens' position.\n", + "\n", + "Attention layers are effectively our way of saying to the transformer, \"don't impose a prior of locality, but instead develop your own algorithm to figure out which tokens are important to which other tokens in any given sequence.\"\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "8WdavD1RfdfN" + }, + "source": [ + "Below is a schematic diagram of the attention layers. We'll go into much more detail during the actual implementation, so don't worry if this doesn't fully make sense yet.\n", + "\n", + "\n", + "\n", + "" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "oZ9QlGm1fdfN" + }, + "source": [ + "### MLP\n", + "\n", + "The MLP layers are just a standard neural network, with a singular hidden layer and a nonlinear activation function. The exact activation isn't conceptually important ([GELU](https://arxiv.org/abs/1606.08415) seems to perform best).\n", + "\n", + "Our hidden dimension is normally `d_mlp = 4 * d_model`. Exactly why the ratios are what they are isn't super important (people basically cargo-cult what GPT did back in the day!).\n", + "\n", + "Importantly, **the MLP operates on positions in the residual stream independently, and in exactly the same way**. It doesn't move information between positions.\n", + "\n", + "Once attention has moved relevant information to a single position in the residual stream, MLPs can actually do computation, reasoning, look up information, etc. *What the hell is going on inside MLPs* is a pretty big open problem in transformer mechanistic interpretability - see the [Toy Model of Superposition Paper](https://transformer-circuits.pub/2022/toy_model/index.html) for more on why this is hard.\n", + "\n", + "To go back to our analogy for transformers, we can essentially view MLPs as the thinking that each person in the line does once they've grabbed the information they need from the people behind them (via attention). Usually the MLP layers make up a much larger fraction of the model's total parameter count than attention layers (often around 2/3 although this varies between architectures), which makes sense since processing the information is a bigger task than just moving it around.\n", + "\n", + "" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "lS_J11YvfdfN" + }, + "source": [ + "Here are a few more intuitions for MLPs, which you might find interesting:\n", + "\n", + "
\n", + "Intuition - MLPs as key-value pairs\n", + "\n", + "We can write the MLP's output as $f(x^T W^{in})W^{out}$, where $W^{in}$ and $W^{out}$ are the different weights of the MLP (ignoring biases), $f$ is the activation function, and $x$ is a vector in the residual stream. This can be rewritten as:\n", + "\n", + "$$\n", + "f(x^T W^{in}) W^{out} = \\sum_{i=1}^{d_{mlp}} f(x^T W^{in}_{[:, i]}) W^{out}_{[i, :]}\n", + "$$\n", + "\n", + "We can view the vectors $W^{in}_{[:, i]}$ as the **input directions**, and $W^{out}_{[i, :]}$ as the **output directions**. We say the input directions are **activated** by certain textual features, and when they are activated, vectors are written in the corresponding output direction. This is very similar to the concept of keys and values in attention layers, which is why these vectors are also sometimes called keys and values (e.g. see the paper [Transformer Feed-Forward Layers Are Key-Value Memories](https://arxiv.org/pdf/2012.14913.pdf)).\n", + "\n", + "Terminology note - sometimes we refer to each of these $d_{mlp}$ input-output pairs as **neurons**.\n", + "\n", + "\n", + "\n", + "---\n", + "\n", + "Here's a step-by-step breakdown of the linear algebra, if it was too fast above. We have:\n", + "\n", + "$$\n", + "\\begin{aligned}\n", + "x^T W^{in} &= x^T [W^{in}_{[:, 1]}\\,, ...\\;, W^{in}_{[:, n]}] \\\\\n", + "&= (x^T W^{in}_{[:, 1]}\\,, \\; ...\\;, \\; x^T W^{in}_{[:, n]})\n", + "\\end{aligned}\n", + "$$\n", + "\n", + "where $W^{in}_{[:, i]}$ are the columns of $W^{in}$. In other words, these values (the pre-GELU activations) are projections of $x$ along the input directions of the neurons.\n", + "\n", + "If we add our activation function and the second matrix, then we get:\n", + "\n", + "$$\n", + "\\begin{aligned}\n", + "f(x^T W^{in})W^{out} &= (f(x^T W^{in}_{[:, 1]})\\,, \\; ...\\;,\\; f(x^T W^{in}_{[:, n]})) \\begin{bmatrix} \\leftarrow W^{out}_{[1, :]} \\rightarrow \\\\ \\vdots \\\\ \\leftarrow W^{out}_{[n, :]} \\rightarrow \\end{bmatrix} \\\\\n", + "&= f(x^T W^{in}_{[:, 1]}) W^{out}_{[1, :]} + \\;...\\; + f(x^T W^{in}_{[:, n]}) W^{out}_{[n, :]} \\\\\n", + "&= \\sum_{i=1}^n f(x^T W^{in}_{[:, i]}) W^{out}_{[i, :]}\n", + "\\end{aligned}\n", + "$$\n", + "\n", + "where $W^{out}_{[i, :]}$ are the rows of $W^{out}$. In other words, our output is a linear combination of the rows of $W^{out}$, with the coefficients of that linear combination given by the projections of $x$ along the columns of $W^{in}$.\n", + "\n", + "
\n", + "\n", + "
\n", + "Intuition - MLPs as knowledge storage\n", + "\n", + "We can think of MLPs as where knowledge gets stored in our transformer. The attention mechanism is what moves information around between sequence positions, but the MLPs are where this information is processed, and new information is written into the residual stream which is a function of the old information.\n", + "\n", + "This is deeply connected to the key-value pairs model, since you can treat key-value pairs as a kind of associative memory system (where the key serves as a unique identifier, and the value holds the related information).\n", + "\n", + "Another related intuition (for which there is some evidence) is **MLPs as memory management**. In an idealized case, we might find that the $i$-th neuron satisfies $W^{in}_{[:, i]} \\approx - W^{out}_{[i, :]} \\approx \\vec v$ for some unit vector $\\vec v$, meaning it may be responsible for erasing the positive component of vector $\\vec x$ in the direction $\\vec v$ (exercise - can you show why this is the case?). This can free up space in the residual stream for other components to write to.\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "oKIOzQ99fdfO" + }, + "source": [ + "Lastly, here's a schematic diagram of the MLP layers. Again, we'll go into much more detail during the actual implementation, so don't worry if this doesn't fully make sense yet.\n", + "\n", + "" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "XV3YK04sfdfO" + }, + "source": [ + "### Unembedding\n", + "\n", + "Finally, we unembed!\n", + "\n", + "This just consists of applying a linear map $W_U$, going from final residual stream to a vector of logits - this is the output.\n", + "\n", + "
\n", + "Aside - tied embeddings\n", + "\n", + "Note - sometimes we use something called a **tied embedding** - this is where we use the same weights for our $W_E$ and $W_U$ matrices. In other words, to get the logit score for a particular token at some sequence position, we just take the vector in the residual stream at that sequence position and take the inner product with the corresponding token embedding vector. This is more training-efficient (because there are fewer parameters in our model), and it might seem principled at first. After all, if two words have very similar meanings, shouldn't they have similar embedding vectors because the model will treat them the same, and similar unembedding vectors because they could both be substituted for each other in most output?\n", + "\n", + "However, this is actually not very principled, for the following main reason: **the direct path involving the embedding and unembedding should approximate bigram frequencies**.\n", + "\n", + "Let's break down this claim. **Bigram frequencies** refers to the frequencies of pairs of words in the English language (e.g. the bigram frequency of \"Barack Obama\" is much higher than the product of the individual frequencies of the words \"Barack\" and \"Obama\"). If our model had no attention heads or MLP layers, then all we have is a linear map from our one-hot encoded token `T` to a probability distribution over the token following `T`. This map is represented by the linear transformation $t \\to t^T W_E W_U$ (where $t$ is our one-hot encoded token vector). Since the output of this transformation can only be a function of the token `T` (and no earlier tokens), the best we can do is have this map approximate the true frequency of bigrams starting with `T`, which appear in the training data. Importantly, **this is not a symmetric map**. We want `T = \"Barack\"` to result in a high probability of the next token being `\"Obama\"`, but not the other way around!\n", + "\n", + "Even in multi-layer models, a similar principle applies. There will be more paths through the model than just the \"direct path\" $W_E W_U$, but because of the residual connections there will always exist a direct path, so there will always be some incentive for $W_E W_U$ to approximate bigram frequencies.\n", + "\n", + "That being said, smaller (<8B parameter) LLMs still often use tied embeddings to improve training and inference efficiency. It can be easier to start from tied weights and then use MLP0 to break the symmetry than to initialize encoder and decoder with no shared structure at all.\n", + "\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "tsNkJH9bfdfO" + }, + "source": [ + "### LayerNorm\n", + "\n", + "* Simple normalization function applied at the start of each layer (i.e. before each MLP, attention layer, and before the unembedding)\n", + "* Converts each input vector (independently in parallel for each `(batch, seq)` residual stream vector) to have mean zero and variance 1.\n", + "* Then applies an elementwise scaling and translation\n", + "* Cool maths tangent: The scale ($\\odot \\gamma$) & translate ($+ \\beta$) is just a linear map. LayerNorm is only applied immediately before another linear map (either the MLP, or the query/key/value linear maps in the attention head, or the unembedding $W_U$). Linear compose linear = linear, so we can just fold this into a single effective linear layer and ignore it.\n", + " * `fold_ln=True` flag in `from_pretrained` does this for you.\n", + "* LayerNorm is annoying for interpretability - it would be linear if not for the fact we divide by the variance, so you can't decompose the contributions of the input to the output independently. But it's *almost* linear - if you're changing a small part of the input you can pretend $\\sqrt{\\text{Var}[x] + \\epsilon}$ is constant, so the LayerNorm operation is linear, but if you're changing $x$ enough to alter the norm substantially it's not linear.\n", + "\n", + "\n", + "\n", + "\n", + "### Positional embeddings\n", + "\n", + "* **Problem:** Attention operates over all pairs of positions. This means it's symmetric with regard to position - the attention calculation from token 5 to token 1 and token 5 to token 2 are the same by default\n", + " * This is dumb because nearby tokens are more relevant.\n", + "* There's a lot of dumb hacks for this.\n", + "* We'll focus on **learned, absolute positional embeddings**. This means we learn a lookup table mapping the index of the position of each token to a residual stream vector, and add this to the embed.\n", + " * Note that we *add* rather than concatenate. This is because the residual stream is shared memory, and likely under significant superposition (the model compresses more features in there than the model has dimensions)\n", + " * We basically never concatenate inside a transformer, unless doing weird shit like generating text efficiently.\n", + "* This connects to **attention as generalized convolution**\n", + " * We argued that language does still have locality, and so it's helpful for transformers to have access to the positional information so they \"know\" two tokens are next to each other (and hence probably relevant to each other)." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "mzvVga3WfdfO" + }, + "source": [ + "## Actual Code!\n", + "\n", + "Model architecture table (this will be helpful for understanding the results you get when running the code block below):\n", + "\n", + "| Parameter | Value |\n", + "|-------------|----------------|\n", + "| batch | 1 |\n", + "| position | 35 |\n", + "| d_model | 768 |\n", + "| n_heads | 12 |\n", + "| n_layers | 12 |\n", + "| d_mlp | 3072 (= 4 * `d_model`) |\n", + "| d_head | 64 (= `d_model / n_heads`) |" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "yErW0QyXfdfO" + }, + "source": [ + "### Parameters and Activations\n", + "\n", + "It's important to distinguish between parameters and activations in the model.\n", + "\n", + "* **Parameters** are the weights and biases that are learned during training.\n", + " * These don't change when the model input changes.\n", + "* **Activations** are temporary numbers calculated during a forward pass, that are functions of the input.\n", + " * We can think of these values as only existing for the duration of a single forward pass, and disappearing afterwards.\n", + " * We can use hooks to access these values during a forward pass (more on hooks later), but it doesn't make sense to talk about a model's activations outside the context of some particular input.\n", + " * Attention scores and patterns are activations (this is slightly non-intuitive because they're used in a matrix multiplication with another activation).\n", + "\n", + "#### Print All Activation Shapes of Reference Model\n", + "\n", + "Run the following code to print all the activation shapes of the reference model:" + ] + }, + { + "cell_type": "code", + "execution_count": 167, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "JK-VmZjifdfO", + "outputId": "b7f51079-7802-4c6a-ef53-e37329b39f48" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "hook_embed (1, 35, 768)\n", + "hook_pos_embed (1, 35, 768)\n", + "blocks.0.hook_resid_pre (1, 35, 768)\n", + "blocks.0.ln1.hook_scale (1, 35, 1)\n", + "blocks.0.ln1.hook_normalized (1, 35, 768)\n", + "blocks.0.attn.hook_q (1, 35, 12, 64)\n", + "blocks.0.attn.hook_k (1, 35, 12, 64)\n", + "blocks.0.attn.hook_v (1, 35, 12, 64)\n", + "blocks.0.attn.hook_attn_scores (1, 12, 35, 35)\n", + "blocks.0.attn.hook_pattern (1, 12, 35, 35)\n", + "blocks.0.attn.hook_z (1, 35, 12, 64)\n", + "blocks.0.hook_attn_out (1, 35, 768)\n", + "blocks.0.hook_resid_mid (1, 35, 768)\n", + "blocks.0.ln2.hook_scale (1, 35, 1)\n", + "blocks.0.ln2.hook_normalized (1, 35, 768)\n", + "blocks.0.mlp.hook_pre (1, 35, 3072)\n", + "blocks.0.mlp.hook_post (1, 35, 3072)\n", + "blocks.0.hook_mlp_out (1, 35, 768)\n", + "blocks.0.hook_resid_post (1, 35, 768)\n", + "ln_final.hook_scale (1, 35, 1)\n", + "ln_final.hook_normalized (1, 35, 768)\n" + ] + } + ], + "source": [ + "for activation_name, activation in cache.items():\n", + " # Only print for first layer\n", + " if \".0.\" in activation_name or \"blocks\" not in activation_name:\n", + " print(f\"{activation_name:30} {tuple(activation.shape)}\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "qL9pZMrmfdfO" + }, + "source": [ + "#### Print All Parameter Shapes of Reference Model" + ] + }, + { + "cell_type": "code", + "execution_count": 168, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "-LQLe6GmfdfO", + "outputId": "7e1f1114-36b6-48e7-9da3-eaa86b4fef14" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "embed.W_E (50257, 768)\n", + "pos_embed.W_pos (1024, 768)\n", + "blocks.0.ln1.w (768,)\n", + "blocks.0.ln1.b (768,)\n", + "blocks.0.ln2.w (768,)\n", + "blocks.0.ln2.b (768,)\n", + "blocks.0.attn.W_Q (12, 768, 64)\n", + "blocks.0.attn.W_O (12, 64, 768)\n", + "blocks.0.attn.b_Q (12, 64)\n", + "blocks.0.attn.b_O (768,)\n", + "blocks.0.attn.W_K (12, 768, 64)\n", + "blocks.0.attn.W_V (12, 768, 64)\n", + "blocks.0.attn.b_K (12, 64)\n", + "blocks.0.attn.b_V (12, 64)\n", + "blocks.0.mlp.W_in (768, 3072)\n", + "blocks.0.mlp.b_in (3072,)\n", + "blocks.0.mlp.W_out (3072, 768)\n", + "blocks.0.mlp.b_out (768,)\n", + "ln_final.w (768,)\n", + "ln_final.b (768,)\n", + "unembed.W_U (768, 50257)\n", + "unembed.b_U (50257,)\n" + ] + } + ], + "source": [ + "for name, param in reference_gpt2.named_parameters():\n", + " # Only print for first layer\n", + " if \".0.\" in name or \"blocks\" not in name:\n", + " print(f\"{name:18} {tuple(param.shape)}\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "igouy9uPfdfO" + }, + "source": [ + "[This diagram](https://raw.githubusercontent.com/chloeli-15/ARENA_img/main/img/full-merm.svg) shows the names of all activations and parameters in a fully general transformer model from transformerlens (except for a few at the start and end, like the embedding and unembedding). Lots of this won't make sense at first, but you can return to this diagram later and check that you understand most/all parts of it.\n", + "\n", + "There's also an annotated version [here](https://raw.githubusercontent.com/chloeli-15/ARENA_img/main/img/transformer-full-updated.png)." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "G23lcMZKfdfO" + }, + "source": [ + "### Config\n", + "\n", + "The config object contains all the hyperparameters of the model. We can print the config of the reference model to see what it contains:" + ] + }, + { + "cell_type": "code", + "execution_count": 169, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "KJA2wjbxfdfO", + "outputId": "6cab6a3b-1203-4aee-dda8-a6da125379b2" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "HookedTransformerConfig:\n", + "{'NTK_by_parts_factor': 8.0,\n", + " 'NTK_by_parts_high_freq_factor': 4.0,\n", + " 'NTK_by_parts_low_freq_factor': 1.0,\n", + " 'NTK_original_ctx_len': 8192,\n", + " 'act_fn': 'gelu_new',\n", + " 'attention_dir': 'causal',\n", + " 'attn_only': False,\n", + " 'attn_scale': 8.0,\n", + " 'attn_scores_soft_cap': -1.0,\n", + " 'attn_types': None,\n", + " 'checkpoint_index': None,\n", + " 'checkpoint_label_type': None,\n", + " 'checkpoint_value': None,\n", + " 'd_head': 64,\n", + " 'd_mlp': 3072,\n", + " 'd_model': 768,\n", + " 'd_vocab': 50257,\n", + " 'd_vocab_out': 50257,\n", + " 'decoder_start_token_id': None,\n", + " 'default_prepend_bos': True,\n", + " 'device': device(type='cpu'),\n", + " 'dtype': torch.float32,\n", + " 'eps': 1e-05,\n", + " 'experts_per_token': None,\n", + " 'final_rms': False,\n", + " 'from_checkpoint': False,\n", + " 'gated_mlp': False,\n", + " 'init_mode': 'gpt2',\n", + " 'init_weights': False,\n", + " 'initializer_range': 0.02886751345948129,\n", + " 'load_in_4bit': False,\n", + " 'model_name': 'gpt2',\n", + " 'n_ctx': 1024,\n", + " 'n_devices': 1,\n", + " 'n_heads': 12,\n", + " 'n_key_value_heads': None,\n", + " 'n_layers': 12,\n", + " 'n_params': 84934656,\n", + " 'normalization_type': 'LN',\n", + " 'num_experts': None,\n", + " 'original_architecture': 'GPT2LMHeadModel',\n", + " 'output_logits_soft_cap': -1.0,\n", + " 'parallel_attn_mlp': False,\n", + " 'positional_embedding_type': 'standard',\n", + " 'post_embedding_ln': False,\n", + " 'relative_attention_max_distance': None,\n", + " 'relative_attention_num_buckets': None,\n", + " 'rotary_adjacent_pairs': False,\n", + " 'rotary_base': 10000,\n", + " 'rotary_base_local': None,\n", + " 'rotary_dim': None,\n", + " 'scale_attn_by_inverse_layer_idx': False,\n", + " 'seed': None,\n", + " 'tie_word_embeddings': False,\n", + " 'tokenizer_name': 'gpt2',\n", + " 'tokenizer_prepends_bos': True,\n", + " 'trust_remote_code': False,\n", + " 'ungroup_grouped_query_attention': False,\n", + " 'use_NTK_by_parts_rope': False,\n", + " 'use_attn_in': False,\n", + " 'use_attn_result': False,\n", + " 'use_attn_scale': True,\n", + " 'use_hook_mlp_in': False,\n", + " 'use_hook_tokens': False,\n", + " 'use_local_attn': False,\n", + " 'use_normalization_before_and_after': False,\n", + " 'use_qk_norm': False,\n", + " 'use_split_qkv_input': False,\n", + " 'window_size': None}\n" + ] + } + ], + "source": [ + "# As a reference - note there's a lot of stuff we don't care about in here, to do with library internals or other architectures\n", + "print(reference_gpt2.cfg)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "KOeKi3K3fdfO" + }, + "source": [ + "We define a stripped down config for our model:" + ] + }, + { + "cell_type": "code", + "execution_count": 170, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "IjUVtNKOfdfO", + "outputId": "13600e5f-e552-4e27-8e4c-1bf410060e56" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "Config(d_model=768, debug=True, layer_norm_eps=1e-05, d_vocab=50257, init_range=0.02, n_ctx=1024, d_head=64, d_mlp=3072, n_heads=12, n_layers=12)\n" + ] + } + ], + "source": [ + "@dataclass\n", + "class Config:\n", + " d_model: int = 768\n", + " debug: bool = True\n", + " layer_norm_eps: float = 1e-5\n", + " d_vocab: int = 50257\n", + " init_range: float = 0.02\n", + " n_ctx: int = 1024\n", + " d_head: int = 64\n", + " d_mlp: int = 3072\n", + " n_heads: int = 12\n", + " n_layers: int = 12\n", + "\n", + "\n", + "cfg = Config()\n", + "print(cfg)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "BlpE2t33fdfO" + }, + "source": [ + "### Exercise - implement `LayerNorm`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴🔴⚪⚪\n", + "> Importance: 🔵🔵🔵⚪⚪\n", + ">\n", + "> You should spend up to 10-15 minutes on this exercise.\n", + "> ```\n", + "\n", + "You should fill in the code below, and then run the tests to verify that your layer is working correctly.\n", + "\n", + "Your LayerNorm should do the following:\n", + "\n", + "* Make mean 0\n", + "* Normalize to have variance 1\n", + "* Scale with learned weights\n", + "* Translate with learned bias\n", + "\n", + "You can use the PyTorch [LayerNorm documentation](https://pytorch.org/docs/stable/generated/torch.nn.LayerNorm.html) as a reference. A few more notes:\n", + "\n", + "* Your layernorm implementation always has `affine=True`, i.e. you do learn parameters `w` and `b` (which are represented as $\\gamma$ and $\\beta$ respectively in the PyTorch documentation).\n", + "* Remember that, after the centering and normalization, each vector of length `d_model` in your input should have mean 0 and variance 1.\n", + "* As the PyTorch documentation page says, your variance should be computed using `unbiased=False`.\n", + "* The `layer_norm_eps` argument in your config object corresponds to the $\\epsilon$ term in the PyTorch documentation (it is included to avoid division-by-zero errors).\n", + "* We've given you a `debug` argument in your config. If `debug=True`, then you can print output like the shape of objects in your `forward` function to help you debug (this is a very useful trick to improve your coding speed).\n", + "\n", + "Fill in the function, where it says `raise NotImplementedError()` (this will be the basic pattern for most other exercises in this section)." + ] + }, + { + "cell_type": "code", + "source": [ + "x = t.tensor([1., 2., 3.])\n", + "\n", + "print(t.var(x, unbiased=False))\n", + "print(t.var(x, unbiased=True))" + ], + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "5J2ybNmio5-6", + "outputId": "ebbe57b0-7055-401d-961d-fc3035c68cd9" + }, + "execution_count": 171, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "tensor(0.6667)\n", + "tensor(1.)\n" + ] + } + ] + }, + { + "cell_type": "code", + "execution_count": 172, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "nnQYof2TfdfO", + "outputId": "5ef8b939-5517-4ea2-a0f2-3ca94e8e69b3" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "Input shape: torch.Size([2, 4, 768])\n", + "Output shape: torch.Size([2, 4, 768]) \n", + "\n", + "Input shape: torch.Size([1, 35, 768])\n", + "Output shape: torch.Size([1, 35, 768])\n", + "Reference output shape: torch.Size([1, 35, 768]) \n", + "\n", + "100.00% of the values are correct\n", + "\n", + "All tests in `test_layer_norm_epsilon` passed!\n" + ] + } + ], + "source": [ + "import math\n", + "class LayerNorm(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.w = nn.Parameter(t.ones(cfg.d_model))\n", + " self.b = nn.Parameter(t.zeros(cfg.d_model))\n", + "\n", + " def forward(self, residual: Float[Tensor, \"batch posn d_model\"]) -> Float[Tensor, \"batch posn d_model\"]:\n", + " m = residual.mean(dim=2, keepdim=True)\n", + " std = (residual.var(dim=2, keepdim=True, unbiased=False) + self.cfg.layer_norm_eps).sqrt()\n", + "\n", + " residual = (residual-m)/std\n", + " return residual*self.w + self.b\n", + "\n", + "\n", + "\n", + "\n", + "\n", + "tests.rand_float_test(LayerNorm, [2, 4, 768])\n", + "tests.load_gpt2_test(LayerNorm, reference_gpt2.ln_final, cache[\"resid_post\", 11])\n", + "tests.test_layer_norm_epsilon(LayerNorm, cache[\"resid_post\", 11])" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "DyK78SCpfdfO" + }, + "source": [ + "
Solution\n", + "\n", + "```python\n", + "class LayerNorm(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.w = nn.Parameter(t.ones(cfg.d_model))\n", + " self.b = nn.Parameter(t.zeros(cfg.d_model))\n", + "\n", + " def forward(self, residual: Float[Tensor, \"batch posn d_model\"]) -> Float[Tensor, \"batch posn d_model\"]:\n", + " residual_mean = residual.mean(dim=-1, keepdim=True)\n", + " residual_std = (residual.var(dim=-1, keepdim=True, unbiased=False) + self.cfg.layer_norm_eps).sqrt()\n", + "\n", + " residual = (residual - residual_mean) / residual_std\n", + " return residual * self.w + self.b\n", + "```\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "v58X_BOKfdfO" + }, + "source": [ + "### Exercise - implement `Embed`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴⚪⚪⚪\n", + "> Importance: 🔵🔵🔵⚪⚪\n", + ">\n", + "> You should spend up to 5-10 minutes on this exercise.\n", + "> ```\n", + "\n", + "This is basically a lookup table from tokens to residual stream vectors.\n", + "\n", + "(Hint - you can implement this in just one line, without any complicated functions. If you've been working on it for >10 mins, you're probably overthinking it!)" + ] + }, + { + "cell_type": "code", + "execution_count": 173, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "Wq4XNTbgfdfO", + "outputId": "67286486-1742-4156-b8b5-cd589fa38e32" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "All tests in `test_embed` passed!\n", + "Input shape: torch.Size([2, 4])\n", + "Output shape: torch.Size([2, 4, 768]) \n", + "\n", + "Input shape: torch.Size([1, 55])\n", + "Output shape: torch.Size([1, 55, 768])\n", + "Reference output shape: torch.Size([1, 55, 768]) \n", + "\n", + "100.00% of the values are correct\n", + "\n" + ] + } + ], + "source": [ + "class Embed(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.W_E = nn.Parameter(t.empty((cfg.d_vocab, cfg.d_model)))\n", + " nn.init.normal_(self.W_E, std=self.cfg.init_range)\n", + "\n", + " def forward(self, tokens: Int[Tensor, \"batch position\"]) -> Float[Tensor, \"batch position d_model\"]:\n", + " return self.W_E[tokens]\n", + "\n", + "\n", + "\n", + "tests.test_embed(Embed)\n", + "tests.rand_int_test(Embed, [2, 4])\n", + "tests.load_gpt2_test(Embed, reference_gpt2.embed, tokens)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "uoRXyuMVfdfO" + }, + "source": [ + "
\n", + "Help - I keep getting RuntimeError: CUDA error: device-side assert triggered.\n", + "\n", + "This is a uniquely frustrating type of error message, because it (1) forces you to restart the kernel, and (2) often won't tell you where the error message actually originated from!\n", + "\n", + "You can fix the second problem by adding the line `os.environ['CUDA_LAUNCH_BLOCKING'] = \"1\"` to the very top of your file (after importing `os`). This won't fix your bug, but it makes sure the correct origin point is identified.\n", + "\n", + "As for actually fixing the bug, this error usually ends up being the result of bad indexing, e.g. you're trying to apply an embedding layer to tokens which are larger than your maximum embedding.\n", + "
\n", + "\n", + "\n", + "
Solution\n", + "\n", + "```python\n", + "class Embed(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.W_E = nn.Parameter(t.empty((cfg.d_vocab, cfg.d_model)))\n", + " nn.init.normal_(self.W_E, std=self.cfg.init_range)\n", + "\n", + " def forward(self, tokens: Int[Tensor, \"batch position\"]) -> Float[Tensor, \"batch position d_model\"]:\n", + " return self.W_E[tokens]\n", + "```\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "ZWrhjEtPfdfO" + }, + "source": [ + "### Exercise - implement `PosEmbed`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴⚪⚪⚪\n", + "> Importance: 🔵🔵🔵⚪⚪\n", + ">\n", + "> You should spend up to 10-15 minutes on this exercise.\n", + "> ```\n", + "\n", + "Positional embedding can also be thought of as a lookup table, but rather than the indices being our token IDs, the indices are just the numbers `0`, `1`, `2`, ..., `seq_len-1` (i.e. the position indices of the tokens in the sequence)." + ] + }, + { + "cell_type": "code", + "execution_count": 174, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "8B6lqs4SfdfO", + "outputId": "da198503-eafe-4578-8b32-f143a385defb" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "All tests in `test_pos_embed` passed!\n", + "Input shape: torch.Size([2, 4])\n", + "Output shape: torch.Size([2, 4, 768]) \n", + "\n", + "Input shape: torch.Size([1, 55])\n", + "Output shape: torch.Size([1, 55, 768])\n", + "Reference output shape: torch.Size([1, 55, 768]) \n", + "\n", + "100.00% of the values are correct\n", + "\n" + ] + } + ], + "source": [ + "class PosEmbed(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.W_pos = nn.Parameter(t.empty((cfg.n_ctx, cfg.d_model)))\n", + " nn.init.normal_(self.W_pos, std=self.cfg.init_range)\n", + "\n", + " def forward(self, tokens: Int[Tensor, \"batch position\"]) -> Float[Tensor, \"batch position d_model\"]:\n", + " positions = torch.arange(tokens.size(1)).expand_as(tokens)\n", + " return self.W_pos[positions]\n", + "\n", + "\n", + "tests.test_pos_embed(PosEmbed)\n", + "tests.rand_int_test(PosEmbed, [2, 4])\n", + "tests.load_gpt2_test(PosEmbed, reference_gpt2.pos_embed, tokens)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "ihjNFjgmfdfP" + }, + "source": [ + "
Solution\n", + "\n", + "```python\n", + "class PosEmbed(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.W_pos = nn.Parameter(t.empty((cfg.n_ctx, cfg.d_model)))\n", + " nn.init.normal_(self.W_pos, std=self.cfg.init_range)\n", + "\n", + " def forward(self, tokens: Int[Tensor, \"batch position\"]) -> Float[Tensor, \"batch position d_model\"]:\n", + " batch, seq_len = tokens.shape\n", + " return einops.repeat(self.W_pos[:seq_len], \"seq d_model -> batch seq d_model\", batch=batch)\n", + "```\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "AkBBXtiffdfP" + }, + "source": [ + "### Exercise - implement `apply_causal_mask`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴⚪⚪⚪\n", + "> Importance: 🔵🔵🔵🔵🔵\n", + ">\n", + "> You should spend up to 10-15 minutes on this exercise.\n", + "> ```\n", + "\n", + "The causal mask function will be a method of the `Attention` class.\n", + "It will take in attention scores, and apply a mask to them so that the model\n", + "can only attend to previous positions (i.e. the model can't cheat by looking at future positions).\n", + "We will implement this function first, and test it, before moving on to the `forward` method\n", + "of the `Attention` class.\n", + "\n", + "A few hints:\n", + "\n", + "* You can use [`torch.where`](https://pytorch.org/docs/stable/generated/torch.where.html), or the [`torch.masked_fill_`](https://pytorch.org/docs/stable/generated/torch.Tensor.masked_fill.html) function when masking the attention scores.\n", + "* The [`torch.triu`](https://pytorch.org/docs/stable/generated/torch.triu.html) function is useful for creating a mask that is True for all positions we want to set probabilities to zero for.\n", + "* Make sure to use the `self.IGNORE` attribute to set the masked positions to negative infinity.\n", + "\n", + "
\n", + "Question - why do you think we mask the attention scores by setting them to negative infinity, rather than the attention probabilities by setting them to zero?\n", + "\n", + "If we masked the attention probabilities, then the probabilities would no longer sum to 1.\n", + "\n", + "We want to mask the scores and *then* take softmax, so that the probabilities are still valid probabilities (i.e. they sum to 1), and the values in the masked positions have no influence on the model's output.\n", + "
" + ] + }, + { + "cell_type": "code", + "execution_count": 175, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "fLdxoNDlfdfP", + "outputId": "74eb38c4-a042-439a-cfca-5c6546546813" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "All tests in `test_causal_mask` passed!\n" + ] + } + ], + "source": [ + "class Attention(nn.Module):\n", + " IGNORE: Float[Tensor, \"\"]\n", + "\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.register_buffer(\"IGNORE\", t.tensor(float(\"-inf\"), dtype=t.float32, device=device))\n", + "\n", + " def apply_causal_mask(\n", + " self,\n", + " attn_scores: Float[Tensor, \"batch n_heads query_pos key_pos\"],\n", + " ) -> Float[Tensor, \"batch n_heads query_pos key_pos\"]:\n", + " \"\"\"\n", + " Applies a causal mask to attention scores, and returns masked scores.\n", + " \"\"\"\n", + " k = t.arange(attn_scores.size(3))[None, :]\n", + " q = t.arange(attn_scores.size(2))[:, None]\n", + " m = k > q\n", + "\n", + " return t.where(\n", + " m[None, None, :, :],\n", + " float(\"-inf\"),\n", + " attn_scores\n", + " )\n", + "\n", + "\n", + "tests.test_causal_mask(Attention.apply_causal_mask)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "Y29s2zzufdfP" + }, + "source": [ + "
\n", + "Hint (pseudocode)\n", + "\n", + "```python\n", + "def apply_causal_mask(\n", + " self, attn_scores: Float[Tensor, \"batch n_heads query_pos key_pos\"]\n", + " ) -> Float[Tensor, \"batch n_heads query_pos key_pos\"]:\n", + "\n", + " # Define a mask that is True for all positions we want to set probabilities to zero for\n", + "\n", + " # Apply the mask to attention scores, then return the masked scores\n", + "```\n", + "
\n", + "\n", + "\n", + "
Solution\n", + "\n", + "```python\n", + "class Attention(nn.Module):\n", + " IGNORE: Float[Tensor, \"\"]\n", + "\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.register_buffer(\"IGNORE\", t.tensor(float(\"-inf\"), dtype=t.float32, device=device))\n", + "\n", + " def apply_causal_mask(\n", + " self,\n", + " attn_scores: Float[Tensor, \"batch n_heads query_pos key_pos\"],\n", + " ) -> Float[Tensor, \"batch n_heads query_pos key_pos\"]:\n", + " \"\"\"\n", + " Applies a causal mask to attention scores, and returns masked scores.\n", + " \"\"\"\n", + " # Define a mask that is True for all positions we want to set probabilities to zero for\n", + " all_ones = t.ones(attn_scores.size(-2), attn_scores.size(-1), device=attn_scores.device)\n", + " mask = t.triu(all_ones, diagonal=1).bool()\n", + " # Apply the mask to attention scores, then return the masked scores\n", + " attn_scores.masked_fill_(mask, self.IGNORE)\n", + " return attn_scores\n", + "```\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "vpK3nRm4fdfP" + }, + "source": [ + "### Exercise - implement `Attention`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴🔴🔴⚪\n", + "> Importance: 🔵🔵🔵🔵🔵\n", + ">\n", + "> You should spend up to 30-45 minutes on this exercise.\n", + "> ```\n", + "\n", + "* **Step 1:** Produce an attention pattern - for each destination token, probability distribution over previous tokens (including current token)\n", + " * Linear map from input -> query, key shape `[batch, seq_posn, head_index, d_head]`\n", + " * Dot product every *pair* of queries and keys to get attn_scores `[batch, head_index, query_pos, key_pos]` (query = dest, key = source)\n", + " * **Scale** and mask `attn_scores` to make it lower triangular, i.e. causal\n", + " * Softmax along the `key_pos` dimension, to get a probability distribution for each query (destination) token - this is our attention pattern!\n", + "* **Step 2:** Move information from source tokens to destination token using attention pattern (move = apply linear map)\n", + " * Linear map from input -> value `[batch, key_pos, head_index, d_head]`\n", + " * Mix along the `key_pos` with attn pattern to get `z`, which is a weighted average of the value vectors `[batch, query_pos, head_index, d_head]`\n", + " * Map to output, `[batch, position, d_model]` (position = query_pos, we've summed over all heads)\n", + "\n", + "Note - when we say **scale**, we mean dividing by `sqrt(d_head)`. The purpose of this is to avoid vanishing gradients (which is a big problem when we're dealing with a function like softmax - if one of the values is much larger than all the others, the probabilities will be close to 0 or 1, and the gradients will be close to 0).\n", + "\n", + "Below is a much larger, more detailed version of the attention head diagram from earlier. This should give you an idea of the actual tensor operations involved. A few clarifications on this diagram:\n", + "\n", + "* Whenever there is a third dimension shown in the pictures, this refers to the `head_index` dimension. We can see that all operations within the attention layer are done independently for each head.\n", + "* The objects in the box are activations; they have a batch dimension (for simplicity, we assume the batch dimension is 1 in the diagram). The objects to the right of the box are our parameters (weights and biases); they have no batch dimension.\n", + "* We arrange the keys, queries and values as `(batch, seq_pos, head_idx, d_head)`, because the biases have shape `(head_idx, d_head)`, so this makes it convenient to add the biases (recall the rules of array broadcasting!)." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "Fit3iKjNfdfP" + }, + "source": [ + "" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "yfgVV1xffdfP" + }, + "source": [ + "
\n", + "A few extra notes on attention (optional)\n", + "\n", + "\n", + "\n", + "Here, we cover some details related to the mathematical formulation of attention heads (and in particular the separation of **QK** and **OV** circuits), which is something we dive a lot deeper into in the next set of exercises in this chapter.\n", + "\n", + "The **QK** circuit consists of the operation of the $W_Q$ and $W_K$ matrices. In other words, it determines the attention pattern, i.e. where information is moved to and from in the residual stream. The functional form of the attention pattern $A$ is:\n", + "\n", + "$$\n", + "A = \\text{softmax}\\left(\\frac{x W_Q W_K^T x^T}{\\sqrt{d_{head}}}\\right)\n", + "$$\n", + "\n", + "where $x$ is the residual stream (shape `[seq_len, d_model]`), and $W_Q$, $W_K$ are the weight matrices for a single head (i.e. shape `[d_model, d_head]`).\n", + "\n", + "The **OV** circuit consists of the operation of the $W_V$ and $W_O$ matrices. Once attention patterns are fixed, these matrices operate on the residual stream at the source position, and their output is the thing which gets moved from source to destination position.\n", + "\n", + "The diagram below shows the functional form of the OV circuit. The QK circuit (pink) is responsible for causing the destination token to attend to the source token, and the OV circuit (light brown) is what actually maps the source token data into the information we'll send to the destination token.\n", + "\n", + "\n", + "\n", + "The functional form of an entire attention head is:\n", + "\n", + "$$\n", + "\\begin{aligned}\n", + "\\text{output} &= \\text{softmax}\\left(\\frac{x W_Q W_K^T x^T}{\\sqrt{d_{head}}}\\right) (x W_V W_O) \\\\\n", + " &= Ax W_V W_O\n", + "\\end{aligned}\n", + "$$\n", + "\n", + "where $W_V$ has shape `[d_model, d_head]`, and $W_O$ has shape `[d_head, d_model]`.\n", + "\n", + "Here, we can clearly see that the **QK circuit** and **OV circuit** are doing conceptually different things, and should be thought of as two distinct parts of the attention head.\n", + "\n", + "Again, don't worry if you don't follow all of this right now - we'll go into **much** more detail on all of this in subsequent exercises. The purpose of the discussion here is just to give you a flavour of what's to come!\n", + "\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "KEZ92SzOfdfP" + }, + "source": [ + "First, it's useful to visualize and play around with attention patterns - what exactly are we looking at here? (Click on a head to lock onto just showing that head's pattern, it'll make it easier to interpret)" + ] + }, + { + "cell_type": "code", + "execution_count": 176, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/", + "height": 357 + }, + "id": "bLBDFvd7fdfP", + "outputId": "98ebfe3d-5b2e-4589-9100-b0f0c6f75a27" + }, + "outputs": [ + { + "output_type": "display_data", + "data": { + "text/plain": [ + "" + ], + "text/html": [ + "
\n", + " " + ] + }, + "metadata": {} + } + ], + "source": [ + "import circuitsvis as cv\n", + "from IPython.display import display\n", + "\n", + "display(\n", + " cv.attention.attention_patterns(\n", + " tokens=reference_gpt2.to_str_tokens(reference_text), attention=cache[\"pattern\", 0][0]\n", + " )\n", + ")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "0h4_A67JfdfP" + }, + "source": [ + "You can also use the `attention_heads` function, which presents the data in a different way (the syntax is exactly the same as `attention_patterns`). Note, if you display this in VSCode then it may exhibit a bug where the main plot continually shrinks in size - if this happens, you should instead save the HTML (i.e. with `html = cv.attention.attention_heads(...); with open(\"attn_heads.html\", \"w\") as f: f.write(str(html))`) and open the plot in your browser.\n", + "\n", + "" + ] + }, + { + "cell_type": "code", + "execution_count": 177, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/", + "height": 386 + }, + "id": "9154K71bfdfP", + "outputId": "09f7464b-22b7-4c22-d5fc-ef6492e9a997" + }, + "outputs": [ + { + "output_type": "display_data", + "data": { + "text/plain": [ + "" + ], + "text/html": [ + "
\n", + " " + ] + }, + "metadata": {} + } + ], + "source": [ + "display(\n", + " cv.attention.attention_heads(\n", + " tokens=reference_gpt2.to_str_tokens(reference_text), attention=cache[\"pattern\", 0][0]\n", + " )\n", + ")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "n468cOWufdfP" + }, + "source": [ + "You should fill in the forward method for `Attention` below. You should also copy your code for `apply_causal_mask` to this new implementation of `Attention` (you can delete the rest of the old implementation code).\n", + "\n", + "Note, this implementation will probably be the most challenging exercise on this page, so don't worry if it takes you some time! You should look at parts of the solution if you're stuck. A few tips:\n", + "\n", + "* Don't forget the attention score scaling (this should come before the masking).\n", + "* Try not to combine a large number of operations into a single line of code.\n", + "* Try to make your variable names descriptive (i.e. it's not just `x = some_fn_of(x), x = some_other_fn_of(x), ...`)." + ] + }, + { + "cell_type": "code", + "execution_count": 178, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "9E8dKXz3fdfP", + "outputId": "3d1a2a31-af59-4936-9b2a-c25f34b3f640" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "All tests in `test_causal_mask` passed!\n", + "Input shape: torch.Size([2, 4, 768])\n", + "torch.Size([2, 4, 768])\n", + "torch.Size([12, 768, 64])\n", + "torch.Size([12, 768, 64])\n", + "torch.Size([12, 768, 64])\n", + "Output shape: torch.Size([2, 4, 768]) \n", + "\n", + "Input shape: torch.Size([1, 35, 768])\n", + "torch.Size([1, 35, 768])\n", + "torch.Size([12, 768, 64])\n", + "torch.Size([12, 768, 64])\n", + "torch.Size([12, 768, 64])\n", + "Output shape: torch.Size([1, 35, 768])\n", + "Reference output shape: torch.Size([1, 35, 768]) \n", + "\n", + "100.00% of the values are correct\n", + "\n" + ] + } + ], + "source": [ + "class Attention(nn.Module):\n", + " IGNORE: Float[Tensor, \"\"]\n", + "\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.W_Q = nn.Parameter(t.empty((cfg.n_heads, cfg.d_model, cfg.d_head)))\n", + " self.W_K = nn.Parameter(t.empty((cfg.n_heads, cfg.d_model, cfg.d_head)))\n", + " self.W_V = nn.Parameter(t.empty((cfg.n_heads, cfg.d_model, cfg.d_head)))\n", + " self.W_O = nn.Parameter(t.empty((cfg.n_heads, cfg.d_head, cfg.d_model)))\n", + " self.b_Q = nn.Parameter(t.zeros((cfg.n_heads, cfg.d_head)))\n", + " self.b_K = nn.Parameter(t.zeros((cfg.n_heads, cfg.d_head)))\n", + " self.b_V = nn.Parameter(t.zeros((cfg.n_heads, cfg.d_head)))\n", + " self.b_O = nn.Parameter(t.zeros((cfg.d_model)))\n", + " nn.init.normal_(self.W_Q, std=self.cfg.init_range)\n", + " nn.init.normal_(self.W_K, std=self.cfg.init_range)\n", + " nn.init.normal_(self.W_V, std=self.cfg.init_range)\n", + " nn.init.normal_(self.W_O, std=self.cfg.init_range)\n", + " self.register_buffer(\"IGNORE\", t.tensor(float(\"-inf\"), dtype=t.float32, device=device))\n", + "\n", + " def forward(self, r: Float[Tensor, \"batch posn d_model\"]) -> Float[Tensor, \"batch posn d_model\"]:\n", + " print(r.shape)\n", + " print(self.W_Q.shape)\n", + " print(self.W_K.shape)\n", + " print(self.W_V.shape)\n", + " q = t.matmul(r.unsqueeze(1), self.W_Q) + self.b_Q.unsqueeze(1)\n", + " k = t.matmul(r.unsqueeze(1), self.W_K) + self.b_K.unsqueeze(1)\n", + " v = t.matmul(r.unsqueeze(1), self.W_V) + self.b_V.unsqueeze(1)\n", + "\n", + " x = t.matmul(q, k.transpose(-2, -1))\n", + "\n", + " ascore = self.apply_causal_mask(\n", + " x / math.sqrt(self.cfg.d_head)\n", + " ).softmax(dim=-1)\n", + "\n", + " weighted_avg = t.matmul(ascore, v)\n", + "\n", + " return t.matmul(weighted_avg, self.W_O).sum(dim=1) + self.b_O\n", + "\n", + " def apply_causal_mask(\n", + " self, attn_scores: Float[Tensor, \"batch n_heads query_pos key_pos\"]\n", + " ) -> Float[Tensor, \"batch n_heads query_pos key_pos\"]:\n", + " \"\"\"\n", + " Applies a causal mask to attention scores, and returns masked scores.\n", + " \"\"\"\n", + " # You should copy your solution from earlier\n", + " k = t.arange(attn_scores.size(3))[None, :]\n", + " q = t.arange(attn_scores.size(2))[:, None]\n", + " m = k > q\n", + "\n", + " return t.where(\n", + " m[None, None, :, :],\n", + " float(\"-inf\"),\n", + " attn_scores\n", + " )\n", + "\n", + "\n", + "tests.test_causal_mask(Attention.apply_causal_mask)\n", + "tests.rand_float_test(Attention, [2, 4, 768])\n", + "tests.load_gpt2_test(Attention, reference_gpt2.blocks[0].attn, cache[\"normalized\", 0, \"ln1\"])" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "lY8l9ZWbfdfP" + }, + "source": [ + "
\n", + "Hint (pseudocode for the forward method)\n", + "\n", + "```python\n", + "def forward(\n", + " self, normalized_resid_pre: Float[Tensor, \"batch posn d_model\"]\n", + ") -> Float[Tensor, \"batch posn d_model\"]:\n", + "\n", + " # Calculate query, key and value vectors\n", + " q, k, v = ...\n", + "\n", + " # Calculate attention scores, then scale and mask, and apply softmax to get probabilities\n", + " attn_scores = ...\n", + " attn_scores_masked = ...\n", + " attn_pattern = ...\n", + "\n", + " # Take weighted sum of value vectors, according to attention probabilities\n", + " z = ...\n", + "\n", + " # Calculate output (by applying matrix W_O and summing over heads, then adding bias b_O)\n", + " attn_out = ...\n", + " return attn_out\n", + "```\n", + "
\n", + "\n", + "\n", + "
Solution\n", + "\n", + "```python\n", + "class Attention(nn.Module):\n", + " IGNORE: Float[Tensor, \"\"]\n", + "\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.W_Q = nn.Parameter(t.empty((cfg.n_heads, cfg.d_model, cfg.d_head)))\n", + " self.W_K = nn.Parameter(t.empty((cfg.n_heads, cfg.d_model, cfg.d_head)))\n", + " self.W_V = nn.Parameter(t.empty((cfg.n_heads, cfg.d_model, cfg.d_head)))\n", + " self.W_O = nn.Parameter(t.empty((cfg.n_heads, cfg.d_head, cfg.d_model)))\n", + " self.b_Q = nn.Parameter(t.zeros((cfg.n_heads, cfg.d_head)))\n", + " self.b_K = nn.Parameter(t.zeros((cfg.n_heads, cfg.d_head)))\n", + " self.b_V = nn.Parameter(t.zeros((cfg.n_heads, cfg.d_head)))\n", + " self.b_O = nn.Parameter(t.zeros((cfg.d_model)))\n", + " nn.init.normal_(self.W_Q, std=self.cfg.init_range)\n", + " nn.init.normal_(self.W_K, std=self.cfg.init_range)\n", + " nn.init.normal_(self.W_V, std=self.cfg.init_range)\n", + " nn.init.normal_(self.W_O, std=self.cfg.init_range)\n", + " self.register_buffer(\"IGNORE\", t.tensor(float(\"-inf\"), dtype=t.float32, device=device))\n", + "\n", + " def forward(self, normalized_resid_pre: Float[Tensor, \"batch posn d_model\"]) -> Float[Tensor, \"batch posn d_model\"]:\n", + " # Calculate query, key and value vectors\n", + " q = (\n", + " einops.einsum(\n", + " normalized_resid_pre,\n", + " self.W_Q,\n", + " \"batch posn d_model, nheads d_model d_head -> batch posn nheads d_head\",\n", + " )\n", + " + self.b_Q\n", + " )\n", + " k = (\n", + " einops.einsum(\n", + " normalized_resid_pre,\n", + " self.W_K,\n", + " \"batch posn d_model, nheads d_model d_head -> batch posn nheads d_head\",\n", + " )\n", + " + self.b_K\n", + " )\n", + " v = (\n", + " einops.einsum(\n", + " normalized_resid_pre,\n", + " self.W_V,\n", + " \"batch posn d_model, nheads d_model d_head -> batch posn nheads d_head\",\n", + " )\n", + " + self.b_V\n", + " )\n", + "\n", + " # Calculate attention scores, then scale and mask, and apply softmax to get probabilities\n", + " attn_scores = einops.einsum(\n", + " q,\n", + " k,\n", + " \"batch posn_Q nheads d_head, batch posn_K nheads d_head -> batch nheads posn_Q posn_K\",\n", + " )\n", + " attn_scores_masked = self.apply_causal_mask(attn_scores / self.cfg.d_head**0.5)\n", + " attn_pattern = attn_scores_masked.softmax(-1)\n", + "\n", + " # Take weighted sum of value vectors, according to attention probabilities\n", + " z = einops.einsum(\n", + " v,\n", + " attn_pattern,\n", + " \"batch posn_K nheads d_head, batch nheads posn_Q posn_K -> batch posn_Q nheads d_head\",\n", + " )\n", + "\n", + " # Calculate output (by applying matrix W_O and summing over heads, then adding bias b_O)\n", + " attn_out = (\n", + " einops.einsum(\n", + " z,\n", + " self.W_O,\n", + " \"batch posn_Q nheads d_head, nheads d_head d_model -> batch posn_Q d_model\",\n", + " )\n", + " + self.b_O\n", + " )\n", + "\n", + " return attn_out\n", + "\n", + " def apply_causal_mask(\n", + " self, attn_scores: Float[Tensor, \"batch n_heads query_pos key_pos\"]\n", + " ) -> Float[Tensor, \"batch n_heads query_pos key_pos\"]:\n", + " \"\"\"\n", + " Applies a causal mask to attention scores, and returns masked scores.\n", + " \"\"\"\n", + " # Define a mask that is True for all positions we want to set probabilities to zero for\n", + " all_ones = t.ones(attn_scores.size(-2), attn_scores.size(-1), device=attn_scores.device)\n", + " mask = t.triu(all_ones, diagonal=1).bool()\n", + " # Apply the mask to attention scores, then return the masked scores\n", + " attn_scores.masked_fill_(mask, self.IGNORE)\n", + " return attn_scores\n", + "```\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "Oa9QZ3KMfdfP" + }, + "source": [ + "### Exercise - implement `MLP`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴⚪⚪⚪\n", + "> Importance: 🔵🔵🔵🔵⚪\n", + ">\n", + "> You should spend up to 10-15 minutes on this exercise.\n", + "> ```\n", + "\n", + "Next, you should implement the MLP layer, which consists of:\n", + "\n", + "* A linear layer, with weight `W_in`, bias `b_in`\n", + "* A nonlinear function (we usually use GELU; the function `gelu_new` has been imported for this purpose)\n", + "* A linear layer, with weight `W_out`, bias `b_out`" + ] + }, + { + "cell_type": "code", + "execution_count": 197, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "YDJtDZDyfdfP", + "outputId": "09ed0d0d-bf42-45ac-f357-b7216bc13d7d" + }, + "outputs": [ + { + "output_type": "stream", + "name": "stdout", + "text": [ + "Input shape: torch.Size([2, 4, 768])\n", + "Output shape: torch.Size([2, 4, 768]) \n", + "\n", + "Input shape: torch.Size([1, 35, 768])\n", + "Output shape: torch.Size([1, 35, 768])\n", + "Reference output shape: torch.Size([1, 35, 768]) \n", + "\n", + "100.00% of the values are correct\n", + "\n" + ] + } + ], + "source": [ + "class MLP(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.W_in = nn.Parameter(t.empty((cfg.d_model, cfg.d_mlp)))\n", + " self.W_out = nn.Parameter(t.empty((cfg.d_mlp, cfg.d_model)))\n", + " self.b_in = nn.Parameter(t.zeros((cfg.d_mlp)))\n", + " self.b_out = nn.Parameter(t.zeros((cfg.d_model)))\n", + " nn.init.normal_(self.W_in, std=self.cfg.init_range)\n", + " nn.init.normal_(self.W_out, std=self.cfg.init_range)\n", + "\n", + " def forward(self, normalized_resid_mid: Float[Tensor, \"batch posn d_model\"]) -> Float[Tensor, \"batch posn d_model\"]:\n", + " l1 = t.matmul(normalized_resid_mid, self.W_in) + self.b_in\n", + " acti = gelu_new(l1)\n", + " l2 = t.matmul(acti, self.W_out) + self.b_out\n", + " return l2\n", + "\n", + "\n", + "tests.rand_float_test(MLP, [2, 4, 768])\n", + "tests.load_gpt2_test(MLP, reference_gpt2.blocks[0].mlp, cache[\"normalized\", 0, \"ln2\"])" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "a-znDk_ifdfP" + }, + "source": [ + "
Solution\n", + "\n", + "```python\n", + "class MLP(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.W_in = nn.Parameter(t.empty((cfg.d_model, cfg.d_mlp)))\n", + " self.W_out = nn.Parameter(t.empty((cfg.d_mlp, cfg.d_model)))\n", + " self.b_in = nn.Parameter(t.zeros((cfg.d_mlp)))\n", + " self.b_out = nn.Parameter(t.zeros((cfg.d_model)))\n", + " nn.init.normal_(self.W_in, std=self.cfg.init_range)\n", + " nn.init.normal_(self.W_out, std=self.cfg.init_range)\n", + "\n", + " def forward(self, normalized_resid_mid: Float[Tensor, \"batch posn d_model\"]) -> Float[Tensor, \"batch posn d_model\"]:\n", + " pre = (\n", + " einops.einsum(\n", + " normalized_resid_mid,\n", + " self.W_in,\n", + " \"batch position d_model, d_model d_mlp -> batch position d_mlp\",\n", + " )\n", + " + self.b_in\n", + " )\n", + " post = gelu_new(pre)\n", + " mlp_out = (\n", + " einops.einsum(post, self.W_out, \"batch position d_mlp, d_mlp d_model -> batch position d_model\")\n", + " + self.b_out\n", + " )\n", + " return mlp_out\n", + "```\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "R1g466wafdfP" + }, + "source": [ + "### Exercise - implement `TransformerBlock`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴⚪⚪⚪\n", + "> Importance: 🔵🔵🔵⚪⚪\n", + ">\n", + "> You should spend up to 10-15 minutes on this exercise.\n", + "> ```\n", + "\n", + "Now, we can put together the attention, MLP and layernorms into a single transformer block. Remember to implement the residual connections correctly!" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "4Qwi3GknfdfP" + }, + "outputs": [], + "source": [ + "class TransformerBlock(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.ln1 = LayerNorm(cfg)\n", + " self.attn = Attention(cfg)\n", + " self.ln2 = LayerNorm(cfg)\n", + " self.mlp = MLP(cfg)\n", + "\n", + " def forward(self, resid_pre: Float[Tensor, \"batch position d_model\"]) -> Float[Tensor, \"batch position d_model\"]:\n", + " raise NotImplementedError()\n", + "\n", + "\n", + "tests.rand_float_test(TransformerBlock, [2, 4, 768])\n", + "tests.load_gpt2_test(TransformerBlock, reference_gpt2.blocks[0], cache[\"resid_pre\", 0])" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "gOMY4f3afdfP" + }, + "source": [ + "
\n", + "Help - I'm getting 100% accuracy on all modules before this point, but only about 90% accuracy on this one.\n", + "\n", + "This might be because your layernorm implementation divides by `std + eps` rather than `(var + eps).sqrt()`. The latter matches the implementation used by GPT-2 (and this error only shows up in these tests).\n", + "\n", + "
\n", + "\n", + "\n", + "
Solution\n", + "\n", + "```python\n", + "class TransformerBlock(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.ln1 = LayerNorm(cfg)\n", + " self.attn = Attention(cfg)\n", + " self.ln2 = LayerNorm(cfg)\n", + " self.mlp = MLP(cfg)\n", + "\n", + " def forward(self, resid_pre: Float[Tensor, \"batch position d_model\"]) -> Float[Tensor, \"batch position d_model\"]:\n", + " resid_mid = self.attn(self.ln1(resid_pre)) + resid_pre\n", + " resid_post = self.mlp(self.ln2(resid_mid)) + resid_mid\n", + " return resid_post\n", + "```\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "JyUtV_mtfdfQ" + }, + "source": [ + "### Exercise - implement `Unembed`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴⚪⚪⚪\n", + "> Importance: 🔵🔵🔵⚪⚪\n", + ">\n", + "> You should spend up to ~10 minutes on this exercise.\n", + "> ```\n", + "\n", + "The unembedding is just a linear layer (with weight `W_U` and bias `b_U`)." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "Bhco9htwfdfQ" + }, + "outputs": [], + "source": [ + "class Unembed(nn.Module):\n", + " def __init__(self, cfg):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.W_U = nn.Parameter(t.empty((cfg.d_model, cfg.d_vocab)))\n", + " nn.init.normal_(self.W_U, std=self.cfg.init_range)\n", + " self.b_U = nn.Parameter(t.zeros((cfg.d_vocab)), requires_grad=False)\n", + "\n", + " def forward(\n", + " self, normalized_resid_final: Float[Tensor, \"batch position d_model\"]\n", + " ) -> Float[Tensor, \"batch position d_vocab\"]:\n", + " raise NotImplementedError()\n", + "\n", + "\n", + "tests.test_unembed(Unembed)\n", + "tests.rand_float_test(Unembed, [2, 4, 768])\n", + "tests.load_gpt2_test(Unembed, reference_gpt2.unembed, cache[\"ln_final.hook_normalized\"])" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "n136eaYufdfQ" + }, + "source": [ + "
\n", + "Why have `requires_grad=False` for the bias? \n", + "\n", + "GPT-2 doesn't have a bias in the unembedding layer, but the TransformerLens copy\n", + "of GPT-2 that we're loading in has a uniform format for the unembedding layer\n", + "where the bias for GPT-2 is set to zero. We preserve the same structure here so\n", + "we don't get any mismatch when we load in the pre-trained weights.\n", + "Neel Nanda explains this in the [original video](https://youtu.be/dsjUDacBw8o?si=pP31Srg9XzQ5A-yj&t=3437) on which this material was based.\n", + "\n", + "For now, we will write the Unembed as if it has a bias, and the test will check\n", + "that you added the bias, for full generality.\n", + "\n", + "
\n", + "\n", + "\n", + "
Solution\n", + "\n", + "```python\n", + "class Unembed(nn.Module):\n", + " def __init__(self, cfg):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.W_U = nn.Parameter(t.empty((cfg.d_model, cfg.d_vocab)))\n", + " nn.init.normal_(self.W_U, std=self.cfg.init_range)\n", + " self.b_U = nn.Parameter(t.zeros((cfg.d_vocab)), requires_grad=False)\n", + "\n", + " def forward(\n", + " self, normalized_resid_final: Float[Tensor, \"batch position d_model\"]\n", + " ) -> Float[Tensor, \"batch position d_vocab\"]:\n", + " return (\n", + " einops.einsum(\n", + " normalized_resid_final,\n", + " self.W_U,\n", + " \"batch posn d_model, d_model d_vocab -> batch posn d_vocab\",\n", + " )\n", + " + self.b_U\n", + " )\n", + "```\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "6Zo0eTbLfdfQ" + }, + "source": [ + "### Exercise - implement `DemoTransformer`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴⚪⚪⚪\n", + "> Importance: 🔵🔵🔵⚪⚪\n", + ">\n", + "> You should spend up to 10-15 minutes on this exercise.\n", + "> ```" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "hyzYZjZhfdfQ" + }, + "outputs": [], + "source": [ + "class DemoTransformer(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.embed = Embed(cfg)\n", + " self.pos_embed = PosEmbed(cfg)\n", + " self.blocks = nn.ModuleList([TransformerBlock(cfg) for _ in range(cfg.n_layers)])\n", + " self.ln_final = LayerNorm(cfg)\n", + " self.unembed = Unembed(cfg)\n", + "\n", + " def forward(self, tokens: Int[Tensor, \"batch position\"]) -> Float[Tensor, \"batch position d_vocab\"]:\n", + " raise NotImplementedError()\n", + "\n", + "\n", + "tests.rand_int_test(DemoTransformer, [2, 4])\n", + "tests.load_gpt2_test(DemoTransformer, reference_gpt2, tokens)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "WI2G-SilfdfQ" + }, + "source": [ + "
Solution\n", + "\n", + "```python\n", + "class DemoTransformer(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.embed = Embed(cfg)\n", + " self.pos_embed = PosEmbed(cfg)\n", + " self.blocks = nn.ModuleList([TransformerBlock(cfg) for _ in range(cfg.n_layers)])\n", + " self.ln_final = LayerNorm(cfg)\n", + " self.unembed = Unembed(cfg)\n", + "\n", + " def forward(self, tokens: Int[Tensor, \"batch position\"]) -> Float[Tensor, \"batch position d_vocab\"]:\n", + " residual = self.embed(tokens) + self.pos_embed(tokens)\n", + " for block in self.blocks:\n", + " residual = block(residual)\n", + " logits = self.unembed(self.ln_final(residual))\n", + " return logits\n", + "```\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "wAcpYaHufdfQ" + }, + "source": [ + "**Try it out!**" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "jhEVY-tafdfQ" + }, + "outputs": [], + "source": [ + "demo_gpt2 = DemoTransformer(Config(debug=False)).to(device)\n", + "demo_gpt2.load_state_dict(reference_gpt2.state_dict(), strict=False)\n", + "\n", + "demo_logits = demo_gpt2(tokens)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "mi-Ow0mIfdfQ" + }, + "source": [ + "Let's take a test string, and calculate the loss!\n", + "\n", + "We're using the formula for **cross-entropy loss**. The cross entropy loss between a modelled distribution $Q$ and target distribution $P$ is:\n", + "\n", + "$$\n", + "-\\sum_x P(x) \\log Q(x)\n", + "$$\n", + "\n", + "In the case where $P$ is just the empirical distribution from target classes (i.e. $P(x^*) = 1$ for the correct class $x^*$) then this becomes:\n", + "\n", + "$$\n", + "-\\log Q(x^*)\n", + "$$\n", + "\n", + "in other words, the negative log prob of the true classification." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "1GlGzmXdfdfQ" + }, + "outputs": [], + "source": [ + "def get_log_probs(\n", + " logits: Float[Tensor, \"batch posn d_vocab\"], tokens: Int[Tensor, \"batch posn\"]\n", + ") -> Float[Tensor, \"batch posn-1\"]:\n", + " log_probs = logits.log_softmax(dim=-1)\n", + " # Get logprobs for the first seq_len-1 predictions (so we can compare them with the actual next tokens)\n", + " log_probs_for_tokens = log_probs[:, :-1].gather(dim=-1, index=tokens[:, 1:].unsqueeze(-1)).squeeze(-1)\n", + "\n", + " return log_probs_for_tokens\n", + "\n", + "\n", + "pred_log_probs = get_log_probs(demo_logits, tokens)\n", + "print(f\"Avg cross entropy loss: {-pred_log_probs.mean():.4f}\")\n", + "print(f\"Avg cross entropy loss for uniform distribution: {math.log(demo_gpt2.cfg.d_vocab):4f}\")\n", + "print(f\"Avg probability assigned to correct token: {pred_log_probs.exp().mean():4f}\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "Jcg2Ea7vfdfQ" + }, + "source": [ + "We can also greedily generate text, by taking the most likely next token and continually appending it to our prompt before feeding it back into the model:" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "oajhdbIqfdfQ" + }, + "outputs": [], + "source": [ + "test_string = \"\"\"Mitigating the risk of extinction from AI should be a global priority alongside other societal-scale risks such as\"\"\"\n", + "for i in tqdm(range(100)):\n", + " test_tokens = reference_gpt2.to_tokens(test_string).to(device)\n", + " demo_logits = demo_gpt2(test_tokens)\n", + " test_string += reference_gpt2.tokenizer.decode(demo_logits[-1, -1].argmax())\n", + "\n", + "print(test_string)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "4piWjeMHfdfQ" + }, + "source": [ + "In section 4️⃣ we'll learn to generate text in slightly more interesting ways than just argmaxing the output (which can lead to unnatural patterns like repetition, or text which is just less natural-sounding)." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "Ju07csQ9fdfQ" + }, + "source": [ + "# 3️⃣ Training a Transformer\n", + "\n", + "> ##### Learning Objectives\n", + ">\n", + "> * Understand how to train a transformer from scratch\n", + "> * Write a basic transformer training loop\n", + "> * Interpret the transformer's falling cross entropy loss with reference to features of the training data (e.g. bigram frequencies)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "hlgX6R4JfdfQ" + }, + "source": [ + "Now that we've built our transformer, and verified that it performs as expected when we load in weights, let's try training it from scratch!\n", + "\n", + "This is a lightweight demonstration of how you can actually train your own GPT-2 with this code! Here we train a tiny model on a tiny dataset, but it's fundamentally the same code for training a larger/more real model (though you'll need beefier GPUs and data parallelism to do it remotely efficiently, and fancier parallelism for much bigger ones).\n", + "\n", + "For our purposes, we'll train a 4 layer model with 16 heads per layer, with context length 128, for 10*500 steps of batch size 32, just to show what it looks like (and so the notebook doesn't melt your colab / machine!)." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "9B_wRH9afdfQ" + }, + "source": [ + "## Create Model" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "-E057frYfdfQ" + }, + "outputs": [], + "source": [ + "model_cfg = Config(\n", + " debug=False,\n", + " d_model=32,\n", + " n_heads=16,\n", + " d_head=2,\n", + " d_mlp=32 * 4,\n", + " n_layers=4,\n", + " n_ctx=128,\n", + " d_vocab=reference_gpt2.cfg.d_vocab,\n", + ")\n", + "model = DemoTransformer(model_cfg)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "Nl3xruqDfdfQ" + }, + "source": [ + "## Training Args\n", + "\n", + "\n", + "Note, for this optimization we'll be using **weight decay**." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "-oQXZp6WfdfQ" + }, + "outputs": [], + "source": [ + "@dataclass\n", + "class TransformerTrainingArgs:\n", + " batch_size: int = 32\n", + " epochs: int = 10\n", + " max_steps_per_epoch: int = 500\n", + " lr: float = 1e-3\n", + " weight_decay: float = 1e-2\n", + " wandb_project: str | None = \"day1-demotransformer\"\n", + " wandb_name: str | None = None\n", + " eval_prompt: str = \"Once upon a time\"\n", + "\n", + "\n", + "args = TransformerTrainingArgs()" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "v3wINJExfdfQ" + }, + "source": [ + "## Create Data\n", + "\n", + "We load in the [TinyStories dataset](https://huggingface.co/datasets/roneneldan/TinyStories), a dataset of synthetically generated simple stories only using a small vocabulary of words that typical 3 to 4-year-olds can understand. This dataset was designed for [exploring how small an LLM can be](https://arxiv.org/pdf/2305.07759) that can still generate coherent text." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "Ul7CiMm7fdfQ" + }, + "outputs": [], + "source": [ + "dataset = datasets.load_dataset(\"roneneldan/TinyStories\", split=\"train\")\n", + "print(dataset)\n", + "print(dataset[0][\"text\"])" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "UmTGBsZxfdfQ" + }, + "source": [ + "`tokenize_and_concatenate` is a useful function which takes our dataset of strings, and returns a dataset of token IDs ready to feed into the model. We then create a dataloader from this tokenized dataset. The useful method `train_test_split` can give us a training and testing set." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "VGh0W8GpfdfQ" + }, + "outputs": [], + "source": [ + "tokenized_dataset = tokenize_and_concatenate(\n", + " dataset,\n", + " reference_gpt2.tokenizer,\n", + " streaming=False,\n", + " max_length=model.cfg.n_ctx,\n", + " column_name=\"text\",\n", + " add_bos_token=True,\n", + " num_proc=8,\n", + ")\n", + "\n", + "dataset_dict = tokenized_dataset.train_test_split(test_size=1000)\n", + "train_loader = DataLoader(\n", + " dataset_dict[\"train\"],\n", + " batch_size=args.batch_size,\n", + " shuffle=True,\n", + " #num_workers=4, runs faster, but kills the kernel if interrupted\n", + " #pin_memory=False\n", + ")\n", + "test_loader = DataLoader(\n", + " dataset_dict[\"test\"],\n", + " batch_size=args.batch_size,\n", + " shuffle=False,\n", + " #num_workers=4,\n", + " #pin_memory=False\n", + ")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "dtShQKN0fdfQ" + }, + "source": [ + "When we iterate through these dataloaders, we will find dictionaries with the single key `'tokens'`, which maps to a tensor of token IDs with shape `(batch, seq_len)`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "Cv6sEPWYfdfQ" + }, + "outputs": [], + "source": [ + "first_batch = train_loader.dataset[: args.batch_size]\n", + "\n", + "print(first_batch.keys())\n", + "print(first_batch[\"tokens\"].shape)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "LBArCkuSfdfQ" + }, + "source": [ + "## Training Loop\n", + "\n", + "If you did the material on [training loops](https://learn.arena.education/chapter0_fundamentals/02_cnns/#2-training-neural-networks) during the first week, this should all be familiar to you. If not, you can skim that section for an overview of the key concepts. The start of the **Training loop** section is most important, and the subsections on [Modularisation](https://learn.arena.education/chapter0_fundamentals/02_cnns/#1-making-your-own-modules) and [dataclasses](https://learn.arena.education/chapter0_fundamentals/02_cnns/#2-training-neural-networks) are also very useful. Lastly, we'll also be using Weights and Biases to train our model - you can read about how to use it [here](https://learn.arena.education/chapter0_fundamentals/03_optimization/#2-weights-and-biases). Here are (roughly) all the things you should know for the following exercises:\n", + " \n", + "* The key parts of a gradient update step are:\n", + " * Calculating the (cross-entropy) loss between a model's output and the true labels,\n", + " * `loss.backward()` - calculate gradients of the loss with respect to the model parameters,\n", + " * `optimizer.step()` - update the model parameters using the gradients,\n", + " * `optimizer.zero_grad()` - zero the gradients so they don't accumulate.\n", + "* We can nicely package up training loops into a class, which includes methods for training and validation steps among other things. This helps with writing code that can be reused in different contexts.\n", + "* We can use dataclasses to store all the arguments relevant to training in one place, and then pass them to our trainer class. Autocompletion is one nice bonus of this!\n", + " * Be careful of scope here, you want to make sure you're referring to `self.args` within the trainer class, rather than the global `args`.\n", + "* You can use Weights and Biases to track experiments and log relevant variables. The three essential functions are:\n", + " * `wandb.init()` - initialize a new run, takes arguments `project`, `name` and `config` (among others).\n", + " * `wandb.log()` - log a dictionary of variables, e.g. `{\"loss\": loss}`. Also takes a `step` argument.\n", + " * `wandb.finish()` - called at the end of training (no arguments)." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "w-yYdUfcfdfR" + }, + "source": [ + "### Exercise - write training loop\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴🔴⚪⚪\n", + "> Importance: 🔵🔵🔵🔵⚪\n", + ">\n", + "> You should spend up to 10-20 minutes on this exercise.\n", + "> ```\n", + "\n", + "You should fill in the methods below. Some guidance:\n", + "\n", + "* Remember we were able to calculate cross entropy loss using the `get_log_probs` function in the previous section.\n", + "* You should use the optimizer `t.optim.AdamW` (Adam with weight decay), and with hyperparameters `lr` and `weight_decay` taken from your `TransformerTrainingArgs` dataclass instance.\n", + "* We've given you the argument `max_steps_per_epoch`, a hacky way of making sure the training phase in each epoch doesn't go on for too long. You can terminate each training phase after this many steps. It's set to a default value that should lead to a very short run demonstrating nontrivial model performance.\n", + "* Remember to move tokens to your device, via `tokens.to(device)` (this should be a global variable, defined at the top of your notebook).\n", + "* You can refer back to the training loops from the [previous chapter of the course](https://learn.arena.education/chapter0_fundamentals/02_cnns/#2-training-neural-networks) if you'd like.\n", + "* We've also provided an instance of the `TransformerSampler` class so you can generate text from your model during training to see how it's doing. We will cover how sampling works in the next section." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "z_GrJ1POfdfR" + }, + "outputs": [], + "source": [ + "class TransformerTrainer:\n", + " def __init__(self, args: TransformerTrainingArgs, model: DemoTransformer):\n", + " super().__init__()\n", + " self.model = model\n", + " self.args = args\n", + " self.sampler = solutions.TransformerSampler(self.model, reference_gpt2.tokenizer)\n", + " self.optimizer = t.optim.AdamW(self.model.parameters(), lr=args.lr, weight_decay=args.weight_decay)\n", + " self.step = 0\n", + "\n", + " self.train_loader = DataLoader(\n", + " dataset_dict[\"train\"],\n", + " batch_size=args.batch_size,\n", + " shuffle=True,\n", + " #num_workers=4, runs faster, but kills the kernel if interrupted\n", + " #pin_memory=False,\n", + " )\n", + " self.test_loader = DataLoader(\n", + " dataset_dict[\"test\"],\n", + " batch_size=args.batch_size,\n", + " shuffle=False,\n", + " #num_workers=4,\n", + " #pin_memory=False,\n", + " )\n", + "\n", + " def training_step(self, batch: dict[str, Int[Tensor, \"batch seq\"]]) -> Float[Tensor, \"\"]:\n", + " \"\"\"\n", + " Calculates the loss on the tokens in the batch, performs a gradient update step, and logs the loss.\n", + "\n", + " Remember that `batch` is a dictionary with the single key 'tokens'.\n", + " \"\"\"\n", + " raise NotImplementedError()\n", + " return loss\n", + "\n", + " @t.inference_mode()\n", + " def evaluate(self) -> float:\n", + " \"\"\"\n", + " Evaluate the model on the test set and return the accuracy.\n", + " \"\"\"\n", + " self.model.eval()\n", + " #\n", + " # YOUR CODE HERE - fill in the `evaluate` method\n", + " #\n", + " self.model.train()\n", + " return accuracy\n", + "\n", + " def train(self):\n", + " \"\"\"\n", + " Trains the model, for `self.args.epochs` epochs. Also handles wandb initialisation, and early stopping\n", + " for each epoch at `self.args.max_steps_per_epoch` steps.\n", + " \"\"\"\n", + " wandb.init(project=self.args.wandb_project, name=self.args.wandb_name, config=self.args)\n", + " accuracy = np.nan\n", + "\n", + " progress_bar = tqdm(total=self.args.max_steps_per_epoch * self.args.epochs)\n", + "\n", + " print(self.sampler.sample(self.args.eval_prompt, max_tokens_generated=50))\n", + " for epoch in range(self.args.epochs):\n", + " for i, batch in enumerate(self.train_loader):\n", + " loss = self.training_step(batch)\n", + " progress_bar.update()\n", + " progress_bar.set_description(f\"Epoch {epoch + 1}, loss: {loss:.3f}, accuracy: {accuracy:.3f}\")\n", + " if i >= self.args.max_steps_per_epoch:\n", + " break\n", + "\n", + " accuracy = self.evaluate()\n", + " print(self.sampler.sample(self.args.eval_prompt, max_tokens_generated=50))\n", + "\n", + " wandb.finish()\n", + "\n", + "\n", + "# See the full run here: https://api.wandb.ai/links/dquarel/nrxuwnv7\n", + "model = DemoTransformer(model_cfg).to(device)\n", + "args = TransformerTrainingArgs()\n", + "trainer = TransformerTrainer(args, model)\n", + "trainer.train()" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "Gpo69bvgfdfR" + }, + "source": [ + "
Solution\n", + "\n", + "```python\n", + "class TransformerTrainer:\n", + " def __init__(self, args: TransformerTrainingArgs, model: DemoTransformer):\n", + " super().__init__()\n", + " self.model = model\n", + " self.args = args\n", + " self.sampler = solutions.TransformerSampler(self.model, reference_gpt2.tokenizer)\n", + " self.optimizer = t.optim.AdamW(self.model.parameters(), lr=args.lr, weight_decay=args.weight_decay)\n", + " self.step = 0\n", + "\n", + " self.train_loader = DataLoader(\n", + " dataset_dict[\"train\"],\n", + " batch_size=args.batch_size,\n", + " shuffle=True,\n", + " #num_workers=4, runs faster, but kills the kernel if interrupted\n", + " #pin_memory=False,\n", + " )\n", + " self.test_loader = DataLoader(\n", + " dataset_dict[\"test\"],\n", + " batch_size=args.batch_size,\n", + " shuffle=False,\n", + " #num_workers=4,\n", + " #pin_memory=False,\n", + " )\n", + "\n", + " def training_step(self, batch: dict[str, Int[Tensor, \"batch seq\"]]) -> Float[Tensor, \"\"]:\n", + " \"\"\"\n", + " Calculates the loss on the tokens in the batch, performs a gradient update step, and logs the loss.\n", + "\n", + " Remember that `batch` is a dictionary with the single key 'tokens'.\n", + " \"\"\"\n", + " tokens = batch[\"tokens\"].to(device)\n", + " logits = self.model(tokens)\n", + " loss = -get_log_probs(logits, tokens).mean()\n", + " loss.backward()\n", + " self.optimizer.step()\n", + " self.optimizer.zero_grad()\n", + " self.step += 1\n", + " wandb.log({\"train_loss\": loss}, step=self.step)\n", + " return loss\n", + "\n", + " @t.inference_mode()\n", + " def evaluate(self) -> float:\n", + " \"\"\"\n", + " Evaluate the model on the test set and return the accuracy.\n", + " \"\"\"\n", + " self.model.eval()\n", + " total_correct, total_samples = 0, 0\n", + "\n", + " for batch in tqdm(self.test_loader, desc=\"Evaluating\"):\n", + " tokens = batch[\"tokens\"].to(device)\n", + " logits: Tensor = self.model(tokens)[:, :-1]\n", + " predicted_tokens = logits.argmax(dim=-1)\n", + " total_correct += (predicted_tokens == tokens[:, 1:]).sum().item()\n", + " total_samples += tokens.size(0) * (tokens.size(1) - 1)\n", + "\n", + " accuracy = total_correct / total_samples\n", + " wandb.log({\"accuracy\": accuracy}, step=self.step)\n", + " self.model.train()\n", + " return accuracy\n", + "\n", + " def train(self):\n", + " \"\"\"\n", + " Trains the model, for `self.args.epochs` epochs. Also handles wandb initialisation, and early stopping\n", + " for each epoch at `self.args.max_steps_per_epoch` steps.\n", + " \"\"\"\n", + " wandb.init(project=self.args.wandb_project, name=self.args.wandb_name, config=self.args)\n", + " accuracy = np.nan\n", + "\n", + " progress_bar = tqdm(total=self.args.max_steps_per_epoch * self.args.epochs)\n", + "\n", + " print(self.sampler.sample(self.args.eval_prompt, max_tokens_generated=50))\n", + " for epoch in range(self.args.epochs):\n", + " for i, batch in enumerate(self.train_loader):\n", + " loss = self.training_step(batch)\n", + " progress_bar.update()\n", + " progress_bar.set_description(f\"Epoch {epoch + 1}, loss: {loss:.3f}, accuracy: {accuracy:.3f}\")\n", + " if i >= self.args.max_steps_per_epoch:\n", + " break\n", + "\n", + " accuracy = self.evaluate()\n", + " print(self.sampler.sample(self.args.eval_prompt, max_tokens_generated=50))\n", + "\n", + " wandb.finish()\n", + "```\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "u-Ciqj46fdfR" + }, + "source": [ + "\n", + "\n", + "When you run the code for the first time, you'll have to login to Weights and Biases, and paste an API key into VSCode. After this is done, your Weights and Biases training run will start. It'll give you a lot of output text, one line of which will look like:\n", + "\n", + "```\n", + "View run at https://wandb.ai///runs/\n", + "```\n", + "\n", + "which you can click on to visit the run page.\n", + "\n", + "> Note - to see the plots more clearly in Weights and Biases, you can click on the **edit panel** of your plot (the small pencil symbol at the top-right), then move the **smoothing** slider to the right." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "giqLPUDNfdfR" + }, + "source": [ + "### A note on this loss curve (optional)\n", + "\n", + "\n", + "What's up with the shape of our loss curve? It seems like we start at around 10-11, drop down very fast, but then level out. It turns out, this is all to do with the kinds of algorithms the model learns during training.\n", + "\n", + "When it starts out, your model will be outputting random noise, which might look a lot like \"predict each token with approximately uniform probability\", i.e. $Q(x) = 1/d_\\text{vocab}$ for all $x$. This gives us a cross entropy loss of $\\log (d_\\text{vocab})$." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "jTXmDSjkfdfR" + }, + "outputs": [], + "source": [ + "d_vocab = model.cfg.d_vocab\n", + "\n", + "print(f\"d_vocab = {d_vocab}\")\n", + "print(f\"Cross entropy loss on uniform distribution = {math.log(d_vocab):.3f}\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "ZdigFHhGfdfR" + }, + "source": [ + "The next thing we might expect the model to learn is the frequencies of words in the English language. After all, small common tokens like `\" and\"` or `\" the\"` might appear much more frequently than others. This would give us an average cross entropy loss of:\n", + "\n", + "$$\n", + "- \\sum_x p_x \\log p_x\n", + "$$\n", + "\n", + "where $p_x$ is the actual frequency of the word in our training data.\n", + "\n", + "We can evaluate this quantity as follows:" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "Ev6EZUJTfdfR" + }, + "outputs": [], + "source": [ + "toks = tokenized_dataset[:][\"tokens\"].flatten()\n", + "\n", + "d_vocab = model.cfg.d_vocab\n", + "freqs = t.bincount(toks, minlength=d_vocab)\n", + "probs = freqs.float() / freqs.sum()\n", + "\n", + "distn = t.distributions.categorical.Categorical(probs=probs)\n", + "entropy = distn.entropy()\n", + "\n", + "print(f\"Entropy of training data = {entropy:.3f}\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "n27SDo4GfdfR" + }, + "source": [ + "After unigram frequencies, the next thing our model usually learns is **bigram frequencies** (i.e. the frequency of pairs of adjacent tokens in the training data). For instance, `\"I\"` and `\" am\"` are common tokens, but their bigram frequency is much higher than it would be if they occurred independently. Bigram frequencies actually take you pretty far, since they also help with:\n", + "\n", + "* Some simple grammatical rules (e.g. a full stop being followed by a capitalized word)\n", + "* Weird quirks of tokenization (e.g. `\" manip\"` being followed by `\"ulative\"`)\n", + "* Common names (e.g. `\"Barack\"` being followed by `\" Obama\"`)\n", + "\n", + "\n", + "After approximating bigram frequencies, we need to start using smarter techniques, like trigrams (which can only be implemented using attention heads), **induction heads** (which we'll learn a lot more about in the next set of exercises!), and fact memorization or more basic grammar and syntax rules. Marginal improvements start getting harder around this point, leading to a flattening of our loss curve." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "itNObPfHfdfR" + }, + "source": [ + "### Exercise (optional) - log completions\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴🔴🔴⚪\n", + "> Importance: 🔵⚪⚪⚪⚪\n", + ">\n", + "> You should spend up to 20-40 minutes on this exercise, if you choose to attempt it.\n", + "> Note, you might want to come back to this exercise *after* you learn how sampling works.\n", + "> ```\n", + "\n", + "Choose a handful of prompts, and log the model's completions on those sentences. We recommend you do this with a lower frequency than loss is logged (e.g. once every 10-100 batches).\n", + "\n", + "The `wandb` syntax for logging text is pretty simple. Firstly, you can just print output as stdout and this is also logged to Weights & Biases (you can find it under the \"Logs\" section of your run). Alternatively, you can log data in the form of a table, and have it appear next to your other charts:\n", + "\n", + "```python\n", + "wandb.log({\"completions_table\": wandb.Table(\n", + " data = data,\n", + " columns = [\"epoch\", \"step\", \"text\"]\n", + ")})\n", + "```\n", + "\n", + "where `data` is a list of length-3 lists, with each list containing (epoch, step, text). If you choose this option, we recommend logging the table less frequently than you're sampling from the model, to make sure you're not sending too much data (because unfortunately wandb doesn't have methods to incrementally update the table during logging).\n", + "\n", + "If you want to try this before going through the sampling exercises (which are quite long!), you can use the code below to sample output from the model. Note that the `TransformerSampler` object is already in inference mode, so you don't need to worry about this." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "2jWJlcVRfdfR" + }, + "outputs": [], + "source": [ + "def sampling_fn(model: DemoTransformer, prompt: str) -> str:\n", + " sampler = solutions.TransformerSampler(model, reference_gpt2.tokenizer)\n", + " output = sampler.sample(prompt, temperature=0.7, top_p=0.95, max_tokens_generated=16)\n", + " return output\n", + "\n", + "\n", + "model = DemoTransformer(model_cfg).to(device)\n", + "\n", + "# Should be entirely random, because it uses a newly initialized model\n", + "print(sampling_fn(model, prompt=\"John and Mary went to the\"))" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "XU3oelcZfdfR" + }, + "outputs": [], + "source": [ + "# YOUR CODE HERE - rewrite the TransformerTrainer.train method, so that it logs completions\n", + "\n", + "\n", + "prompt_list = [\n", + " \"Eliezer Shlomo Yudkowsky (born September 11, 1979) is an American decision and artificial intelligence (AI) theorist and writer, best known for\",\n", + " \"In a shocking finding, scientist discovered a herd of unicorns living in a remote, previously unexplored valley, in the Andes Mountains. Even more surprising to the researchers was the fact that the unicorns spoke perfect English.\",\n", + " \"John and Mary went to the\",\n", + "]\n", + "\n", + "model = DemoTransformer(model_cfg).to(device)\n", + "args = TransformerTrainingArgsLogText()\n", + "trainer = TransformerTrainer(args, model)\n", + "trainer.train(sampling_fn, prompt_list)\n", + "# Read full report here - https://api.wandb.ai/links/callum-mcdougall/5ex16e5w" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "bWNb0Mk3fdfR" + }, + "source": [ + "
Solution\n", + "\n", + "```python\n", + "@dataclass\n", + "class TransformerTrainingArgsLogText(TransformerTrainingArgs):\n", + " text_sample_freq: int = 20\n", + " table_log_freq: int = 200\n", + "\n", + " def __post_init__(self):\n", + " assert self.table_log_freq >= self.text_sample_freq, (\n", + " \"You should log the table less frequently than you add text to it.\"\n", + " )\n", + "\n", + "\n", + "def train_log_text(self: TransformerTrainer, sampling_fn: Callable, prompt_list: list[str]):\n", + " \"\"\"\n", + " Trains the model, for `self.args.epochs` epochs. Also handles wandb initialisation, and early stopping\n", + " for each epoch at `self.args.max_steps_per_epoch` steps.\n", + "\n", + " This also takes 2 extra arguments:\n", + " sampling_fn: function which takes model & a single prompt (i.e. text string) and returns text string output\n", + " prompt_list: list of prompts we'll log output on\n", + " \"\"\"\n", + " wandb.init(project=self.args.wandb_project, name=self.args.wandb_name, config=self.args)\n", + " accuracy = np.nan\n", + " progress_bar = tqdm(total=self.args.max_steps_per_epoch * self.args.epochs)\n", + "\n", + " # Create a list for storing data\n", + " completions_list = []\n", + "\n", + " for epoch in range(self.args.epochs):\n", + " for i, batch in enumerate(self.train_loader):\n", + " loss = self.training_step(batch)\n", + " progress_bar.update()\n", + " progress_bar.set_description(f\"Epoch {epoch + 1}, loss: {loss:.3f}, accuracy: {accuracy:.3f}\")\n", + "\n", + " # Control the adding of text to the table, and the logging of text\n", + " if self.step % self.args.text_sample_freq == 0:\n", + " text_completions = [sampling_fn(self.model, prompt) for prompt in prompt_list]\n", + " completions_list.append([epoch, self.step, *text_completions])\n", + " if self.step % self.args.table_log_freq == 0:\n", + " wandb.log(\n", + " {\n", + " \"completions_table\": wandb.Table(\n", + " data=completions_list,\n", + " columns=[\n", + " \"epoch\",\n", + " \"step\",\n", + " *[f\"prompt_{i}\" for i in range(len(prompt_list))],\n", + " ],\n", + " )\n", + " }\n", + " )\n", + "\n", + " if i >= self.args.max_steps_per_epoch:\n", + " break\n", + "\n", + " accuracy = self.evaluate()\n", + "\n", + " wandb.finish()\n", + "\n", + "\n", + "TransformerTrainer.train = train_log_text\n", + "\n", + "\n", + "prompt_list = [\n", + " \"Eliezer Shlomo Yudkowsky (born September 11, 1979) is an American decision and artificial intelligence (AI) theorist and writer, best known for\",\n", + " \"In a shocking finding, scientist discovered a herd of unicorns living in a remote, previously unexplored valley, in the Andes Mountains. Even more surprising to the researchers was the fact that the unicorns spoke perfect English.\",\n", + " \"John and Mary went to the\",\n", + "]\n", + "\n", + "model = DemoTransformer(model_cfg).to(device)\n", + "args = TransformerTrainingArgsLogText()\n", + "trainer = TransformerTrainer(args, model)\n", + "trainer.train(sampling_fn, prompt_list)\n", + "# Read full report here - https://api.wandb.ai/links/callum-mcdougall/5ex16e5w\n", + "```\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "DNbSguJxfdfR" + }, + "source": [ + "You shouldn't expect to see perfect logical coherence from your model, but you should at least see that it respects basic word frequencies, and follows basic rules of grammar some of the time. Hopefully this gives some perspective on how difficult training a transformer can be!" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "PBBUdNbufdfR" + }, + "source": [ + "# 4️⃣ Sampling from a Transformer\n", + "\n", + "> ##### Learning Objectives\n", + ">\n", + "> * Learn how to sample from a transformer\n", + "> * This includes basic methods like greedy search or top-k, and more advanced methods like beam search\n", + "> * Learn how to cache the output of a transformer, so that it can be used to generate text more efficiently\n", + "> * Optionally, rewrite your sampling functions to make use of your caching methods" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "PVlMQOCwfdfR" + }, + "source": [ + "Let's discuss how we might go about producing output from a transformer.\n", + "\n", + "One obvious method to sample tokens from a distribution would be to always take the token assigned the highest probability. But this can lead to some boring and repetitive outcomes, and at worst it can lock our transformer's output into a loop.\n", + "\n", + "First, you should read HuggingFace's blog post [How to generate text: using different decoding methods for language generation with Transformers](https://huggingface.co/blog/how-to-generate). Once you've done that, you can start the exercises below." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "j6sK-yiffdfR" + }, + "source": [ + "## `TransformerSampler` class\n", + "\n", + "Below, we've given you the `TransformerSampler` class. This contains the following important methods:\n", + "\n", + "- `sample`, which is the highest-level method. It repeatedly calls `sample_next_token` to generate new tokens, until one of the termination criteria is met.\n", + "- `sample_next_token`, which samples a single new token based on some hyperparameters. This might involve various different sampling methods and techniques e.g. temperature scaling, top-k sampling, top-p sampling, etc.\n", + "- A set of other methods, which apply the previously mentioned sampling methods and techniques.\n", + "\n", + "You can see how `sample_next_token` works, and as an example how greedy sampling is implemented via `greedy_search` - we just continually take the tokens with the highest logits at each step.\n", + "\n", + "
\n", + "Question - why do you think temperature=0.0 corresponds to greedy sampling?\n", + "\n", + "To apply a temperature to our sampling (as we'll see later) means to scale all logits by `(1 / temperature)`. The basic intuition here is:\n", + "\n", + "* A higher temperature means a smaller scale factor, so the logits all approach zero, i.e. uniform distribution, and the sampling process is a lot more random (producing more diverse and varied outputs)\n", + "* A lower temperature means a larger scale factor, so the logits all approach infinity, i.e. a Dirac delta function, and the sampling process is a lot more deterministic (producing less varied output)\n", + "\n", + "As temperature gets close to zero, the difference between the largest logit and second largest logit becomes very large, so the distribution tends to \"probability of 1 on the highest-likelihood token\", i.e. greedy sampling. You can derive this formally if you prefer.\n", + "
\n", + "\n", + "In the next exercise you'll implement the `sample` method, and then you'll go on to implement all the other methods." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "MYFLw3K5fdfR" + }, + "source": [ + "### Exercise - implement `sample`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴🔴🔴⚪\n", + "> Importance: 🔵🔵🔵⚪⚪\n", + ">\n", + "> You should spend up to 25-40 minutes on this exercise.\n", + "> ```\n", + "\n", + "The `sample` method generates new tokens autoregressively, by repeatedly:\n", + "\n", + "- Passing the current sequence of tokens through the model to get logits,\n", + "- Using some sampling technique to select a new token, i.e. `sample_next_token(input_ids, logits, **kwargs)`,\n", + "- Appending this new token to the input sequence,\n", + "- Repeating the process until one of the termination criteria is met: either we generate `max_tokens_generated` new tokens, or we generate the end-of-sequence token (which we can access via `self.tokenizer.eos_token_id`).\n", + "\n", + "Lastly, we use the `tokenizer.decode` method to return the sampled string. You're also invited to use the `verbose` argument, for printing the decoded sequences while they're being generated (this can help with debugging).\n", + "\n", + "Below is some code which tests your sampling function by performing greedy sampling (which means always choosing the most likely next token at each step).\n", + "\n", + "A few hints:\n", + "\n", + "- Don't forget about tensor shapes! Your model's input should always have a batch dimension, i.e. it should be shape `(1, seq_len)`.\n", + "- The `sample_next_token` method will return an integer, so make sure you wrap this in a tensor before concatenating it to the end of your input IDs.\n", + "- Also remember to have your tensors be on the same device (we have a global `device` variable).\n", + "- Remember to put your model in evaluation mode, using `model.eval()`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "-96keyBEfdfR" + }, + "outputs": [], + "source": [ + "class TransformerSampler:\n", + " def __init__(self, model: DemoTransformer, tokenizer: GPT2TokenizerFast):\n", + " self.model = model\n", + " self.cfg = model.cfg\n", + " self.tokenizer = tokenizer\n", + "\n", + " @t.inference_mode()\n", + " def sample(self, prompt: str, max_tokens_generated=100, verbose=False, **kwargs) -> str:\n", + " \"\"\"\n", + " Returns a string of autoregressively generated text, starting from the prompt.\n", + "\n", + " Sampling terminates at max_tokens_generated, or when the model generates an end-of-sequence token. kwargs are\n", + " passed to sample_next_token, to give detailed instructions on how new tokens are chosen.\n", + " Pass `seed` to make generation reproducible.\n", + " \"\"\"\n", + " self.model.eval()\n", + " seed = kwargs.pop(\"seed\", None)\n", + " if seed is not None:\n", + " t.manual_seed(seed)\n", + " np.random.seed(seed)\n", + " raise NotImplementedError()\n", + "\n", + " @staticmethod\n", + " def sample_next_token(\n", + " input_ids: Int[Tensor, \" seq_len\"],\n", + " logits: Float[Tensor, \"d_vocab\"],\n", + " temperature=1.0,\n", + " top_k=0,\n", + " top_p=0.0,\n", + " frequency_penalty=0.0,\n", + " ) -> int:\n", + " assert input_ids.ndim == 1, \"input_ids should be a 1D sequence of token ids\"\n", + " assert logits.ndim == 1, \"logits should be a 1D tensor of shape (d_vocab,)\"\n", + " assert temperature >= 0, \"Temperature should be non-negative\"\n", + " assert 0 <= top_p <= 1.0, \"Top-p must be a probability\"\n", + " assert 0 <= top_k, \"Top-k must be non-negative\"\n", + " assert not (top_p != 0 and top_k != 0), \"At most one of top-p and top-k supported\"\n", + "\n", + " # Apply all the specialized sampling methods\n", + " if temperature == 0:\n", + " return TransformerSampler.greedy_search(logits)\n", + " elif temperature != 1.0:\n", + " logits = TransformerSampler.apply_temperature(logits, temperature)\n", + " if frequency_penalty != 0.0:\n", + " logits = TransformerSampler.apply_frequency_penalty(input_ids, logits, frequency_penalty)\n", + " if top_k > 0:\n", + " return TransformerSampler.sample_top_k(logits, top_k)\n", + " if top_p > 0.0:\n", + " return TransformerSampler.sample_top_p(logits, top_p)\n", + " return TransformerSampler.sample_basic(logits)\n", + "\n", + " @staticmethod\n", + " def greedy_search(logits: Float[Tensor, \"d_vocab\"]) -> int:\n", + " \"\"\"\n", + " Returns the most likely token (as an int).\n", + " \"\"\"\n", + " raise NotImplementedError()\n", + "\n", + " @staticmethod\n", + " def apply_temperature(logits: Float[Tensor, \"d_vocab\"], temperature: float) -> Float[Tensor, \"d_vocab\"]:\n", + " \"\"\"\n", + " Applies temperature scaling to the logits.\n", + " \"\"\"\n", + " raise NotImplementedError()\n", + "\n", + " @staticmethod\n", + " def apply_frequency_penalty(\n", + " input_ids: Int[Tensor, \" seq_len\"], logits: Float[Tensor, \"d_vocab\"], freq_penalty: float\n", + " ) -> Float[Tensor, \"d_vocab\"]:\n", + " \"\"\"\n", + " Applies a frequency penalty to the logits.\n", + " \"\"\"\n", + " raise NotImplementedError()\n", + "\n", + " @staticmethod\n", + " def sample_basic(logits: Float[Tensor, \"d_vocab\"]) -> int:\n", + " \"\"\"\n", + " Samples from the distribution defined by the logits.\n", + " \"\"\"\n", + " raise NotImplementedError()\n", + "\n", + " @staticmethod\n", + " def sample_top_k(logits: Float[Tensor, \"d_vocab\"], k: int) -> int:\n", + " \"\"\"\n", + " Samples from the top k most likely tokens.\n", + " \"\"\"\n", + " raise NotImplementedError()\n", + "\n", + " @staticmethod\n", + " def sample_top_p(logits: Float[Tensor, \"d_vocab\"], top_p: float, min_tokens_to_keep: int = 1) -> int:\n", + " \"\"\"\n", + " Samples from the most likely tokens which make up at least p cumulative probability.\n", + " \"\"\"\n", + " raise NotImplementedError()\n", + "\n", + " @t.inference_mode()\n", + " def beam_search(\n", + " self,\n", + " prompt: str,\n", + " num_return_sequences: int,\n", + " num_beams: int,\n", + " max_new_tokens: int,\n", + " no_repeat_ngram_size: int | None = None,\n", + " ) -> list[tuple[float, str]]:\n", + " \"\"\"\n", + " Implements a beam search, by repeatedly performing the `generate` and `filter` steps (starting from the initial\n", + " prompt) until either of the two stopping criteria are met: (1) we've generated `max_new_tokens` tokens, or (2)\n", + " we've generated `num_return_sequences` terminating sequences.\n", + " \"\"\"\n", + " raise NotImplementedError()\n", + "\n", + "\n", + "model = DemoTransformer(Config()).to(device)\n", + "model.load_state_dict(reference_gpt2.state_dict(), strict=False)\n", + "tokenizer = reference_gpt2.tokenizer\n", + "sampler = TransformerSampler(model, tokenizer)\n", + "\n", + "prompt = \"Jingle bells, jingle bells, jingle all the way\"\n", + "print(f\"Testing greedy decoding\\nPrompt: {prompt!r}\")\n", + "\n", + "expected = \"Jingle bells, jingle bells, jingle all the way up to the top of the mountain.\"\n", + "output = sampler.sample(prompt, max_tokens_generated=8, temperature=0.0)\n", + "\n", + "print(f\"Expected: {expected!r}\\nActual: {output!r}\\n\")\n", + "assert output == expected\n", + "\n", + "print(\"Tests passed!\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "iOJeGjITfdfS" + }, + "source": [ + "
\n", + "Solution\n", + "\n", + "```python\n", + "@t.inference_mode()\n", + "def sample(self, prompt: str, max_tokens_generated=100, verbose=False, **kwargs):\n", + " \"\"\"\n", + " Returns a string of autoregressively generated text, starting from the prompt.\n", + "\n", + " Sampling terminates at max_tokens_generated, or when the model generates an end-of-sequence token. kwargs are\n", + " passed to sample_next_token, to give detailed instructions on how new tokens are chosen. Pass `seed` to make\n", + " generation reproducible — it's set once before the loop here, not forwarded to sample_next_token.\n", + " \"\"\"\n", + " self.model.eval()\n", + " seed = kwargs.pop(\"seed\", None)\n", + " if seed is not None:\n", + " t.manual_seed(seed)\n", + " np.random.seed(seed)\n", + " input_ids = self.tokenizer.encode(prompt, return_tensors=\"pt\").to(device)[0]\n", + "\n", + " for i in range(max_tokens_generated):\n", + " # Get new logits (make sure we don't pass in more tokens than the model's context length)\n", + " logits = self.model(input_ids[None, -self.cfg.n_ctx :])\n", + " # We only take logits for the last token, because this is what we're sampling\n", + " logits = logits[0, -1]\n", + " # Get next token (as a tensor of size (1, 1) so we can concat it to input_ids)\n", + " next_token = t.tensor([TransformerSampler.sample_next_token(input_ids, logits, **kwargs)], device=device)\n", + " # Create new input ids string, with shape (1, old_seq_len + 1)\n", + " input_ids = t.cat([input_ids, next_token], dim=-1)\n", + " # Print out results, if required\n", + " if verbose:\n", + " print(self.tokenizer.decode(input_ids), end=\"\\r\")\n", + " # If our new token was the end-of-text token, stop\n", + " if next_token == getattr(self.tokenizer, \"eos_token_id\", None):\n", + " break\n", + "\n", + " return self.tokenizer.decode(input_ids)\n", + "```\n", + "\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "JNIsIurKfdfS" + }, + "source": [ + "## Sampling with Categorical\n", + "\n", + "Now, we'll move into implementing specific sampling methods. In each of these cases, you should return to the class definition above and fill in the corresponding method.\n", + "\n", + "PyTorch provides a [`distributions`](https://pytorch.org/docs/stable/distributions.html#distribution) package with a number of convenient methods for sampling from various distributions.\n", + "\n", + "For now, we just need [`t.distributions.categorical.Categorical`](https://pytorch.org/docs/stable/distributions.html#categorical). Use this to implement `sample_basic`, which just samples from the provided logits (which may have already been modified by the temperature and frequency penalties).\n", + "\n", + "Note that this will be slow since we aren't batching the samples, but don't worry about speed for now." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "4X5oh_HVfdfS" + }, + "source": [ + "### Exercise - `sample_basic`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴⚪⚪⚪\n", + "> Importance: 🔵🔵⚪⚪⚪\n", + ">\n", + "> You should spend up to 5-15 minutes on this exercise.\n", + "> ```\n", + "\n", + "Implement basic sampling in the `TransformerSampler` class above (i.e. the `sample_basic` method), then run the code below to verify your solution works." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "6QOUCOBrfdfS" + }, + "outputs": [], + "source": [ + "tests.test_sample_basic(TransformerSampler.sample_basic)\n", + "\n", + "prompt = \"John and Mary went to the\"\n", + "input_ids = tokenizer.encode(prompt, return_tensors=\"pt\").to(device)\n", + "logits = model(input_ids)[0, -1]\n", + "\n", + "expected_top_5 = {\n", + " \" church\": 0.0648,\n", + " \" house\": 0.0367,\n", + " \" temple\": 0.0145,\n", + " \" same\": 0.0104,\n", + " \" Church\": 0.0097,\n", + "}\n", + "frequency_of_top_5 = defaultdict(int)\n", + "\n", + "N = 10_000\n", + "for _ in tqdm(range(N)):\n", + " token = TransformerSampler.sample_next_token(input_ids.squeeze(), logits)\n", + " frequency_of_top_5[tokenizer.decode(token)] += 1\n", + "\n", + "for word in expected_top_5:\n", + " expected_freq = expected_top_5[word]\n", + " observed_freq = frequency_of_top_5[word] / N\n", + " print(f\"Word: {word!r:<9}. Expected freq {expected_freq:.4f}, observed freq {observed_freq:.4f}\")\n", + " assert abs(observed_freq - expected_freq) < 0.01, \"Try increasing N if this fails by a small amount.\"\n", + "\n", + "print(\"Tests passed!\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "Azzc1uWNfdfS" + }, + "source": [ + "
\n", + "Solution\n", + "\n", + "```python\n", + "@staticmethod\n", + "def sample_basic(logits: Float[Tensor, \"d_vocab\"]) -> int:\n", + " \"\"\"\n", + " Samples from the distribution defined by the logits.\n", + " \"\"\"\n", + " sampled_token = t.distributions.categorical.Categorical(logits=logits).sample()\n", + " return sampled_token.item()\n", + "```\n", + "\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "s9RHXCU6fdfS" + }, + "source": [ + "### Exercise - `apply_temperature`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴⚪⚪⚪⚪\n", + "> Importance: 🔵🔵⚪⚪⚪\n", + ">\n", + "> You should spend up to 5-10 minutes on this exercise.\n", + "> ```\n", + "\n", + "Temperature sounds fancy, but it's literally just dividing the logits by the temperature. You should implement this in your `TransformerSampler` class now." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "iBCrA4PlfdfS" + }, + "outputs": [], + "source": [ + "tests.test_apply_temperature(TransformerSampler.apply_temperature)\n", + "\n", + "logits = t.tensor([1, 2]).log()\n", + "\n", + "cold_logits = TransformerSampler.apply_temperature(logits, temperature=0.001)\n", + "print('A low temperature \"sharpens\" or \"peaks\" the distribution: ', cold_logits)\n", + "t.testing.assert_close(cold_logits, 1000.0 * logits)\n", + "\n", + "hot_logits = TransformerSampler.apply_temperature(logits, temperature=1000.0)\n", + "print(\"A high temperature flattens the distribution: \", hot_logits)\n", + "t.testing.assert_close(hot_logits, 0.001 * logits)\n", + "\n", + "print(\"Tests passed!\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "goodIVAFfdfS" + }, + "source": [ + "
\n", + "Solution\n", + "\n", + "```python\n", + "@staticmethod\n", + "def apply_temperature(logits: Float[Tensor, \"d_vocab\"], temperature: float) -> Float[Tensor, \"d_vocab\"]:\n", + " \"\"\"\n", + " Applies temperature scaling to the logits.\n", + " \"\"\"\n", + " return logits / temperature\n", + "```\n", + "\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "VcCEA8opfdfS" + }, + "source": [ + "### Exercise - `apply_frequency_penalty`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴⚪⚪⚪\n", + "> Importance: 🔵⚪⚪⚪⚪\n", + ">\n", + "> You should spend up to 10-15 minutes on this exercise.\n", + "> ```\n", + "\n", + "The frequency penalty is simple as well: count the number of occurrences of each token, then subtract `freq_penalty` for each occurrence. Hint: use `t.bincount` (documentation [here](https://pytorch.org/docs/stable/generated/torch.bincount.html)) to do this in a vectorized way.\n", + "\n", + "You should implement the `apply_frequency_penalty` method in your `TransformerSampler` class now, then run the cell below to check your solution." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "xNTRgoy8fdfS" + }, + "source": [ + "
\n", + "Help - I'm getting a RuntimeError; my tensor sizes don't match.\n", + "\n", + "Look at the documentation page for `t.bincount`. You might need to use the `minlength` argument - why?\n", + "
" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "KVZcyRYcfdfS" + }, + "outputs": [], + "source": [ + "tests.test_apply_frequency_penalty(TransformerSampler.apply_frequency_penalty)\n", + "\n", + "bieber_prompt = \"And I was like Baby, baby, baby, oh Like, Baby, baby, baby, no Like, Baby, baby, baby, oh I thought you'd always be mine, mine\"\n", + "input_ids = tokenizer.encode(bieber_prompt, return_tensors=\"pt\")\n", + "logits = t.ones(tokenizer.vocab_size)\n", + "penalized_logits = TransformerSampler.apply_frequency_penalty(input_ids.squeeze(), logits, 2.0)\n", + "\n", + "assert penalized_logits[5156].item() == -11, \"Expected 6 occurrences of ' baby' with leading space, 1-2*6=-11\"\n", + "assert penalized_logits[14801].item() == -5, \"Expected 3 occurrences of ' Baby' with leading space, 1-2*3=-5\"\n", + "\n", + "print(\"Tests passed!\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "qSFNOetbfdfS" + }, + "source": [ + "
\n", + "Solution\n", + "\n", + "```python\n", + "@staticmethod\n", + "def apply_frequency_penalty(\n", + " input_ids: Int[Tensor, \"seq_len\"], logits: Float[Tensor, \"d_vocab\"], freq_penalty: float\n", + ") -> Float[Tensor, \"d_vocab\"]:\n", + " \"\"\"\n", + " Applies a frequency penalty to the logits.\n", + " \"\"\"\n", + " d_vocab = logits.size(0)\n", + " id_freqs = t.bincount(input_ids, minlength=d_vocab)\n", + " return logits - freq_penalty * id_freqs\n", + "```\n", + "\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "suboGNK5fdfS" + }, + "source": [ + "### Sampling - Manual Testing\n", + "\n", + "Run the below cell to get a sense for the `temperature` and `freq_penalty` arguments. Play with your own prompt and try other values.\n", + "\n", + "Note: your model can generate newlines or non-printing characters, so calling `print` on generated text sometimes looks awkward on screen. You can call `repr` on the string before printing to have the string escaped nicely." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "azGPY8ydfdfS" + }, + "outputs": [], + "source": [ + "sampler = TransformerSampler(model, tokenizer)\n", + "\n", + "N_RUNS = 1\n", + "your_prompt = \"Jingle bells, jingle bells, jingle all the way\"\n", + "cases = [\n", + " (\"High freq penalty\", dict(frequency_penalty=100.0)),\n", + " (\"Negative freq penalty\", dict(frequency_penalty=-3.0)),\n", + " (\"Too hot!\", dict(temperature=2.0)),\n", + " (\"Pleasantly cool\", dict(temperature=0.7)),\n", + " (\"Pleasantly warm\", dict(temperature=0.9)),\n", + " (\"Too cold!\", dict(temperature=0.01)),\n", + "]\n", + "\n", + "table = Table(\"Name\", \"Kwargs\", \"Output\", title=\"Sampling - Manual Testing\")\n", + "\n", + "for name, kwargs in cases:\n", + " for i in range(N_RUNS):\n", + " output = sampler.sample(your_prompt, max_tokens_generated=24, **kwargs)\n", + " table.add_row(name, str(kwargs), repr(output) + \"\\n\")\n", + "\n", + "rprint(table)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "Qd_0SL4KfdfS" + }, + "source": [ + "## Top-K Sampling\n", + "\n", + "Conceptually, the steps in top-k sampling are:\n", + "- Find the `top_k` largest probabilities (you can use [`torch.topk`](https://pytorch.org/docs/stable/generated/torch.topk.html))\n", + "- Set all other probabilities to zero\n", + "- Normalize and sample" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "4QNrDWYNfdfS" + }, + "source": [ + "### Exercise - `sample_top_k`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴⚪⚪⚪\n", + "> Importance: 🔵⚪⚪⚪⚪\n", + ">\n", + "> You should spend up to 5-10 minutes on this exercise.\n", + "> ```\n", + "\n", + "Implement the method `sample_top_k` now. Your implementation should stay in log-space throughout (don't exponentiate to obtain probabilities). This means you don't actually need to worry about normalizing, because `Categorical` accepts unnormalised logits." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "_jvqoB7ofdfS" + }, + "outputs": [], + "source": [ + "tests.test_sample_top_k(TransformerSampler.sample_top_k)\n", + "\n", + "prompt = \"John and Mary went to the\"\n", + "input_ids = tokenizer.encode(prompt, return_tensors=\"pt\").to(device)\n", + "logits = model(input_ids)[0, -1]\n", + "\n", + "expected_top_5 = {\n", + " \" church\": 0.0648,\n", + " \" house\": 0.0367,\n", + " \" temple\": 0.0145,\n", + " \" same\": 0.0104,\n", + " \" Church\": 0.0097,\n", + "}\n", + "topk_5_sum = sum(expected_top_5.values())\n", + "\n", + "observed_freqs = defaultdict(int)\n", + "\n", + "N = 10000\n", + "for _ in tqdm(range(N)):\n", + " token = TransformerSampler.sample_next_token(input_ids.squeeze(), logits, top_k=5)\n", + " observed_freqs[tokenizer.decode(token)] += 1\n", + "\n", + "for word in expected_top_5:\n", + " expected_freq = expected_top_5[word] / topk_5_sum\n", + " observed_freq = observed_freqs[word] / N\n", + " print(f\"Word: {word!r:<9}. Expected freq = {expected_freq:.4f}, observed freq = {observed_freq:.4f}\")\n", + " assert abs(observed_freq - expected_freq) < 0.01" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "edbKR_SkfdfS" + }, + "source": [ + "
\n", + "Solution\n", + "\n", + "```python\n", + "@staticmethod\n", + "def sample_top_k(logits: Float[Tensor, \"d_vocab\"], k: int) -> int:\n", + " \"\"\"\n", + " Samples from the top k most likely tokens.\n", + " \"\"\"\n", + " top_k_logits, top_k_token_ids = logits.topk(k)\n", + " # Get sampled token (which is an index corresponding to the list of top-k tokens)\n", + " sampled_token_idx = t.distributions.categorical.Categorical(logits=top_k_logits).sample()\n", + " # Get the actual token id, as an int\n", + " return top_k_token_ids[sampled_token_idx].item()\n", + "```\n", + "\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "RAqsaI-3fdfS" + }, + "source": [ + "The [GPT-2 paper](https://d4mucfpksywv.cloudfront.net/better-language-models/language_models_are_unsupervised_multitask_learners.pdf) famously included an example prompt about unicorns. Now it's your turn to see just how cherry picked this example was.\n", + "\n", + "The paper claims they used `top_k=40` and best of 10 samples." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "AlegHri9fdfS" + }, + "outputs": [], + "source": [ + "sampler = TransformerSampler(model, tokenizer)\n", + "\n", + "your_prompt = \"In a shocking finding, scientist discovered a herd of unicorns living in a remote, previously unexplored valley, in the Andes Mountains. Even more surprising to the researchers was the fact that the unicorns spoke perfect English.\"\n", + "\n", + "output = sampler.sample(your_prompt, temperature=0.7, top_k=40, max_tokens_generated=64)\n", + "\n", + "rprint(f\"Your model said:\\n\\n[bold dark_orange]{output}\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "86suJ6xgfdfS" + }, + "source": [ + "This is pretty incredible! For some perspective on how much of a paradigm shift even basic models like this represented, we recommend reading [this section from Simulators](https://www.lesswrong.com/posts/vJFdjigzmcXMhNTsx/simulators#The_limit_of_sequence_modeling)." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "olHAxAK0fdfS" + }, + "source": [ + "## Top-p aka Nucleus Sampling\n", + "\n", + "The basic idea is that we choose the most likely words, up until the total probability of words we've chosen crosses some threshold. Then we sample from those chosen words based on their logits.\n", + "\n", + "The steps are:\n", + "\n", + "- Sort the probabilities from largest to smallest\n", + "- Find the cutoff point where the cumulative probability first equals or exceeds `top_p`. We do the cutoff inclusively, keeping the first probability above the threshold.\n", + "- If the number of kept probabilities is less than `min_tokens_to_keep`, keep that many tokens instead.\n", + "- Set all other probabilities to zero\n", + "- Normalize and sample\n", + "\n", + "For example, if our probabilities were `(0.4, 0.3, 0.2, 0.1)` and our cutoff was `top_p=0.8`, then we'd sample from the first three elements (because their total probability is `0.9` which is over the threshold, but the first two only have a total prob of `0.7` which is under the threshold). Once we've chosen to sample from those three, we would renormalise them by dividing by their sum, so the probabilities we use when sampling are `(0.4/0.9, 0.3/0.9, 0.2/0.9)`.\n", + "\n", + "Optionally, refer to the paper [The Curious Case of Neural Text Degeneration](https://arxiv.org/pdf/1904.09751.pdf) for some comparison of different methods." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "-2bFGYSWfdfS" + }, + "source": [ + "### Exercise - `sample_top_p`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴🔴⚪⚪\n", + "> Importance: 🔵⚪⚪⚪⚪\n", + ">\n", + "> You should spend up to 15-20 minutes on this exercise.\n", + "> ```" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "bbYRpnkZfdfS" + }, + "outputs": [], + "source": [ + "tests.test_sample_top_p(TransformerSampler.sample_top_p)\n", + "\n", + "prompt = \"John and Mary went to the\"\n", + "input_ids = tokenizer.encode(prompt, return_tensors=\"pt\").to(device)\n", + "logits = model(input_ids)[0, -1]\n", + "\n", + "expected_top_10pct = {\n", + " \" church\": 0.0648,\n", + " \" house\": 0.0367, # These are the two most likely tokens, and add up to >10%\n", + "}\n", + "top_10pct_sum = sum(expected_top_10pct.values())\n", + "\n", + "observed_freqs = defaultdict(int)\n", + "\n", + "N = 10_000\n", + "for _ in tqdm(range(N)):\n", + " token = TransformerSampler.sample_next_token(input_ids.squeeze(), logits, top_p=0.1)\n", + " observed_freqs[tokenizer.decode(token)] += 1\n", + "\n", + "for word in expected_top_10pct:\n", + " expected_freq = expected_top_10pct[word] / top_10pct_sum\n", + " observed_freq = observed_freqs[word] / N\n", + " print(f\"Word: {word!r:<9}. Expected freq {expected_freq:.4f}, observed freq {observed_freq:.4f}\")\n", + " assert abs(observed_freq - expected_freq) < 0.01, \"Try increasing N if this fails by a small amount.\"" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "YnHiJ4grfdfS" + }, + "source": [ + "
\n", + "Help - I'm stuck on how to implement this function.\n", + "\n", + "First, sort the logits using the `sort(descending=True)` method (this returns values and indices). Then you can get `cumulative_probs` by applying softmax to these logits and taking the cumsum. Then, you can decide how many probabilities to keep by using the `t.searchsorted` function.\n", + "\n", + "Once you've decided which probabilities to keep, it's easiest to sample from them using the original logits (you should have preserved the indices when you called `logits.sort`). This way, you don't need to worry about renormalising like you would if you were using probabilities.\n", + "
\n", + "\n", + "
\n", + "Solution\n", + "\n", + "```python\n", + "@staticmethod\n", + "def sample_top_p(logits: Float[Tensor, \"d_vocab\"], top_p: float, min_tokens_to_keep: int = 1) -> int:\n", + " \"\"\"\n", + " Samples from the most likely tokens which make up at least p cumulative probability.\n", + " \"\"\"\n", + " # Sort logits, and get cumulative probabilities\n", + " logits_sorted, indices = logits.sort(descending=True, stable=True)\n", + " cumul_probs = logits_sorted.softmax(-1).cumsum(-1)\n", + " # Choose which tokens to keep, in the set we sample from\n", + " n_keep = t.searchsorted(cumul_probs, top_p, side=\"left\").item() + 1\n", + " n_keep = max(n_keep, min_tokens_to_keep)\n", + " keep_idx = indices[:n_keep]\n", + " keep_logits = logits[keep_idx]\n", + " # Perform the sampling\n", + " sample = t.distributions.categorical.Categorical(logits=keep_logits).sample()\n", + " return keep_idx[sample].item()\n", + "```\n", + "\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "6b8a9Y51fdfS" + }, + "source": [ + "Now, an example of top-p sampling:" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "0tXwmxydfdfS" + }, + "outputs": [], + "source": [ + "sampler = TransformerSampler(model, tokenizer)\n", + "\n", + "your_prompt = \"Eliezer Shlomo Yudkowsky (born September 11, 1979) is an American decision and artificial intelligence (AI) theorist and writer, best known for\"\n", + "output = sampler.sample(your_prompt, temperature=0.7, top_p=0.95, max_tokens_generated=64)\n", + "rprint(f\"Your model said:\\n\\n[bold dark_orange]{output}\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "JtdAqDhHfdfT" + }, + "source": [ + "## Beam search\n", + "\n", + "Finally, we'll implement a more advanced way of searching over output: **beam search**. You should read the [HuggingFace page](https://huggingface.co/blog/how-to-generate#beam-search) on beam search before moving on.\n", + "\n", + "In beam search, we maintain a list of size `num_beams` completions which are the most likely completions so far as measured by the product of their probabilities. Since this product can become very small, we use the sum of log probabilities instead. Note - log probabilities are *not* the same as your model's output. We get log probabilities by first taking softmax of our output and then taking log. You can do this with the [`log_softmax`](https://pytorch.org/docs/stable/generated/torch.nn.functional.log_softmax.html) function / tensor method.\n", + "\n", + "
\n", + "Log probabilities are equal to the logit output after being translated by some amount X (where X is a function of the original logit output). Can you prove this?\n", + "\n", + "Suppose our vector of logits is $x$, and we take softmax to get a vector of probabilities $p$, then log again to get a vector of log probabilities $l$. Then the $i$-th element of this vector of logprobs is:\n", + "\n", + "$$\n", + "\\begin{align}\n", + "l_i &= \\log p_i \\\\\n", + "&= \\log \\frac{\\exp(x_i)}{\\sum_j \\exp(x_j)} \\\\\n", + "&= x_i - \\log \\sum_j \\exp(x_j) \\\\\n", + "&= x_i - C\n", + "\\end{align}\n", + "$$\n", + "\n", + "where $C = \\log \\sum_j \\exp(x_j)$ is the same for all elements. So we can see that $l_i$ is equal to the logit output $x_i$ after being translated by $C$.\n", + "\n", + "It's important not to mix up logits and logprobs!\n", + "
\n", + "\n", + "
\n", + "Why do you think we use log softmax rather than logit output?\n", + "\n", + "Logit output is translation invariant. If we had two different beams and we were generating the next tokens in those beams, there would be no reasonable way to compare the two beams to each other, because we could shift the logit vector for one beam by a constant amount without changing the distribution.\n", + "\n", + "
\n", + "\n", + "At each iteration, we run the batch of completions through the model and take the log-softmax to obtain `d_vocab` log-probs for each completion, or `num_beams * d_vocab` possible next completions in total.\n", + "\n", + "If we kept all of these, then we would have `num_beams * d_vocab * d_vocab` completions after the next iteration which is way too many, so instead we sort them by their score and loop through from best (highest) log probability to worst (lowest).\n", + "\n", + "The illustration below might help (based on real results from this method). Here, we have the following hyperparameters:\n", + "\n", + "```python\n", + "num_beams = 3\n", + "max_new_tokens = 3\n", + "num_return_sequences = 2\n", + "```\n", + "\n", + "\n", + "\n", + "Note how after each \"generate\" stage, we have `num_beams ** 2` possible completions, which we then filter down to `num_beams`. This is because we need this many in order to find the best `num_beams` completions overall - for example, it's possible that all the best beams of length `n+1` come from the same beam of length `n`, in which case we'll need to keep all `num_beams` that we generated from that single beam.\n", + "\n", + "How do we deal with sequences that terminate early (i.e. by generating an EOS token)? Answer - we append them to the list of completions which we'll return at the end, and remove them from the generation tree. Our algorithm terminates when either all our sequences have length `max_new_tokens` larger than the initial prompt length, or we've generated `num_return_sequences` terminating sequences." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "ITK8_L43fdfT" + }, + "source": [ + "### Exercise - implement `beam_search`\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴🔴🔴🔴\n", + "> Importance: 🔵⚪⚪⚪⚪\n", + ">\n", + "> You should spend up to 30-50 minutes on this exercise.\n", + "> ```\n", + "\n", + "We've given you one implementation of `beam_search` below, which calls the `generate` and `filter` methods of the `Beams` class (these correspond to the two stages in the diagram above). The `beam_search` method works as follows:\n", + "\n", + "- Create a list `final_logprobs_and_completions` for storing the final output, as tuples of (logprob sum, string completion).\n", + "- Perform `max_new_tokens` steps of generation (producing a new set of beams) and filtering (getting the best beams from these combinations), while also adding terminated beams to the list of best beams\n", + "- Return these terminated beams plus the best ones we have at the end of the steps.\n", + "\n", + "So all you need to do is fill in the `generate` and `filter` methods. Below, you'll find some unit tests for the `generate` and `filter` methods. When you've passed these tests, you should be able to run the full `beam_search` function.\n", + "\n", + "**Important note** - by default, beam search produces a lot of repeated words / phrases / sentences. This makes sense - if the model finds some completion with a much higher logit sum than most completions in its beam search space, then it will want to repeat this completion even if it doesn't make a lot of sense in context. A common solution is to ban repetition of n-grams, which you should also implement in the function below. In other words, rather than sampling tokens from each sequence by taking `logprobs.topk(k)` in your `generate` method, you should take the `k` top tokens after filtering out those that give you repeated n-grams of length `no_repeat_ngram_size`. Good values of this parameter to try are 2 or 3 (although we recommend you try without this parameter first, so you can see how much of a difference it makes!)." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "YsVbdmYCfdfT" + }, + "outputs": [], + "source": [ + "@dataclass\n", + "class Beams:\n", + " \"\"\"Class to store beams during beam search.\"\"\"\n", + "\n", + " model: DemoTransformer\n", + " tokenizer: GPT2TokenizerFast\n", + " logprob_sums: Float[Tensor, \" batch\"]\n", + " tokens: Int[Tensor, \"batch seq\"]\n", + "\n", + " def __getitem__(self, batch_idx) -> \"Beams\":\n", + " \"\"\"Allows you to create new beams from old beams by slicing along batch dim (useful for `filter`).\"\"\"\n", + " return Beams(self.model, self.tokenizer, self.logprob_sums[batch_idx], self.tokens[batch_idx])\n", + "\n", + " @property\n", + " def logprobs_and_completions(self) -> list[tuple[float, str]]:\n", + " \"\"\"Returns self as a list of logprob sums and completions (useful for getting final output).\"\"\"\n", + " return [\n", + " (logprob_sum.item(), self.tokenizer.decode(tokens))\n", + " for (logprob_sum, tokens) in zip(self.logprob_sums, self.tokens)\n", + " ]\n", + "\n", + " def generate(self, k: int, no_repeat_ngram_size: int | None = None) -> \"Beams\":\n", + " \"\"\"\n", + " Starts from the current set of beams (i.e. self.tokens) and returns a new set of `len(self.tokens) * k` beams,\n", + " containing the best `k` continuations for each of the original beams.\n", + "\n", + " Optional argument `no_repeat_ngram_size` means your model won't generate any sequences with a repeating n-gram\n", + " of this length.\n", + " \"\"\"\n", + " raise NotImplementedError()\n", + "\n", + " def filter(self, k: int) -> tuple[\"Beams\", \"Beams\"]:\n", + " \"\"\"\n", + " Returns:\n", + " best_beams: Beams\n", + " filtered version of self, containing all best `k` which are also not terminated.\n", + " early_terminations: Beams\n", + " filtered version of self, containing all best `k` which are also terminated.\n", + " \"\"\"\n", + " raise NotImplementedError()\n", + "\n", + "\n", + " def print(self, title=\"Best completions\", max_print_chars=80) -> None:\n", + " \"\"\"\n", + " Prints out a set of sequences with their corresponding logprob sums.\n", + " \"\"\"\n", + " if len(self.tokens) == 0:\n", + " return\n", + " table = Table(\"logprob sum\", \"completion\", title=title)\n", + " for logprob_sum, tokens in zip(self.logprob_sums, self.tokens):\n", + " text = self.tokenizer.decode(tokens)\n", + " if len(repr(text)) > max_print_chars:\n", + " text = text[: int(0.3 * max_print_chars)] + \" ... \" + text[-int(0.7 * max_print_chars) :]\n", + " table.add_row(f\"{logprob_sum:>8.3f}\", repr(text))\n", + " rprint(table)\n", + "\n", + "\n", + "@t.inference_mode()\n", + "def beam_search(\n", + " self: TransformerSampler,\n", + " prompt: str,\n", + " num_return_sequences: int,\n", + " num_beams: int,\n", + " max_new_tokens: int,\n", + " no_repeat_ngram_size: int | None = None,\n", + ") -> list[tuple[float, str]]:\n", + " \"\"\"\n", + " Implements a beam search, by repeatedly performing the `generate` and `filter` steps (starting from the initial\n", + " prompt) until either of the two stopping criteria are met: (1) we've generated `max_new_tokens` tokens, or (2)\n", + " we've generated `num_return_sequences` terminating sequences.\n", + " \"\"\"\n", + " assert num_return_sequences <= num_beams\n", + " self.model.eval()\n", + "\n", + " tokens = self.tokenizer.encode(prompt, return_tensors=\"pt\").to(device)\n", + "\n", + " final_logprobs_and_completions = [] # we add to this list as we get terminated beams\n", + " best_beams = Beams(self.model, self.tokenizer, t.tensor([0.0]).to(device), tokens) # start with just 1 beam\n", + "\n", + " for _ in tqdm(range(max_new_tokens)):\n", + " t.cuda.empty_cache()\n", + "\n", + " # Generate & filter beams\n", + " best_beams = best_beams.generate(k=num_beams, no_repeat_ngram_size=no_repeat_ngram_size)\n", + " best_beams, best_beams_terminated = best_beams.filter(k=num_beams)\n", + "\n", + " # Add terminated beams to our list, and return early if we have enough\n", + " final_logprobs_and_completions.extend(best_beams_terminated.logprobs_and_completions)\n", + " if len(final_logprobs_and_completions) >= num_return_sequences:\n", + " return final_logprobs_and_completions[:num_return_sequences]\n", + "\n", + " # Return terminated beams plus the best ongoing beams of length `orig_len + max_new_tokens`\n", + " final_logprobs_and_completions.extend(best_beams.logprobs_and_completions)\n", + " return final_logprobs_and_completions[:num_return_sequences]\n", + "\n", + "\n", + "TransformerSampler.beam_search = beam_search" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "6HAW7f5qfdfT" + }, + "source": [ + "
\n", + "Help - I'm stuck on the implementation of no_repeat_ngram_size.\n", + "\n", + "Here's a method, which you can use in your `generate` function in place of `logprobs.topk(k)`, which filters out the ngrams of length `no_repeat_ngram_size` which have already appeared in `self.tokens`:\n", + "\n", + "```python\n", + "def get_topk_non_repeating(\n", + " self,\n", + " logprobs: Float[Tensor, \"batch d_vocab\"],\n", + " no_repeat_ngram_size: int | None,\n", + " k: int,\n", + ") -> tuple[Float[Tensor, \"k\"], Int[Tensor, \"k\"]]:\n", + " \"\"\"\n", + " logprobs:\n", + " tensor of the log-probs for the next token\n", + " no_repeat_ngram_size:\n", + " size of ngram to avoid repeating\n", + " k:\n", + " number of top logits to return, for each beam in our collection\n", + "\n", + " Returns:\n", + " equivalent to the output of `logprobs.topk(dim=-1)`, but makes sure that no returned tokens would produce an\n", + " ngram of size `no_repeat_ngram_size` which has already appeared in `self.tokens`.\n", + " \"\"\"\n", + " batch, seq_len = self.tokens.shape\n", + "\n", + " # If completion isn't long enough for a repetition, or we have no restrictions, just return topk\n", + " if (no_repeat_ngram_size is not None) and (seq_len > no_repeat_ngram_size - 1):\n", + " # Otherwise, we need to check for ngram repetitions\n", + " # First, get the most recent `no_repeat_ngram_size-1` tokens\n", + " last_ngram_prefix = self.tokens[:, seq_len - (no_repeat_ngram_size - 1) :]\n", + " # Next, find all the tokens we're not allowed to generate, by checking all past ngrams for a match\n", + " for i in range(seq_len - (no_repeat_ngram_size - 1)):\n", + " ngrams = self.tokens[:, i : i + no_repeat_ngram_size] # (batch, ngram)\n", + " ngrams_are_repeated = (ngrams[:, :-1] == last_ngram_prefix).all(-1) # (batch,)\n", + " ngram_end_tokens = ngrams[:, [-1]] # (batch, 1)\n", + " # Fill logprobs with neginf wherever the ngrams are repeated\n", + " logprobs[range(batch), ngram_end_tokens] = t.where(\n", + " ngrams_are_repeated, -1.0e4, logprobs[range(batch), ngram_end_tokens]\n", + " )\n", + "\n", + " # Finally, get our actual tokens\n", + " return logprobs.topk(k=k, dim=-1)\n", + "```\n", + "\n", + "
\n", + "\n", + "
\n", + "Solution\n", + "\n", + "```python\n", + "def generate(self, k: int, no_repeat_ngram_size: int | None = None) -> \"Beams\":\n", + " \"\"\"\n", + " Starts from the current set of beams (i.e. self.tokens) and returns a new set of `len(self.tokens) * k` beams,\n", + " containing the best `k` continuations for each of the original beams.\n", + "\n", + " Optional argument `no_repeat_ngram_size` means your model won't generate any sequences with a repeating n-gram\n", + " of this length.\n", + " \"\"\"\n", + " # Get the output logprobs for the next token (for every sequence in current beams)\n", + " logprobs = self.model(self.tokens)[:, -1, :].log_softmax(-1)\n", + "\n", + " # Get the top-k tokens for each sequence\n", + " topk_logprobs, topk_tokenIDs = self.get_topk_non_repeating(logprobs, no_repeat_ngram_size, k=k)\n", + "\n", + " # Add new logprobs & concat new tokens. When doing this, we need to add an extra `k` dimension since our current\n", + " # logprobs & tokens have shape (batch,) and (batch, seq), but our new ones both have shape (batch, k)\n", + " new_logprob_sums = einops.repeat(self.logprob_sums, \"b -> b k\", k=k) + topk_logprobs\n", + " new_tokens = t.concat([einops.repeat(self.tokens, \"b s -> b k s\", k=k), topk_tokenIDs.unsqueeze(-1)], dim=-1)\n", + "\n", + " return Beams(self.model, self.tokenizer, new_logprob_sums.flatten(), new_tokens.flatten(0, 1))\n", + "\n", + "def filter(self, k: int) -> tuple[\"Beams\", \"Beams\"]:\n", + " \"\"\"\n", + " Returns:\n", + " best_beams: Beams\n", + " filtered version of self, containing all best `k` which are also not terminated.\n", + " early_terminations: Beams\n", + " filtered version of self, containing all best `k` which are also terminated.\n", + " \"\"\"\n", + " # Get the indices of top `k` beams\n", + " top_beam_indices = self.logprob_sums.topk(k=k, dim=0).indices.tolist()\n", + " # Get the indices of terminated sequences\n", + " new_tokens = self.tokens[:, -1]\n", + " terminated_indices = t.nonzero(new_tokens == self.tokenizer.eos_token_id)\n", + "\n", + " # Get the indices of the `k` best sequences (some terminated, some not terminated)\n", + " best_continuing = [i for i in top_beam_indices if i not in terminated_indices]\n", + " best_terminated = [i for i in top_beam_indices if i in terminated_indices]\n", + "\n", + " # Return the beam objects from these indices\n", + " return self[best_continuing], self[best_terminated]\n", + "\n", + "def get_topk_non_repeating(\n", + " self,\n", + " logprobs: Float[Tensor, \"batch d_vocab\"],\n", + " no_repeat_ngram_size: int | None,\n", + " k: int,\n", + ") -> tuple[Float[Tensor, \"k\"], Int[Tensor, \"k\"]]:\n", + " \"\"\"\n", + " logprobs:\n", + " tensor of the log-probs for the next token\n", + " no_repeat_ngram_size:\n", + " size of ngram to avoid repeating\n", + " k:\n", + " number of top logits to return, for each beam in our collection\n", + "\n", + " Returns:\n", + " equivalent to the output of `logprobs.topk(dim=-1)`, but makes sure that no returned tokens would produce an\n", + " ngram of size `no_repeat_ngram_size` which has already appeared in `self.tokens`.\n", + " \"\"\"\n", + " batch, seq_len = self.tokens.shape\n", + "\n", + " # If completion isn't long enough for a repetition, or we have no restrictions, just return topk\n", + " if (no_repeat_ngram_size is not None) and (seq_len > no_repeat_ngram_size - 1):\n", + " # Otherwise, we need to check for ngram repetitions\n", + " # First, get the most recent `no_repeat_ngram_size-1` tokens\n", + " last_ngram_prefix = self.tokens[:, seq_len - (no_repeat_ngram_size - 1) :]\n", + " # Next, find all the tokens we're not allowed to generate, by checking all past ngrams for a match\n", + " for i in range(seq_len - (no_repeat_ngram_size - 1)):\n", + " ngrams = self.tokens[:, i : i + no_repeat_ngram_size] # (batch, ngram)\n", + " ngrams_are_repeated = (ngrams[:, :-1] == last_ngram_prefix).all(-1) # (batch,)\n", + " ngram_end_tokens = ngrams[:, [-1]] # (batch, 1)\n", + " # Fill logprobs with neginf wherever the ngrams are repeated\n", + " logprobs[range(batch), ngram_end_tokens] = t.where(\n", + " ngrams_are_repeated, -1.0e4, logprobs[range(batch), ngram_end_tokens]\n", + " )\n", + "\n", + " # Finally, get our actual tokens\n", + " return logprobs.topk(k=k, dim=-1)\n", + "```\n", + "\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "rjXvw_iRfdfT" + }, + "source": [ + "Example usage of the `Beams` class, and the `print` method, corresponding to the diagram above:" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "dbgpHXuvfdfT" + }, + "outputs": [], + "source": [ + "# Start with prompt \"When I was\", get top 3 tokens (and their logprobs), and use that to create & display the top 3 beams\n", + "prompt = \"When I was\"\n", + "tokens = tokenizer.encode(prompt, return_tensors=\"pt\").to(device)\n", + "logprobs = model(tokens)[0, -1].log_softmax(-1)\n", + "top_logprobs, top_tokens = logprobs.topk(k=3, dim=-1)\n", + "\n", + "new_tokens = t.concat([tokens.repeat(3, 1), top_tokens.unsqueeze(-1)], dim=-1)\n", + "\n", + "beams = Beams(model, tokenizer, logprob_sums=top_logprobs, tokens=new_tokens)\n", + "beams.print()" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "4_vHkNyofdfT" + }, + "source": [ + "And here are some unit tests for your `generate` and `filter` methods, starting from the prompt `\"When I was\"` (so your output should match the diagram above)." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "GtTnOmRdfdfT" + }, + "outputs": [], + "source": [ + "print(\"Testing generate...\")\n", + "new_beams = beams.generate(k=3, no_repeat_ngram_size=1)\n", + "new_beams.print()\n", + "\n", + "expected_values = [\n", + " (-3.1, \"When I was a kid\"),\n", + " (-4.8, \"When I was a child\"),\n", + " (-4.9, \"When I was a little\"),\n", + "]\n", + "\n", + "for i, (logprob_sum, completion) in enumerate(new_beams.logprobs_and_completions[:3]):\n", + " assert abs(logprob_sum - expected_values[i][0]) < 0.1, f\"{i}\"\n", + " assert completion == expected_values[i][1], f\"{i}\"\n", + "\n", + "print(\"All tests for `generate` passed!\")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "mdNeIYxNfdfT" + }, + "outputs": [], + "source": [ + "print(\"Testing `filter`...\")\n", + "\n", + "best_beams, terminated_beams = new_beams.filter(3)\n", + "best_beams.print()\n", + "\n", + "expected_values = [\n", + " (-3.1, \"When I was a kid\"),\n", + " (-3.2, \"When I was growing up\"),\n", + " (-4.6, \"When I was in the\"),\n", + "]\n", + "\n", + "for i, (logprob_sum, completion) in enumerate(best_beams.logprobs_and_completions):\n", + " assert abs(logprob_sum - expected_values[i][0]) < 0.1, f\"{i}\"\n", + " assert completion == expected_values[i][1], f\"{i}\"\n", + "\n", + "assert len(terminated_beams.logprobs_and_completions) == 0\n", + "\n", + "print(\"All tests for `filter` passed!\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "0E9uJll2fdfT" + }, + "source": [ + "Lastly, we'll test the `no_repeat_ngram_size` argument. We do this by continually generating new tokens from our starting beams `beams`, and seeing if the model repeats the `I was` ngram (which it will by default unless we prohibit repeating n-grams)." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "ZOUUnQSKfdfT" + }, + "outputs": [], + "source": [ + "print(\"Testing `no_repeat_ngram_size`...\")\n", + "\n", + "new_beams = beams\n", + "for _ in range(5):\n", + " new_beams = new_beams.generate(k=1)\n", + "new_beams.print(title=\"Completions with no ngram restriction\")\n", + "assert all(\"I was\" in completion.removeprefix(prompt) for _, completion in new_beams.logprobs_and_completions), (\n", + " \"Without restriction, all beams should be completed as '...I was...'\"\n", + ")\n", + "\n", + "new_beams = beams\n", + "for _ in range(5):\n", + " new_beams = new_beams.generate(k=1, no_repeat_ngram_size=2)\n", + "new_beams.print(title=\"Completions with no repeated bigrams\")\n", + "assert all(\"I was\" not in completion.removeprefix(prompt) for _, completion in new_beams.logprobs_and_completions), (\n", + " \"With no repeated bigrams, no beams should contain a second '...I was...'\"\n", + ")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "AhJ7dtIdfdfT" + }, + "source": [ + "Once you've passed all of these unit tests, you can try implementing the full beam search function. It should create a `Beams` object from the initial prompt, and then repeatedly call `generate` and `filter` until the stopping criteria are met." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "QwQZuCyffdfT" + }, + "outputs": [], + "source": [ + "sampler = TransformerSampler(model, tokenizer)\n", + "\n", + "prompt = \"The ships hung in the sky in much the same way that\"\n", + "orig_len = len(tokenizer.encode(prompt))\n", + "\n", + "final_logitsums_and_completions = sampler.beam_search(\n", + " prompt=prompt,\n", + " num_return_sequences=3,\n", + " num_beams=40,\n", + " max_new_tokens=60,\n", + " no_repeat_ngram_size=2,\n", + ")\n", + "\n", + "# Print all the best output\n", + "for logprob_sum, text in final_logitsums_and_completions:\n", + " avg_logprob_as_prob = t.tensor(logprob_sum / (len(tokenizer.encode(text)) - orig_len)).exp()\n", + " rprint(f\"Avg token prob = {avg_logprob_as_prob:.3f}\\nBest output:\\n[bold dark_orange]{text}\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "NTyvflHSfdfT" + }, + "source": [ + "## KV Caching" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "FIpJ-MAYfdfT" + }, + "source": [ + "*This section is also designed to be challenging, and take quite some time. There are many different ways to solve it, and you're expected to try and find your own way (you should think about this for a while before looking at the suggestions in the dropdowns). Additionally, you might not find it as interesting as some of the other sections.*" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "quss9KUffdfT" + }, + "source": [ + "### How can caching help us?\n", + "\n", + "The text generation we've done so far is needlessly re-computing certain values, which is very noticeable when you try to generate longer sequences.\n", + "\n", + "Suppose you're generating text, and you've already run GPT on the sentence \"My life motto:\". Now you want to run the model on the sentence \"My life motto: Always\". Which computations from the first sentence can you reuse?\n", + "\n", + "
\n", + "Answer\n", + "\n", + "At each attention layer, the only things the attention layer needs from the previous sequence positions are the key and value vectors. This is explained in the following diagram, which compares the attention layer with and without caching (it's a big diagram so you might want to open it in a separate window to zoom in).\n", + "\n", + "\n", + "\n", + "\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "vql4Ym00fdfT" + }, + "source": [ + "### Exercise - implement KV caching\n", + "\n", + "> ```yaml\n", + "> Difficulty: 🔴🔴🔴🔴🔴\n", + "> Importance: 🔵⚪⚪⚪⚪\n", + ">\n", + "> You are expected to spend well over an hour on this exercise, if you choose to do it.\n", + "> ```\n", + "\n", + "Modify your GPT-2 to optionally use a cache. When you run your GPT on `\"My life motto:\"`, it should store the necessary values in the cache. Then in the next forward pass with just `\" Always\"` as input, it should load the cached values instead of recomputing them (and update the cache). This only needs to work with a single input sequence (batch size of 1), and you can assume that after the first forward pass, the input will be just one token.\n", + "\n", + "The design of the cache is completely up to you - discuss possible designs with your partner before writing code. It should be possible to have only one GPT2 instance and many different cache instances at one time. Imagine that you want to use one instance to serve multiple users submitting requests for text generation like in [AI Dungeon](https://aidungeon.io/).\n", + "\n", + "You'll also need to rewrite parts of your `DemoTransformer` code, in order to get this to work. The tests have been built to accommodate modules which return their output as the first element in a tuple (i.e. `(output, cache)`) rather than just returning the output, so you should use the tests to verify that your modules still work as expected.\n", + "\n", + "Some example considerations:\n", + "\n", + "* Which GPT-2 classes need to interact with the cache?\n", + " * Will you need to change the positional embedding, and if so then how?\n", + "* Should the cache be mutable and be updated in place, or does updating actually just create a separate instance?\n", + " * *(Hint here - think about how you might use the cache during beam search.)*\n", + "* Is it possible for other programmers to incorrectly use your cache? Is there a way to prevent this failure mode or at least detect this and complain loudly?" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "1n8YNoyAfdfT" + }, + "source": [ + "
\n", + "Cache implementation (example)\n", + "\n", + "This KeyValueCache object is structured as just a fancy tensor (it inherits all the methods from Tensor). The main difference is that it has a few extra helper methods, e.g. constructing an empty cache from a Config object.\n", + "\n", + "There are other ways you could do this, e.g. having your `KeyValueCache` class contain a list of `KeyValueCacheEntry` objects (where each of these corresponds to a different layer).\n", + "\n", + "```python\n", + "# Define a type for a single layer's cache entry (useful for type checking in later functions)\n", + "KeyValueCacheTensor = Float[Tensor, \"2 batch seq_len n_heads d_head\"]\n", + "\n", + "class KeyValueCache(Tensor):\n", + " '''\n", + " This class holds tensors of key and value vectors, to be used for caching.\n", + "\n", + " If we define it using cfg and batch then it's initialized as empty, but\n", + " we can also define it from kv_cache_entries.\n", + " '''\n", + " @classmethod\n", + " def new_empty(cls, cfg: Config, batch: int = 1) -> \"KeyValueCache\":\n", + " '''\n", + " Doing a forward pass on a cache created in this way indicates \"we don't\n", + " yet have a cache, but we want this forward pass to return a cache\".\n", + " Whereas using cache=None in a forward pass indicates we don't want to\n", + " return a cache.\n", + " '''\n", + " shape = (cfg.n_layers, 2, batch, 0, cfg.n_heads, cfg.d_head)\n", + " return cls(*shape).to(device)\n", + "\n", + " # Define a handful of properties, so they can be referenced directly rather than\n", + " # indexing (which is more likely to lead to mistakes)\n", + "\n", + " @property\n", + " def k(self) -> Tensor:\n", + " return self[:, 0]\n", + "\n", + " @property\n", + " def v(self) -> Tensor:\n", + " return self[:, 1]\n", + "\n", + " @property\n", + " def batch(self) -> int:\n", + " return self.shape[2]\n", + "\n", + " @property\n", + " def seq_len(self) -> int:\n", + " return self.shape[3]\n", + "\n", + "\n", + "# Example implementation:\n", + "cfg = model.cfg\n", + "batch = 6\n", + "kv_cache = KeyValueCache.new_empty(cfg, batch)\n", + "\n", + "print(f\"Shape of all kv-cache = {tuple(kv_cache.shape)}\")\n", + "print(f\"Shape of just k-cache = {tuple(kv_cache.k.shape)}\")\n", + "for kv_cache_entry in kv_cache:\n", + " print(f\"Shape of cache entry for one layer = {tuple(kv_cache_entry.shape)}\")\n", + " break\n", + "print(f\"Batch size = {kv_cache.batch}\")\n", + "print(f\"Current sequence length = {kv_cache.seq_len}\")\n", + "```\n", + "\n", + "
\n", + "\n", + "
\n", + "New DemoTransformer components (and testing)\n", + "\n", + "```python\n", + "# Define new model parts where necessary, and create a new model & test it\n", + "# Note that sometimes our modules return a tuple of (tensor output, cache) rather than just output. The\n", + "# tests have been built to accommodate this.\n", + "\n", + "\n", + "class PosEmbed(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.W_pos = nn.Parameter(t.empty((cfg.n_ctx, cfg.d_model)))\n", + " nn.init.normal_(self.W_pos, std=self.cfg.init_range)\n", + "\n", + " def forward(\n", + " self,\n", + " tokens: Int[Tensor, \"batch position\"],\n", + " past_kv_pos_offset: int = 0\n", + " ) -> Float[Tensor, \"batch position d_model\"]:\n", + "\n", + " batch, seq_len = tokens.shape\n", + " return einops.repeat(\n", + " self.W_pos[past_kv_pos_offset: seq_len+past_kv_pos_offset],\n", + " \"seq d_model -> batch seq d_model\",\n", + " batch=batch\n", + " )\n", + "\n", + "\n", + "class Attention(nn.Module):\n", + " IGNORE: Float[Tensor, \"\"]\n", + "\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.W_Q = nn.Parameter(t.empty((cfg.n_heads, cfg.d_model, cfg.d_head)))\n", + " self.W_K = nn.Parameter(t.empty((cfg.n_heads, cfg.d_model, cfg.d_head)))\n", + " self.W_V = nn.Parameter(t.empty((cfg.n_heads, cfg.d_model, cfg.d_head)))\n", + " self.W_O = nn.Parameter(t.empty((cfg.n_heads, cfg.d_head, cfg.d_model)))\n", + " self.b_Q = nn.Parameter(t.zeros((cfg.n_heads, cfg.d_head)))\n", + " self.b_K = nn.Parameter(t.zeros((cfg.n_heads, cfg.d_head)))\n", + " self.b_V = nn.Parameter(t.zeros((cfg.n_heads, cfg.d_head)))\n", + " self.b_O = nn.Parameter(t.zeros((cfg.d_model)))\n", + " nn.init.normal_(self.W_Q, std=self.cfg.init_range)\n", + " nn.init.normal_(self.W_K, std=self.cfg.init_range)\n", + " nn.init.normal_(self.W_V, std=self.cfg.init_range)\n", + " nn.init.normal_(self.W_O, std=self.cfg.init_range)\n", + " self.register_buffer(\"IGNORE\", t.tensor(-1e5, dtype=t.float32, device=device))\n", + "\n", + " def forward(\n", + " self,\n", + " normalized_resid_pre: Float[Tensor, \"batch posn d_model\"],\n", + " kv_cache_entry: KeyValueCacheTensor | None = None,\n", + " ) -> tuple[\n", + " Float[Tensor, \"batch posn d_model\"],\n", + " KeyValueCacheTensor | None\n", + " ]:\n", + " '''\n", + " Returns the result of applying attention layer to normalized_resid_pre, as well as\n", + " the new cached key and value vectors (which we get from concatenating the old cached\n", + " ones with the new key and value vectors).\n", + " '''\n", + " # Calculate the new query, key and value vectors\n", + " q = einops.einsum(\n", + " normalized_resid_pre, self.W_Q,\n", + " \"batch posn d_model, nheads d_model d_head -> batch posn nheads d_head\"\n", + " ) + self.b_Q\n", + " k = einops.einsum(\n", + " normalized_resid_pre, self.W_K,\n", + " \"batch posn d_model, nheads d_model d_head -> batch posn nheads d_head\"\n", + " ) + self.b_K\n", + " v = einops.einsum(\n", + " normalized_resid_pre, self.W_V,\n", + " \"batch posn d_model, nheads d_model d_head -> batch posn nheads d_head\"\n", + " ) + self.b_V\n", + "\n", + " # If cache_entry is not None, this means we use the previous key and value vectors\n", + " # Also we'll need to get a new cache entry which will be used later to construct a new cache\n", + " if kv_cache_entry is not None:\n", + " k = t.concat([kv_cache_entry[0], k], dim=1)\n", + " v = t.concat([kv_cache_entry[1], v], dim=1)\n", + " kv_cache_entry = t.stack([k, v])\n", + "\n", + " # Calculate attention scores, then scale and mask, and apply softmax to get probabilities\n", + " attn_scores = einops.einsum(\n", + " q, k,\n", + " \"batch posn_Q nheads d_head, batch posn_K nheads d_head -> batch nheads posn_Q posn_K\"\n", + " )\n", + " attn_scores_masked = self.apply_causal_mask(attn_scores / self.cfg.d_head ** 0.5)\n", + " attn_pattern = attn_scores_masked.softmax(-1)\n", + "\n", + " # Take weighted sum of value vectors, according to attention probabilities\n", + " z = einops.einsum(\n", + " v, attn_pattern,\n", + " \"batch posn_K nheads d_head, batch nheads posn_Q posn_K -> batch posn_Q nheads d_head\"\n", + " )\n", + "\n", + " # Calculate output (by applying matrix W_O and summing over heads, then adding bias b_O)\n", + " out = einops.einsum(\n", + " z, self.W_O,\n", + " \"batch posn_Q nheads d_head, nheads d_head d_model -> batch posn_Q d_model\"\n", + " ) + self.b_O\n", + "\n", + " return out, kv_cache_entry\n", + "\n", + " def apply_causal_mask(\n", + " self, attn_scores: Float[Tensor, \"batch n_heads query_pos key_pos\"]\n", + " ) -> Float[Tensor, \"batch n_heads query_pos key_pos\"]:\n", + " '''\n", + " Here, attn_scores have shape (batch, n_heads, query_pos, key_pos), where query_pos represents the\n", + " new (non-cached) positions, and key_pos represent all the positions (cached and non-cached).\n", + "\n", + " So when we create our mask, the query indices and key indices will both go up to the same value\n", + " (the full sequence length), but the query indices will start at >0.\n", + " '''\n", + " new_seq_len, full_seq_len = attn_scores.shape[-2:]\n", + " assert new_seq_len <= full_seq_len\n", + " q_posn = einops.repeat(attn_scores.new_tensor(range(full_seq_len-new_seq_len, full_seq_len)), \"q -> q k\", k=full_seq_len)\n", + " k_posn = einops.repeat(attn_scores.new_tensor(range(full_seq_len)), \"k -> q k\", q=new_seq_len)\n", + " mask = q_posn < k_posn\n", + " attn_scores = attn_scores.masked_fill(mask, self.IGNORE)\n", + " return attn_scores\n", + "\n", + "\n", + "class TransformerBlock(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.ln1 = LayerNorm(cfg)\n", + " self.attn = Attention(cfg)\n", + " self.ln2 = LayerNorm(cfg)\n", + " self.mlp = MLP(cfg)\n", + "\n", + " def forward(\n", + " self,\n", + " resid_pre: Float[Tensor, \"batch position d_model\"],\n", + " kv_cache_entry: KeyValueCacheTensor | None = None,\n", + " ) -> Float[Tensor, \"batch position d_model\"]:\n", + "\n", + " attn_out, kv_cache_entry = self.attn(self.ln1(resid_pre), kv_cache_entry)\n", + " resid_mid = attn_out + resid_pre\n", + " resid_post = self.mlp(self.ln2(resid_mid)) + resid_mid\n", + " return resid_post, kv_cache_entry\n", + "\n", + "\n", + "class DemoTransformer(nn.Module):\n", + " def __init__(self, cfg: Config):\n", + " super().__init__()\n", + " self.cfg = cfg\n", + " self.embed = Embed(cfg)\n", + " self.pos_embed = PosEmbed(cfg)\n", + " self.blocks = nn.ModuleList([TransformerBlock(cfg) for _ in range(cfg.n_layers)])\n", + " self.ln_final = LayerNorm(cfg)\n", + " self.unembed = Unembed(cfg)\n", + "\n", + " def forward(\n", + " self,\n", + " tokens: Int[Tensor, \"batch seq_pos\"],\n", + " kv_cache: KeyValueCache | None = None\n", + " ) -> Float[Tensor, \"batch position d_vocab\"]:\n", + "\n", + " using_kv_cache = kv_cache is not None\n", + "\n", + " if using_kv_cache:\n", + " # If using kv_cache, then we only need to pass forward the newest tokens\n", + " # Remember to add positional offset!\n", + " n_cached_tokens = kv_cache.seq_len\n", + " tokens = tokens[:, n_cached_tokens:]\n", + " residual = self.embed(tokens) + self.pos_embed(tokens, n_cached_tokens)\n", + " else:\n", + " # If not using cache, turn it into a list of None's (so we can iterate through it)\n", + " kv_cache = [None for _ in range(self.cfg.n_layers)]\n", + " residual = self.embed(tokens) + self.pos_embed(tokens)\n", + "\n", + " # Apply all layers, and create a (new) kv_cache from the key & value vectors\n", + " new_kv_cache_entries: list[KeyValueCacheTensor] = []\n", + " for block, kv_cache_entry in zip(self.blocks, kv_cache):\n", + " residual, kv_cache_entry = block(residual, kv_cache_entry)\n", + " if using_kv_cache: new_kv_cache_entries.append(kv_cache_entry)\n", + "\n", + " logits = self.unembed(self.ln_final(residual))\n", + "\n", + " if using_kv_cache:\n", + " return logits, KeyValueCache(t.stack(new_kv_cache_entries))\n", + " else:\n", + " return logits, None\n", + "\n", + "\n", + "tokens = reference_gpt2.to_tokens(reference_text).to(device)\n", + "logits, cache = reference_gpt2.run_with_cache(tokens)\n", + "\n", + "tests.rand_int_test(PosEmbed, [2, 4])\n", + "tests.load_gpt2_test(PosEmbed, reference_gpt2.pos_embed, tokens)\n", + "tests.rand_float_test(Attention, [2, 4, 768])\n", + "tests.load_gpt2_test(Attention, reference_gpt2.blocks[0].attn, cache[\"normalized\", 0, \"ln1\"])\n", + "tests.rand_float_test(TransformerBlock, [2, 4, 768])\n", + "tests.load_gpt2_test(TransformerBlock, reference_gpt2.blocks[0], cache[\"resid_pre\", 0])\n", + "tests.rand_int_test(DemoTransformer, [2, 4])\n", + "tests.load_gpt2_test(DemoTransformer, reference_gpt2, tokens)\n", + "```\n", + "\n", + "
\n", + "\n", + "
\n", + "New sampling function\n", + "\n", + "```python\n", + "@t.inference_mode()\n", + "def sample_with_cache(\n", + " self: TransformerSampler,\n", + " prompt: str,\n", + " max_tokens_generated=100,\n", + " kv_cache: KeyValueCache | None = None,\n", + " verbose=False,\n", + " seed: int | None = None,\n", + " **kwargs\n", + ") -> str:\n", + "\n", + " self.model.eval()\n", + " input_ids = self.tokenizer.encode(prompt, return_tensors=\"pt\").to(device)[0]\n", + " if seed is not None:\n", + " np.random.seed(seed)\n", + " t.manual_seed(seed)\n", + "\n", + " for i in tqdm(range(max_tokens_generated)):\n", + " # Get new logits (make sure we don't pass in more tokens than the model's context length)\n", + " logits, kv_cache = self.model(input_ids[None, -self.cfg.n_ctx:], kv_cache)\n", + " # We only take logits for the last token, because this is what we're sampling\n", + " logits = logits[0, -1]\n", + " # Get next token (as a tensor of size (1, 1) so we can concat it to input_ids)\n", + " next_token = t.tensor([TransformerSampler.sample_next_token(input_ids, logits, **kwargs)], device=device)\n", + " # Create new input ids string, with shape (1, old_seq_len + 1)\n", + " input_ids = t.cat([input_ids, next_token], dim=-1)\n", + " # Print out results, if required\n", + " if verbose:\n", + " print(self.tokenizer.decode(input_ids), end=\"\\r\")\n", + " # If our new token was the end-of-text token, stop\n", + " if next_token == getattr(self.tokenizer, \"eos_token_id\", None):\n", + " break\n", + "\n", + " return self.tokenizer.decode(input_ids)\n", + "\n", + "\n", + "TransformerSampler.sample = sample_with_cache\n", + "```\n", + "
\n", + "\n", + "
\n", + "Code to verify that the same output is being produced by cache and no-cache versions (and to compare speeds)\n", + "\n", + "```python\n", + "device = t.device(\"cuda\") # can also try \"cpu\"\n", + "\n", + "model = DemoTransformer(Config()).to(device)\n", + "model.load_state_dict(reference_gpt2.state_dict(), strict=False);\n", + "\n", + "initial_text = \"Eliezer Shlomo Yudkowsky (born September 11, 1979) is an American decision and artificial intelligence (AI) theorist and writer, best known for\"\n", + "# input_ids = tokenizer.encode(initial_text, return_tensors=\"pt\").squeeze()\n", + "\n", + "sampler = TransformerSampler(model, tokenizer)\n", + "\n", + "# Run the noncached version\n", + "t0 = time.time()\n", + "text = sampler.sample(\n", + " initial_text,\n", + " temperature=0.7,\n", + " top_p=0.95,\n", + " seed=0,\n", + ")\n", + "print(f\"Time taken (without cache): {time.time() - t0:.2f} seconds\")\n", + "rprint(f\"Model output:\\n\\n[bold dark_orange]{text}[/]\")\n", + "\n", + "# Run the cached version\n", + "t0 = time.time()\n", + "text_with_cache = sampler.sample(\n", + " initial_text,\n", + " temperature=0.7,\n", + " top_p=0.95,\n", + " seed=0,\n", + " kv_cache=KeyValueCache.new_empty(sampler.cfg)\n", + ")\n", + "print(f\"Time taken (with cache): {time.time() - t0:.2f} seconds\")\n", + "rprint(f\"Model output:\\n\\n[bold dark_orange]{text_with_cache}[/]\")\n", + "\n", + "# # Check they are the same\n", + "assert text == text_with_cache, \"Your outputs are different, meaning you've probably made a mistake in your cache implementation (or failed to use random seeds).\"\n", + "print(\"Tests passed!\")\n", + "```\n", + "\n", + "
" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "maVrhyG6fdfT" + }, + "source": [ + "You may find that your cache implementation provides a modest speedup, but probably not close to the `seq_len`-factor speedup you'd expect from the fact that you only compute one additional token at each step rather than all of them. Why is this? The answer is that, much like everything to do with computational and memory costs in deep learning, it's not so simple. There are a host of different factors which might be bottlenecking our model's forward pass speed. If you try this on the CPU, you should get a much more noticeable speedup.\n", + "\n", + "For a bit more on these topics, see [here](https://kipp.ly/blog/transformer-inference-arithmetic/#kv-cache)." + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "Coqz4E7ofdfU" + }, + "source": [ + "## Bonus - cached beam search\n", + "\n", + "Can you modify your beam search function to use caching?\n", + "\n", + "Depending on how you implemented your cache earlier, you might find that a different form of caching is better suited to beam search.\n", + "\n", + "Again, we've provided an example implementation in a dropdown below, which is based on the cache implementation above and the previous solution for `beam_search`.\n", + "\n", + "
\n", + "Cached beam search function\n", + "\n", + "As we touched on earlier, thanks to our modular code, not a lot needs to be changed when adding cache support.\n", + "\n", + "```python\n", + "@dataclass\n", + "class Beams:\n", + " '''Class to store beams during beam search.'''\n", + " model: DemoTransformer\n", + " tokenizer: GPT2TokenizerFast\n", + " logprob_sums: Float[Tensor, \"batch\"]\n", + " tokens: Int[Tensor, \"batch seq\"]\n", + " kv_cache: KeyValueCache | None = None\n", + "\n", + " def __getitem__(self, idx) -> \"Beams\":\n", + " '''Helpful function allowing you to take a slice of the beams object along the batch dimension.'''\n", + " return Beams(\n", + " self.model,\n", + " self.tokenizer,\n", + " self.logprob_sums[idx],\n", + " self.tokens[idx],\n", + " self.kv_cache[:, :, idx] if self.kv_cache is not None else None\n", + " )\n", + "\n", + " @property\n", + " def logprobs_and_completions(self) -> list[tuple[float, str]]:\n", + " '''Returns self as a list of logprob sums and completions (useful for getting final output).'''\n", + " return [\n", + " (logprob_sum.item(), self.tokenizer.decode(tokens))\n", + " for (logprob_sum, tokens) in zip(self.logprob_sums, self.tokens)\n", + " ]\n", + "\n", + "\n", + " def generate(self, k: int, no_repeat_ngram_size: int | None = None) -> \"Beams\":\n", + " '''\n", + " Starting from the current set of beams (i.e. self.tokens) and returns a new set of `len(self.tokens) * k` beams,\n", + " containing the best `k` continuations for each of the original beams.\n", + "\n", + " Optional argument `no_repeat_ngram_size` means your model won't generate any sequences with a repeating n-gram\n", + " of this length.\n", + " '''\n", + " # Get the output logprobs for the next token (for every sequence in current beams)\n", + " logprobs, kv_cache = self.model(self.tokens, self.kv_cache)\n", + " logprobs = logprobs[:, -1, :].log_softmax(-1)\n", + "\n", + " # Get the top `toks_per_beam` tokens for each sequence\n", + " topk_logprobs, topk_tokenIDs = self.get_topk_non_repeating(logprobs, no_repeat_ngram_size, k=k)\n", + "\n", + " # Add new logprobs & concat new tokens. When doing this, we need to add an extra `k` dimension since our current\n", + " # logprobs & tokens have shape (batch,) and (batch, seq), but our new ones both have shape (batch, k)\n", + " new_logprob_sums = einops.repeat(self.logprob_sums, \"b -> b k\", k=k) + topk_logprobs\n", + " new_tokens = t.concat([einops.repeat(self.tokens, \"b s -> b k s\", k=k), topk_tokenIDs.unsqueeze(-1)], dim=-1)\n", + "\n", + " return Beams(self.model, self.tokenizer, new_logprob_sums.flatten(), new_tokens.flatten(0, 1), new_kv_cache)\n", + "\n", + "\n", + " def filter(self, k: int) -> tuple[\"Beams\", \"Beams\"]:\n", + " '''\n", + " Returns:\n", + " best_beams: Beams\n", + " filtered version of self, containing all best `k` which are also not terminated.\n", + " early_terminations: Beams\n", + " filtered version of self, containing all best `k` which are also terminated.\n", + " '''\n", + " # Get the indices of top `k` beams\n", + " top_beam_indices = self.logprob_sums.topk(k=k, dim=0).indices.tolist()\n", + " # Get the indices of terminated sequences\n", + " new_tokens = self.tokens[:, -1]\n", + " terminated_indices = t.nonzero(new_tokens == self.tokenizer.eos_token_id)\n", + "\n", + " # Get the indices of the `k` best sequences (some terminated, some not terminated)\n", + " best_continuing = [i for i in top_beam_indices if i not in terminated_indices]\n", + " best_terminated = [i for i in top_beam_indices if i in terminated_indices]\n", + "\n", + " # Return the beam objects from these indices\n", + " return self[best_continuing], self[best_terminated]\n", + "\n", + "\n", + " def get_topk_non_repeating(\n", + " self,\n", + " logprobs: Float[Tensor, \"batch d_vocab\"],\n", + " no_repeat_ngram_size: int | None,\n", + " k: int,\n", + " ) -> tuple[Float[Tensor, \"k\"], Int[Tensor, \"k\"]]:\n", + " \"\"\"\n", + " logprobs:\n", + " tensor of the log-probs for the next token\n", + " no_repeat_ngram_size:\n", + " size of ngram to avoid repeating\n", + " k:\n", + " number of top logits to return, for each beam in our collection\n", + "\n", + " Returns:\n", + " equivalent to the output of `logprobs.topk(dim=-1)`, but makes sure that no returned tokens would produce an\n", + " ngram of size `no_repeat_ngram_size` which has already appeared in `self.tokens`.\n", + " \"\"\"\n", + " batch, seq_len = self.tokens.shape\n", + "\n", + " # If completion isn't long enough for a repetition, or we have no restrictions, just return topk\n", + " if (no_repeat_ngram_size is not None) and (seq_len > no_repeat_ngram_size - 1):\n", + " # Otherwise, we need to check for ngram repetitions\n", + " # First, get the most recent `no_repeat_ngram_size-1` tokens\n", + " last_ngram_prefix = self.tokens[:, seq_len - (no_repeat_ngram_size - 1) :]\n", + " # Next, find all the tokens we're not allowed to generate, by checking all past ngrams for a match\n", + " for i in range(seq_len - (no_repeat_ngram_size - 1)):\n", + " ngrams = self.tokens[:, i : i + no_repeat_ngram_size] # (batch, ngram)\n", + " ngrams_are_repeated = (ngrams[:, :-1] == last_ngram_prefix).all(-1) # (batch,)\n", + " ngram_end_tokens = ngrams[:, [-1]] # (batch, 1)\n", + " # Fill logprobs with neginf wherever the ngrams are repeated\n", + " logprobs[range(batch), ngram_end_tokens] = t.where(\n", + " ngrams_are_repeated, -1.0e4, logprobs[range(batch), ngram_end_tokens]\n", + " )\n", + "\n", + " # Finally, get our actual tokens\n", + " return logprobs.topk(k=k, dim=-1)\n", + "\n", + " def print(self, title=\"Best completions\", max_print_chars=80) -> None:\n", + " '''\n", + " Prints out a set of sequences with their corresponding logitsums.\n", + " '''\n", + " if len(self.tokens) == 0:\n", + " return\n", + " table = Table(\"logitsum\", \"completion\", title=title)\n", + " for logprob_sum, tokens in zip(self.logprob_sums, self.tokens):\n", + " text = self.tokenizer.decode(tokens)\n", + " if len(repr(text)) > max_print_chars:\n", + " text = text[:int(0.3 * max_print_chars)] + \" ... \" + text[-int(0.7 * max_print_chars):]\n", + " table.add_row(f\"{logprob_sum:>8.3f}\", repr(text))\n", + " rprint(table)\n", + "\n", + "\n", + " @t.inference_mode()\n", + " def beam_search(\n", + " self,\n", + " prompt: str,\n", + " num_return_sequences: int,\n", + " num_beams: int,\n", + " max_new_tokens: int,\n", + " no_repeat_ngram_size: int | None = None,\n", + " kv_cache: KeyValueCache | None = None,\n", + " ) -> list[tuple[float, Tensor]]:\n", + " '''\n", + " Implements a beam search, by repeatedly performing the `generate` and `filter` steps (starting from the initial\n", + " prompt) until either of the two stopping criteria are met: (1) we've generated `max_new_tokens` tokens, or (2)\n", + " we've generated `num_returns_sequences` terminating sequences.\n", + " '''\n", + " assert num_return_sequences <= num_beams\n", + " self.model.eval()\n", + "\n", + " tokens = self.tokenizer.encode(prompt, return_tensors=\"pt\").to(device)\n", + "\n", + " final_logprobs_and_completions = [] # we add to this list as we get terminated beams\n", + " best_beams = Beams(self.model, self.tokenizer, t.tensor([0.0]).to(device), tokens) # start with just 1 beam\n", + "\n", + " for _ in tqdm(range(max_new_tokens)):\n", + " # Generate & filter beams\n", + " best_beams = best_beams.generate(k=num_beams, no_repeat_ngram_size=no_repeat_ngram_size)\n", + " best_beams, best_beams_terminated = best_beams.filter(k=num_beams)\n", + "\n", + " # Add terminated beams to our list, and return early if we have enough\n", + " final_logprobs_and_completions.extend(best_beams_terminated.logprobs_and_completions)\n", + " if len(final_logprobs_and_completions) >= num_return_sequences:\n", + " return final_logprobs_and_completions[:num_return_sequences]\n", + "\n", + " # Return terminated beams plus the best ongoing beams of length `orig_len + max_new_tokens`\n", + " final_logprobs_and_completions.extend(best_beams.logprobs_and_completions)\n", + " return final_logprobs_and_completions[:num_return_sequences]\n", + "\n", + "\n", + "```\n", + "\n", + "
\n", + "\n", + "
\n", + "Code to verify that the same output is being produced by cache and no-cache versions (and to compare speeds)\n", + "\n", + "```python\n", + "prompt = \"For you, the day Bison graced your village was the most important day of your life. But for me, it was\"\n", + "orig_len = len(tokenizer.encode(prompt))\n", + "\n", + "beam_search_kwargs = dict(\n", + " prompt=prompt,\n", + " num_return_sequences=3,\n", + " num_beams=20,\n", + " max_new_tokens=60,\n", + " no_repeat_ngram_size=2,\n", + " verbose=False\n", + ")\n", + "\n", + "sampler = TransformerSampler(model, tokenizer)\n", + "\n", + "# Run the noncached version\n", + "t0 = time.time()\n", + "final_logitsums_and_completions = sampler.beam_search(**beam_search_kwargs)\n", + "logprob_sum, text = final_logitsums_and_completions[0]\n", + "avg_logprob_as_prob = t.tensor(logprob_sum / (len(tokenizer.encode(text)) - orig_len)).exp().item()\n", + "print(f\"Time (without cache): {time.time() - t0:.2f} seconds\")\n", + "print(f\"Avg logprob (expressed as a probability) = {avg_logprob_as_prob:.3f}\")\n", + "rprint(f\"Output:\\n\\n[bold dark_orange]{text}[/]\\n\\n\")\n", + "\n", + "# Run the cached version\n", + "t0 = time.time()\n", + "beam_search_kwargs[\"kv_cache\"] = KeyValueCache.new_empty(model.cfg)\n", + "final_logitsums_and_completions = sampler.beam_search(**beam_search_kwargs)\n", + "logprob_sum, text_with_cache = final_logitsums_and_completions[0]\n", + "avg_logprob_as_prob = t.tensor(logprob_sum / (len(tokenizer.encode(text)) - orig_len)).exp().item()\n", + "print(f\"Time (with cache): {time.time() - t0:.2f} seconds\")\n", + "print(f\"Avg logprob (as probability) = {avg_logprob_as_prob:.3f}\", end=\"\")\n", + "rprint(f\"Output:\\n\\n[bold dark_orange]{text_with_cache}[/]\\n\\n\")\n", + "\n", + "# Check they are the same\n", + "assert text == text_with_cache, \"Your outputs are different, meaning you've probably made a mistake in your cache implementation.\"\n", + "print(\"Tests passed!\")\n", + "```\n", + "\n", + "
" + ] + } + ], + "metadata": { + "language_info": { + "name": "python" + }, + "colab": { + "provenance": [] + }, + "kernelspec": { + "name": "python3", + "display_name": "Python 3" + }, + "widgets": { + "application/vnd.jupyter.widget-state+json": { + "version_major": 2, + "version_minor": 0, + "state": { + "a76388bdb84c422f833a0015ba4a7ab5": { + "model_module": "@jupyter-widgets/controls", + "model_name": "HBoxModel", + "model_module_version": "1.5.0", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HBoxModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HBoxView", + "box_style": "", + "children": [ + "IPY_MODEL_de3deb7c5ee2471186c2d3d354ef8014", + "IPY_MODEL_5faad48740fd49acb8869dfb35bbc9fb", + "IPY_MODEL_263b3cbc12c746df95145835ba3e39c2" + ], + "layout": "IPY_MODEL_3b3c2c97d7f54f1c9c7bf92d8157c1c6" + } + }, + "de3deb7c5ee2471186c2d3d354ef8014": { + "model_module": "@jupyter-widgets/controls", + "model_name": "HTMLModel", + "model_module_version": "1.5.0", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_60c4eb83080d4333835df32ae061af5e", + "placeholder": "​", + "style": "IPY_MODEL_53891efb605249ad83402d96a3c67dae", + "value": "Loading weights: 100%" + } + }, + "5faad48740fd49acb8869dfb35bbc9fb": { + "model_module": "@jupyter-widgets/controls", + "model_name": "FloatProgressModel", + "model_module_version": "1.5.0", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "FloatProgressModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "ProgressView", + "bar_style": "success", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_a90b519263d844a883816df8a9c8d7cb", + "max": 148, + "min": 0, + "orientation": "horizontal", + "style": "IPY_MODEL_c64111af122341fbace7fb2195d0b519", + "value": 148 + } + }, + "263b3cbc12c746df95145835ba3e39c2": { + "model_module": "@jupyter-widgets/controls", + "model_name": "HTMLModel", + "model_module_version": "1.5.0", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_f7667933820842af82625dfa89b59748", + "placeholder": "​", + "style": "IPY_MODEL_0ee53673a92b4d0a8398e9128c81c988", + "value": " 148/148 [00:00<00:00,  3.34it/s]" + } + }, + "3b3c2c97d7f54f1c9c7bf92d8157c1c6": { + "model_module": "@jupyter-widgets/base", + "model_name": "LayoutModel", + "model_module_version": "1.2.0", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "60c4eb83080d4333835df32ae061af5e": { + "model_module": "@jupyter-widgets/base", + "model_name": "LayoutModel", + "model_module_version": "1.2.0", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "53891efb605249ad83402d96a3c67dae": { + "model_module": "@jupyter-widgets/controls", + "model_name": "DescriptionStyleModel", + "model_module_version": "1.5.0", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "a90b519263d844a883816df8a9c8d7cb": { + "model_module": "@jupyter-widgets/base", + "model_name": "LayoutModel", + "model_module_version": "1.2.0", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "c64111af122341fbace7fb2195d0b519": { + "model_module": "@jupyter-widgets/controls", + "model_name": "ProgressStyleModel", + "model_module_version": "1.5.0", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "ProgressStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "bar_color": null, + "description_width": "" + } + }, + "f7667933820842af82625dfa89b59748": { + "model_module": "@jupyter-widgets/base", + "model_name": "LayoutModel", + "model_module_version": "1.2.0", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "0ee53673a92b4d0a8398e9128c81c988": { + "model_module": "@jupyter-widgets/controls", + "model_name": "DescriptionStyleModel", + "model_module_version": "1.5.0", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + } + } + } + } + }, + "nbformat": 4, + "nbformat_minor": 0 +} \ No newline at end of file