mishformer_lens_v1

建立於 差異永不過期
117 刪除
362 行
218 新增
345 行
# %% [markdown]
# %% [markdown]
# [![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](https://colab.research.google.com/github/TransformerLensOrg/TransformerLens/blob/main/demos/Exploratory_Analysis_Demo.ipynb)
# [![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](https://colab.research.google.com/github/TransformerLensOrg/TransformerLens/blob/main/demos/Exploratory_Analysis_Demo.ipynb)
# %% [markdown]
# %% [markdown]
# # Exploratory Analysis Demo
# # Exploratory Analysis Demo
#
#
# (This is a MishformerLens version of the TransformerLens Exploratory Analysis Demo; and all lines edited have a `# MishformerLens:` comment above them)
#
# This notebook demonstrates how to use the
# This notebook demonstrates how to use the
# [TransformerLens](https://github.com/TransformerLensOrg/TransformerLens/) library to perform exploratory
# [TransformerLens](https://github.com/TransformerLensOrg/TransformerLens/) library to perform exploratory
# analysis. The notebook tries to replicate the analysis of the Indirect Object Identification circuit
# analysis. The notebook tries to replicate the analysis of the Indirect Object Identification circuit
# in the [Interpretability in the Wild](https://arxiv.org/abs/2211.00593) paper.
# in the [Interpretability in the Wild](https://arxiv.org/abs/2211.00593) paper.
# %% [markdown]
# %% [markdown]
# ## Tips for Reading This
# ## Tips for Reading This
#
#
# * If running in Google Colab, go to Runtime > Change Runtime Type and select GPU as the hardware
# * If running in Google Colab, go to Runtime > Change Runtime Type and select GPU as the hardware
# accelerator.
# accelerator.
# * Look up unfamiliar terms in [the mech interp explainer](https://neelnanda.io/glossary)
# * Look up unfamiliar terms in [the mech interp explainer](https://neelnanda.io/glossary)
# * You can run all this code for yourself
# * You can run all this code for yourself
# * The graphs are interactive
# * The graphs are interactive
# * Use the table of contents pane in the sidebar to navigate (in Colab) or VSCode's "Outline" in the
# * Use the table of contents pane in the sidebar to navigate (in Colab) or VSCode's "Outline" in the
#   explorer tab.
#   explorer tab.
# * Collapse irrelevant sections with the dropdown arrows
# * Collapse irrelevant sections with the dropdown arrows
# * Search the page using the search in the sidebar (with Colab) not CTRL+F
# * Search the page using the search in the sidebar (with Colab) not CTRL+F
# %% [markdown]
# %% [markdown]
# ## Setup
# ## Setup
# %% [markdown]
# %% [markdown]
# ### Environment Setup (ignore)
# ### Environment Setup (ignore)
# %% [markdown]
# %% [markdown]
# **You can ignore this part:** It's just for use internally to setup the tutorial in different
# **You can ignore this part:** It's just for use internally to setup the tutorial in different
# environments. You can delete this section if using in your own repo.
# environments. You can delete this section if using in your own repo.
# %%
# %%


# Detect if we're running in Google Colab
# Detect if we're running in Google Colab
try:
try:
import google.colab
    import google.colab
IN_COLAB = True
    IN_COLAB = True
print("Running as a Colab notebook")
    print("Running as a Colab notebook")
except:
except:
IN_COLAB = False
    IN_COLAB = False


# Install if in Colab
# Install if in Colab
if IN_COLAB:
if IN_COLAB:
%pip install transformer_lens
    %pip install transformer_lens
%pip install circuitsvis
    %pip install circuitsvis
# Install a faster Node version
    # Install a faster Node version
!curl -fsSL https://deb.nodesource.com/setup_16.x | sudo -E bash -; sudo apt-get install -y nodejs  # noqa
    !curl -fsSL https://deb.nodesource.com/setup_16.x | sudo -E bash -; sudo apt-get install -y nodejs  # noqa


# Hot reload in development mode & not running on the CD
# Hot reload in development mode & not running on the CD
if not IN_COLAB:
if not IN_COLAB:
from IPython import get_ipython
    from IPython import get_ipython
ip = get_ipython()
    ip = get_ipython()
# MishformerLens: for quality of life I always want this enabled unconditionally here:
    if not ip.extension_manager.loaded:
# if not ip.extension_manager.loaded:
        ip.extension_manager.load('autoreload')
# ip.extension_manager.load('autoreload')
        %autoreload 2
# %autoreload 2
ipython = get_ipython()
if ipython is not None:
ipython.magic(u"%load_ext autoreload")
ipython.magic(u"%autoreload 2")



# %% [markdown]
# %% [markdown]
# ### Imports
# ### Imports
# %%
# %%
from functools import partial
from functools import partial
from typing import List, Optional, Union
from typing import List, Optional, Union


import einops
import einops
import numpy as np
import numpy as np
import plotly.express as px
import plotly.express as px
import plotly.io as pio
import plotly.io as pio
import torch
import torch
from circuitsvis.attention import attention_heads
from circuitsvis.attention import attention_heads
from fancy_einsum import einsum
from fancy_einsum import einsum
from IPython.display import HTML, IFrame
from IPython.display import HTML, IFrame
from jaxtyping import Float
from jaxtyping import Float


import transformer_lens.utils as utils
import transformer_lens.utils as utils
from transformer_lens import ActivationCache
from transformer_lens import ActivationCache, HookedTransformer

# MishformerLens: this replaces the import of HookedTransformer from transformer_lens with our version
from mishformer_lens import HookedTransformer

# %% [markdown]
# %% [markdown]
# ### PyTorch Setup
# ### PyTorch Setup
# %% [markdown]
# %% [markdown]
# We turn automatic differentiation off, to save GPU memory, as this notebook focuses on model inference not model training.
# We turn automatic differentiation off, to save GPU memory, as this notebook focuses on model inference not model training.
# %%
# %%
torch.set_grad_enabled(False)
torch.set_grad_enabled(False)
print("Disabled automatic differentiation")
print("Disabled automatic differentiation")
# %% [markdown]
# %% [markdown]
# ### Plotting Helper Functions (ignore)
# ### Plotting Helper Functions (ignore)
# %% [markdown]
# %% [markdown]
# Some plotting helper functions are included here (for simplicity).
# Some plotting helper functions are included here (for simplicity).
# %%
# %%
def imshow(tensor, **kwargs):
def imshow(tensor, **kwargs):
px.imshow(
    px.imshow(
utils.to_numpy(tensor),
        utils.to_numpy(tensor),
color_continuous_midpoint=0.0,
        color_continuous_midpoint=0.0,
color_continuous_scale="RdBu",
        color_continuous_scale="RdBu",
**kwargs,
        **kwargs,
).show()
    ).show()




def line(tensor, **kwargs):
def line(tensor, **kwargs):
px.line(
    px.line(
y=utils.to_numpy(tensor),
        y=utils.to_numpy(tensor),
**kwargs,
        **kwargs,
).show()
    ).show()




def scatter(x, y, xaxis="", yaxis="", caxis="", **kwargs):
def scatter(x, y, xaxis="", yaxis="", caxis="", **kwargs):
x = utils.to_numpy(x)
    x = utils.to_numpy(x)
y = utils.to_numpy(y)
    y = utils.to_numpy(y)
px.scatter(
    px.scatter(
y=y,
        y=y,
x=x,
        x=x,
labels={"x": xaxis, "y": yaxis, "color": caxis},
        labels={"x": xaxis, "y": yaxis, "color": caxis},
**kwargs,
        **kwargs,
).show()
    ).show()
# %% [markdown]
# %% [markdown]
# ## Introduction
# ## Introduction
#
#
# This is a demo notebook for [TransformerLens](https://github.com/TransformerLensOrg/TransformerLens), a library for mechanistic interpretability of GPT-2 style transformer language models. A core design principle of the library is to enable exploratory analysis - one of the most fun parts of mechanistic interpretability compared to normal ML is the extremely short feedback loops! The point of this library is to keep the gap between having an experiment idea and seeing the results as small as possible, to make it easy for **research to feel like play** and to enter a flow state.
# This is a demo notebook for [TransformerLens](https://github.com/TransformerLensOrg/TransformerLens), a library for mechanistic interpretability of GPT-2 style transformer language models. A core design principle of the library is to enable exploratory analysis - one of the most fun parts of mechanistic interpretability compared to normal ML is the extremely short feedback loops! The point of this library is to keep the gap between having an experiment idea and seeing the results as small as possible, to make it easy for **research to feel like play** and to enter a flow state.
#
#
# The goal of this notebook is to demonstrate what exploratory analysis looks like in practice with the library. I use my standard toolkit of basic mechanistic interpretability techniques to try interpreting a real circuit in GPT-2 small. Check out [the main demo](https://colab.research.google.com/github/TransformerLensOrg/TransformerLens/blob/main/demos/Main_Demo.ipynb) for an introduction to the library and how to use it.
# The goal of this notebook is to demonstrate what exploratory analysis looks like in practice with the library. I use my standard toolkit of basic mechanistic interpretability techniques to try interpreting a real circuit in GPT-2 small. Check out [the main demo](https://colab.research.google.com/github/TransformerLensOrg/TransformerLens/blob/main/demos/Main_Demo.ipynb) for an introduction to the library and how to use it.
#
#
# Stylistically, I will go fairly slowly and explain in detail what I'm doing and why, aiming to help convey how to do this kind of research yourself! But the code itself is written to be simple and generic, and easy to copy and paste into your own projects for different tasks and models.
# Stylistically, I will go fairly slowly and explain in detail what I'm doing and why, aiming to help convey how to do this kind of research yourself! But the code itself is written to be simple and generic, and easy to copy and paste into your own projects for different tasks and models.
#
#
# Details tags contain asides, flavour + interpretability intuitions. These are more in the weeds and you don't need to read them or understand them, but they're helpful if you want to learn how to do mechanistic interpretability yourself! I star the ones I think are most important.
# Details tags contain asides, flavour + interpretability intuitions. These are more in the weeds and you don't need to read them or understand them, but they're helpful if you want to learn how to do mechanistic interpretability yourself! I star the ones I think are most important.
# <details><summary>(*) Example details tag</summary>Example aside!</details>
# <details><summary>(*) Example details tag</summary>Example aside!</details>
# %% [markdown]
# %% [markdown]
# ### Indirect Object Identification
# ### Indirect Object Identification
#
#
# The first step when trying to reverse engineer a circuit in a model is to identify *what* capability
# The first step when trying to reverse engineer a circuit in a model is to identify *what* capability
# I want to reverse engineer. Indirect Object Identification is a task studied in Redwood Research's
# I want to reverse engineer. Indirect Object Identification is a task studied in Redwood Research's
# excellent [Interpretability in the Wild](https://arxiv.org/abs/2211.00593) paper (see [my interview
# excellent [Interpretability in the Wild](https://arxiv.org/abs/2211.00593) paper (see [my interview
# with the authors](https://www.youtube.com/watch?v=gzwj0jWbvbo) or [Kevin Wang's Twitter
# with the authors](https://www.youtube.com/watch?v=gzwj0jWbvbo) or [Kevin Wang's Twitter
# thread](https://threadreaderapp.com/thread/1587601532639494146.html) for an overview). The task is
# thread](https://threadreaderapp.com/thread/1587601532639494146.html) for an overview). The task is
# to complete sentences like "After John and Mary went to the shops, John gave a bottle of milk to"
# to complete sentences like "After John and Mary went to the shops, John gave a bottle of milk to"
# with " Mary" rather than " John".
# with " Mary" rather than " John".
#
#
# In the paper they rigorously reverse engineer a 26 head circuit, with 7 separate categories of heads
# In the paper they rigorously reverse engineer a 26 head circuit, with 7 separate categories of heads
# used to perform this capability. Their rigorous methods are fairly involved, so in this notebook,
# used to perform this capability. Their rigorous methods are fairly involved, so in this notebook,
# I'm going to skimp on rigour and instead try to speed run the process of finding suggestive evidence
# I'm going to skimp on rigour and instead try to speed run the process of finding suggestive evidence
# for this circuit!
# for this circuit!
#
#
# The circuit they found roughly breaks down into three parts:
# The circuit they found roughly breaks down into three parts:
# 1. Identify what names are in the sentence
# 1. Identify what names are in the sentence
# 2. Identify which names are duplicated
# 2. Identify which names are duplicated
# 3. Predict the name that is *not* duplicated
# 3. Predict the name that is *not* duplicated
# %% [markdown]
# %% [markdown]
# The first step is to load in our model, GPT-2 Small, a 12 layer and 80M parameter transformer with `HookedTransformer.from_pretrained`. The various flags are simplifications that preserve the model's output but simplify its internals.
# The first step is to load in our model, GPT-2 Small, a 12 layer and 80M parameter transformer with `HookedTransformer.from_pretrained`. The various flags are simplifications that preserve the model's output but simplify its internals.
# %%
# %%
# NBVAL_IGNORE_OUTPUT
# NBVAL_IGNORE_OUTPUT
model = HookedTransformer.from_pretrained(
model = HookedTransformer.from_pretrained(
"gpt2-small",
    "gpt2-small",
# MishformerLens: center_unembed and center_writing_weights and fold_ln and refactor_factored_attn_matrices are TODO(v1)
    center_unembed=True,
center_unembed=False,
    center_writing_weights=True,
center_writing_weights=False,
    fold_ln=True,
fold_ln=False,
    refactor_factored_attn_matrices=True,
refactor_factored_attn_matrices=False,
# MishformerLens: fold_value_biases is also a TODO(v1)
fold_value_biases=False,
# MisformerLens: we want attention patterns and get a warning without this
attn_implementation='eager',
)
)


# Get the default device used
# Get the default device used
device: torch.device = utils.get_device()
device: torch.device = utils.get_device()
# %% [markdown]
# %% [markdown]
# The next step is to verify that the model can *actually* do the task! Here we use `utils.test_prompt`, and see that the model is significantly better at predicting Mary than John!
# The next step is to verify that the model can *actually* do the task! Here we use `utils.test_prompt`, and see that the model is significantly better at predicting Mary than John!
#
#
# <details><summary>Asides:</summary>
# <details><summary>Asides:</summary>
#
#
# Note: If we were being careful, we'd want to run the model on a range of prompts and find the average performance
# Note: If we were being careful, we'd want to run the model on a range of prompts and find the average performance
#
#
# `prepend_bos` is a flag to add a BOS (beginning of sequence) to the start of the prompt. GPT-2 was not trained with this, but I find that it often makes model behaviour more stable, as the first token is treated weirdly.
# `prepend_bos` is a flag to add a BOS (beginning of sequence) to the start of the prompt. GPT-2 was not trained with this, but I find that it often makes model behaviour more stable, as the first token is treated weirdly.
# </details>
# </details>
# %%
# %%
example_prompt = "After John and Mary went to the store, John gave a bottle of milk to"
example_prompt = "After John and Mary went to the store, John gave a bottle of milk to"
example_answer = " Mary"
example_answer = " Mary"
utils.test_prompt(example_prompt, example_answer, model, prepend_bos=True)
utils.test_prompt(example_prompt, example_answer, model, prepend_bos=True)
# %% [markdown]
# %% [markdown]
# We now want to find a reference prompt to run the model on. Even though our ultimate goal is to reverse engineer how this behaviour is done in general, often the best way to start out in mechanistic interpretability is by zooming in on a concrete example and understanding it in detail, and only *then* zooming out and verifying that our analysis generalises.
# We now want to find a reference prompt to run the model on. Even though our ultimate goal is to reverse engineer how this behaviour is done in general, often the best way to start out in mechanistic interpretability is by zooming in on a concrete example and understanding it in detail, and only *then* zooming out and verifying that our analysis generalises.
#
#
# We'll run the model on 4 instances of this task, each prompt given twice - one with the first name as the indirect object, one with the second name. To make our lives easier, we'll carefully choose prompts with single token names and the corresponding names in the same token positions.
# We'll run the model on 4 instances of this task, each prompt given twice - one with the first name as the indirect object, one with the second name. To make our lives easier, we'll carefully choose prompts with single token names and the corresponding names in the same token positions.
#
#
# <details> <summary>(*) <b>Aside on tokenization</b></summary>
# <details> <summary>(*) <b>Aside on tokenization</b></summary>
#
#
# We want models that can take in arbitrary text, but models need to have a fixed vocabulary. So the solution is to define a vocabulary of **tokens** and to deterministically break up arbitrary text into tokens. Tokens are, essentially, subwords, and are determined by finding the most frequent substrings - this means that tokens vary a lot in length and frequency!
# We want models that can take in arbitrary text, but models need to have a fixed vocabulary. So the solution is to define a vocabulary of **tokens** and to deterministically break up arbitrary text into tokens. Tokens are, essentially, subwords, and are determined by finding the most frequent substrings - this means that tokens vary a lot in length and frequency!
#
#
# Tokens are a *massive* headache and are one of the most annoying things about reverse engineering language models... Different names will be different numbers of tokens, different prompts will have the relevant tokens at different positions, different prompts will have different total numbers of tokens, etc. Language models often devote significant amounts of parameters in early layers to convert inputs from tokens to a more sensible internal format (and do the reverse in later layers). You really, really want to avoid needing to think about tokenization wherever possible when doing exploratory analysis (though, of course, it's relevant later when trying to flesh out your analysis and make it rigorous!). HookedTransformer comes with several helper methods to deal with tokens: `to_tokens, to_string, to_str_tokens, to_single_token, get_token_position`
# Tokens are a *massive* headache and are one of the most annoying things about reverse engineering language models... Different names will be different numbers of tokens, different prompts will have the relevant tokens at different positions, different prompts will have different total numbers of tokens, etc. Language models often devote significant amounts of parameters in early layers to convert inputs from tokens to a more sensible internal format (and do the reverse in later layers). You really, really want to avoid needing to think about tokenization wherever possible when doing exploratory analysis (though, of course, it's relevant later when trying to flesh out your analysis and make it rigorous!). HookedTransformer comes with several helper methods to deal with tokens: `to_tokens, to_string, to_str_tokens, to_single_token, get_token_position`
#
#
# **Exercise:** I recommend using `model.to_str_tokens` to explore how the model tokenizes different strings. In particular, try adding or removing spaces at the start, or changing capitalization - these change tokenization!</details>
# **Exercise:** I recommend using `model.to_str_tokens` to explore how the model tokenizes different strings. In particular, try adding or removing spaces at the start, or changing capitalization - these change tokenization!</details>
# %%
# %%
prompt_format = [
prompt_format = [
"When John and Mary went to the shops,{} gave the bag to",
    "When John and Mary went to the shops,{} gave the bag to",
"When Tom and James went to the park,{} gave the ball to",
    "When Tom and James went to the park,{} gave the ball to",
"When Dan and Sid went to the shops,{} gave an apple to",
    "When Dan and Sid went to the shops,{} gave an apple to",
"After Martin and Amy went to the park,{} gave a drink to",
    "After Martin and Amy went to the park,{} gave a drink to",
]
]
names = [
names = [
(" Mary", " John"),
    (" Mary", " John"),
(" Tom", " James"),
    (" Tom", " James"),
(" Dan", " Sid"),
    (" Dan", " Sid"),
(" Martin", " Amy"),
    (" Martin", " Amy"),
]
]
# List of prompts
# List of prompts
prompts = []
prompts = []
# List of answers, in the format (correct, incorrect)
# List of answers, in the format (correct, incorrect)
answers = []
answers = []
# List of the token (ie an integer) corresponding to each answer, in the format (correct_token, incorrect_token)
# List of the token (ie an integer) corresponding to each answer, in the format (correct_token, incorrect_token)
answer_tokens = []
answer_tokens = []
for i in range(len(prompt_format)):
for i in range(len(prompt_format)):
for j in range(2):
    for j in range(2):
answers.append((names[i][j], names[i][1 - j]))
        answers.append((names[i][j], names[i][1 - j]))
answer_tokens.append(
        answer_tokens.append(
(
            (
model.to_single_token(answers[-1][0]),
                model.to_single_token(answers[-1][0]),
model.to_single_token(answers[-1][1]),
                model.to_single_token(answers[-1][1]),
)
            )
)
        )
# Insert the *incorrect* answer to the prompt, making the correct answer the indirect object.
        # Insert the *incorrect* answer to the prompt, making the correct answer the indirect object.
prompts.append(prompt_format[i].format(answers[-1][1]))
        prompts.append(prompt_format[i].format(answers[-1][1]))
answer_tokens = torch.tensor(answer_tokens).to(device)
answer_tokens = torch.tensor(answer_tokens).to(device)
print(prompts)
print(prompts)
print(answers)
print(answers)
# %% [markdown]
# %% [markdown]
# **Gotcha**: It's important that all of your prompts have the same number of tokens. If they're different lengths, then the position of the "final" logit where you can check logit difference will differ between prompts, and this will break the below code. The easiest solution is just to choose your prompts carefully to have the same number of tokens (you can eg add filler words like The, or newlines to start).
# **Gotcha**: It's important that all of your prompts have the same number of tokens. If they're different lengths, then the position of the "final" logit where you can check logit difference will differ between prompts, and this will break the below code. The easiest solution is just to choose your prompts carefully to have the same number of tokens (you can eg add filler words like The, or newlines to start).
#
#
# There's a range of other ways of solving this, eg you can index more intelligently to get the final logit. A better way is to just use left padding by setting `model.tokenizer.padding_side = 'left'` before tokenizing the inputs and running the model; this way, you can use something like `logits[:, -1, :]` to easily access the final token outputs without complicated indexing. TransformerLens checks the value of `padding_side` of the tokenizer internally, and if the flag is set to be `'left'`, it adjusts the calculation of absolute position embedding and causal masking accordingly.
# There's a range of other ways of solving this, eg you can index more intelligently to get the final logit. A better way is to just use left padding by setting `model.tokenizer.padding_side = 'left'` before tokenizing the inputs and running the model; this way, you can use something like `logits[:, -1, :]` to easily access the final token outputs without complicated indexing. TransformerLens checks the value of `padding_side` of the tokenizer internally, and if the flag is set to be `'left'`, it adjusts the calculation of absolute position embedding and causal masking accordingly.
#
#
# In this demo, though, we stick to using the prompts of the same number of tokens because we want to show some visualisations aggregated along the batch dimension later in the demo.
# In this demo, though, we stick to using the prompts of the same number of tokens because we want to show some visualisations aggregated along the batch dimension later in the demo.
# %%
# %%
for prompt in prompts:
for prompt in prompts:
str_tokens = model.to_str_tokens(prompt)
    str_tokens = model.to_str_tokens(prompt)
print("Prompt length:", len(str_tokens))
    print("Prompt length:", len(str_tokens))
print("Prompt as tokens:", str_tokens)
    print("Prompt as tokens:", str_tokens)
# %% [markdown]
# %% [markdown]
# We now run the model on these prompts and use `run_with_cache` to get both the logits and a cache of all internal activations for later analysis
# We now run the model on these prompts and use `run_with_cache` to get both the logits and a cache of all internal activations for later analysis
# %%
# %%
tokens = model.to_tokens(prompts, prepend_bos=True)
tokens = model.to_tokens(prompts, prepend_bos=True)


# Run the model and cache all activations
# Run the model and cache all activations
original_logits, cache = model.run_with_cache(tokens)
original_logits, cache = model.run_with_cache(tokens)

# MishformerLens: I'm leaving this in for debugging, doesn't run by default
#%%

if False:
# I'm leaving this in for debugging
from transformer_lens import HookedTransformer as TLHookedTransformer
tl_model = TLHookedTransformer.from_pretrained(
"gpt2-small",
center_unembed=False,
center_writing_weights=False,
fold_ln=False,
refactor_factored_attn_matrices=False,
fold_value_biases=False,
attn_implementation='eager',
)
tl_model.set_use_hook_mlp_in(True)
tl_original_logits, tl_cache = tl_model.run_with_cache(tokens)
for x, val in cache.items():
if val.shape != tl_cache[x].shape:
raise ValueError(f"{(x, val.shape, tl_cache[x].shape)}")

# %% [markdown]
# %% [markdown]
# We'll later be evaluating how model performance differs upon performing various interventions, so it's useful to have a metric to measure model performance. Our metric here will be the **logit difference**, the difference in logit between the indirect object's name and the subject's name (eg, `logit(Mary)-logit(John)`).
# We'll later be evaluating how model performance differs upon performing various interventions, so it's useful to have a metric to measure model performance. Our metric here will be the **logit difference**, the difference in logit between the indirect object's name and the subject's name (eg, `logit(Mary)-logit(John)`).
# %%
# %%
def logits_to_ave_logit_diff(logits, answer_tokens, per_prompt=False):
def logits_to_ave_logit_diff(logits, answer_tokens, per_prompt=False):
# Only the final logits are relevant for the answer
    # Only the final logits are relevant for the answer
final_logits = logits[:, -1, :]
    final_logits = logits[:, -1, :]
answer_logits = final_logits.gather(dim=-1, index=answer_tokens)
    answer_logits = final_logits.gather(dim=-1, index=answer_tokens)
answer_logit_diff = answer_logits[:, 0] - answer_logits[:, 1]
    answer_logit_diff = answer_logits[:, 0] - answer_logits[:, 1]
if per_prompt:
    if per_prompt:
return answer_logit_diff
        return answer_logit_diff
else:
    else:
return answer_logit_diff.mean()
        return answer_logit_diff.mean()




print(
print(
"Per prompt logit difference:",
    "Per prompt logit difference:",
logits_to_ave_logit_diff(original_logits, answer_tokens, per_prompt=True)
    logits_to_ave_logit_diff(original_logits, answer_tokens, per_prompt=True)
.detach()
    .detach()
.cpu()
    .cpu()
.round(decimals=3),
    .round(decimals=3),
)
)
original_average_logit_diff = logits_to_ave_logit_diff(original_logits, answer_tokens)
original_average_logit_diff = logits_to_ave_logit_diff(original_logits, answer_tokens)
print(
print(
"Average logit difference:",
    "Average logit difference:",
round(logits_to_ave_logit_diff(original_logits, answer_tokens).item(), 3),
    round(logits_to_ave_logit_diff(original_logits, answer_tokens).item(), 3),
)
)
# %% [markdown]
# %% [markdown]
# We see that the average logit difference is 3.5 - for context, this represents putting an $e^{3.5}\approx 33\times$ higher probability on the correct answer.
# We see that the average logit difference is 3.5 - for context, this represents putting an $e^{3.5}\approx 33\times$ higher probability on the correct answer.
# %% [markdown]
# %% [markdown]
# ## Brainstorm What's Actually Going On (Optional)
# ## Brainstorm What's Actually Going On (Optional)
#
#
# Before diving into running experiments, it's often useful to spend some time actually reasoning about how the behaviour in question could be implemented in the transformer. **This is optional, and you'll likely get the most out of engaging with this section if you have a decent understanding already of what a transformer is and how it works!**
# Before diving into running experiments, it's often useful to spend some time actually reasoning about how the behaviour in question could be implemented in the transformer. **This is optional, and you'll likely get the most out of engaging with this section if you have a decent understanding already of what a transformer is and how it works!**
#
#
# You don't have to do this and forming hypotheses after exploration is also reasonable, but I think it's often easier to explore and interpret results with some grounding in what you might find. In this particular case, I'm cheating somewhat, since I know the answer, but I'm trying to simulate the process of reasoning about it!
# You don't have to do this and forming hypotheses after exploration is also reasonable, but I think it's often easier to explore and interpret results with some grounding in what you might find. In this particular case, I'm cheating somewhat, since I know the answer, but I'm trying to simulate the process of reasoning about it!
#
#
# Note that often your hypothesis will be wrong in some ways and often be completely off. We're doing science here, and the goal is to understand how the model *actually* works, and to form true beliefs! There are two separate traps here at two extremes that it's worth tracking:
# Note that often your hypothesis will be wrong in some ways and often be completely off. We're doing science here, and the goal is to understand how the model *actually* works, and to form true beliefs! There are two separate traps here at two extremes that it's worth tracking:
# * Confusion: Having no hypotheses at all, getting a lot of data and not knowing what to do with it, and just floundering around
# * Confusion: Having no hypotheses at all, getting a lot of data and not knowing what to do with it, and just floundering around
# * Dogmatism: Being overconfident in an incorrect hypothesis and being unwilling to let go of it when reality contradicts you, or flinching away from running the experiments that might disconfirm it.
# * Dogmatism: Being overconfident in an incorrect hypothesis and being unwilling to let go of it when reality contradicts you, or flinching away from running the experiments that might disconfirm it.
#
#
# **Exercise:** Spend some time thinking through how you might imagine this behaviour being implemented in a transformer. Try to think through this for yourself before reading through my thoughts!
# **Exercise:** Spend some time thinking through how you might imagine this behaviour being implemented in a transformer. Try to think through this for yourself before reading through my thoughts!
#
#
# <details> <summary>(*) <b>My reasoning</b></summary>
# <details> <summary>(*) <b>My reasoning</b></summary>
#
#
# <h3>Brainstorming:</h3>
# <h3>Brainstorming:</h3>
#
#
# So, what's hard about the task? Let's focus on the concrete example of the first prompt, "When John and Mary went to the shops, John gave the bag to" -> " Mary".  
# So, what's hard about the task? Let's focus on the concrete example of the first prompt, "When John and Mary went to the shops, John gave the bag to" -> " Mary".  
#
#
# A good starting point is thinking though whether a tiny model could do this, eg a <a href="https://transformer-circuits.pub/2021/framework/index.html">1L Attn-Only model</a>. I'm pretty sure the answer is no! Attention is really good at the primitive operations of looking nearby, or copying information. I can believe a tiny model could figure out that at `to` it should look for names and predict that those names came next (eg the skip trigram " John...to -> John"). But it's much harder to tell how <i>many</i> of each previous name there are - attending 0.3 to each copy of John will look exactly the same as attending 0.6 to a single John token. So this will be pretty hard to figure out on the " to" token!
# A good starting point is thinking though whether a tiny model could do this, eg a <a href="https://transformer-circuits.pub/2021/framework/index.html">1L Attn-Only model</a>. I'm pretty sure the answer is no! Attention is really good at the primitive operations of looking nearby, or copying information. I can believe a tiny model could figure out that at `to` it should look for names and predict that those names came next (eg the skip trigram " John...to -> John"). But it's much harder to tell how <i>many</i> of each previous name there are - attending 0.3 to each copy of John will look exactly the same as attending 0.6 to a single John token. So this will be pretty hard to figure out on the " to" token!
#
#
# The natural place to break this symmetry is on the second " John" token - telling whether there is an earlier copy of the <i>current</i> token should be a much easier task. So I might expect there to be a head which detects duplicate tokens on the second " John" token, and then another head which moves that information from the second " John" token to the " to" token.
# The natural place to break this symmetry is on the second " John" token - telling whether there is an earlier copy of the <i>current</i> token should be a much easier task. So I might expect there to be a head which detects duplicate tokens on the second " John" token, and then another head which moves that information from the second " John" token to the " to" token.
#
#
# The model then needs to learn to predict " Mary" and <i>not</i> " John". I can see two natural ways to do this:
# The model then needs to learn to predict " Mary" and <i>not</i> " John". I can see two natural ways to do this:
# 1. Detect all preceding names and move this information to " to" and then delete the any name corresponding to the duplicate token feature. This feels easier done with a non-linearity, since precisely cancelling out vectors is hard, so I'd imagine an MLP layer deletes the " John" direction of the residual stream
# 1. Detect all preceding names and move this information to " to" and then delete the any name corresponding to the duplicate token feature. This feels easier done with a non-linearity, since precisely cancelling out vectors is hard, so I'd imagine an MLP layer deletes the " John" direction of the residual stream
# 2. Have a head which attends to all previous names, but where the duplicate token features <i>inhibit</i> it from attending to specific names. So this only attends to Mary. And then the output of this head maps to the logits.  
# 2. Have a head which attends to all previous names, but where the duplicate token features <i>inhibit</i> it from attending to specific names. So this only attends to Mary. And then the output of this head maps to the logits.  
#
#
# (Spoiler: It's the second one).
# (Spoiler: It's the second one).
#
#
# <h3>Experiment Ideas</h3>
# <h3>Experiment Ideas</h3>
#
#
# A test that could distinguish these two is to look at which components of the model add directly to the logits - if it's mostly attention heads which attend to " Mary" and to neither " John" it's probably hypothesis 2, if it's mostly MLPs it's probably hypothesis 1.
# A test that could distinguish these two is to look at which components of the model add directly to the logits - if it's mostly attention heads which attend to " Mary" and to neither " John" it's probably hypothesis 2, if it's mostly MLPs it's probably hypothesis 1.
#
#
# And we should be able to identify duplicate token heads by finding ones which attend from " John" to " John", and whose outputs are then moved to the " to" token by V-Composition with another head (Spoiler: It's more complicated than that!)
# And we should be able to identify duplicate token heads by finding ones which attend from " John" to " John", and whose outputs are then moved to the " to" token by V-Composition with another head (Spoiler: It's more complicated than that!)
#
#
# Note that all of the above reasoning is very simplistic and could easily break in a real model! There'll be significant parts of the model that figure out whether to use this circuit at all (we don't want to inhibit duplicated names when, eg, figuring out what goes at the start of the <i>next</i> sentence), and may be parts towards the end of the model that do "post-processing" just before the final output. But it's a good starting point for thinking about what's going on.
# Note that all of the above reasoning is very simplistic and could easily break in a real model! There'll be significant parts of the model that figure out whether to use this circuit at all (we don't want to inhibit duplicated names when, eg, figuring out what goes at the start of the <i>next</i> sentence), and may be parts towards the end of the model that do "post-processing" just before the final output. But it's a good starting point for thinking about what's going on.
# %% [markdown]
# %% [markdown]
# ## Direct Logit Attribution
# ## Direct Logit Attribution
# %% [markdown]
# %% [markdown]
# *Look up unfamiliar terms in the [mech interp explainer](https://neelnanda.io/glossary)*
# *Look up unfamiliar terms in the [mech interp explainer](https://neelnanda.io/glossary)*
#
#
# Further, the easiest part of the model to understand is the output - this is what the model is trained to optimize, and so it can always be directly interpreted! Often the right approach to reverse engineering a circuit is to start at the end, understand how the model produces the right answer, and to then work backwards. The main technique used to do this is called **direct logit attribution**
# Further, the easiest part of the model to understand is the output - this is what the model is trained to optimize, and so it can always be directly interpreted! Often the right approach to reverse engineering a circuit is to start at the end, understand how the model produces the right answer, and to then work backwards. The main technique used to do this is called **direct logit attribution**
#
#
# **Background:** The central object of a transformer is the **residual stream**. This is the sum of the outputs of each layer and of the original token and positional embedding. Importantly, this means that any linear function of the residual stream can be perfectly decomposed into the contribution of each layer of the transformer. Further, each attention layer's output can be broken down into the sum of the output of each head (See [A Mathematical Framework for Transformer Circuits](https://transformer-circuits.pub/2021/framework/index.html) for details), and each MLP layer's output can be broken down into the sum of the output of each neuron (and a bias term for each layer).
# **Background:** The central object of a transformer is the **residual stream**. This is the sum of the outputs of each layer and of the original token and positional embedding. Importantly, this means that any linear function of the residual stream can be perfectly decomposed into the contribution of each layer of the transformer. Further, each attention layer's output can be broken down into the sum of the output of each head (See [A Mathematical Framework for Transformer Circuits](https://transformer-circuits.pub/2021/framework/index.html) for details), and each MLP layer's output can be broken down into the sum of the output of each neuron (and a bias term for each layer).
#
#
# The logits of a model are `logits=Unembed(LayerNorm(final_residual_stream))`. The Unembed is a linear map, and LayerNorm is approximately a linear map, so we can decompose the logits into the sum of the contributions of each component, and look at which components contribute the most to the logit of the correct token! This is called **direct logit attribution**. Here we look at the direct attribution to the logit difference!
# The logits of a model are `logits=Unembed(LayerNorm(final_residual_stream))`. The Unembed is a linear map, and LayerNorm is approximately a linear map, so we can decompose the logits into the sum of the contributions of each component, and look at which components contribute the most to the logit of the correct token! This is called **direct logit attribution**. Here we look at the direct attribution to the logit difference!
#
#
# <details> <summary>(*) <b>Background and motivation of the logit difference</b></summary>
# <details> <summary>(*) <b>Background and motivation of the logit difference</b></summary>
#
#
# Logit difference is actually a *really* nice and elegant metric and is a particularly nice aspect of the setup of Indirect Object Identification. In general, there are two natural ways to interpret the model's outputs: the output logits, or the output log probabilities (or probabilities).
# Logit difference is actually a *really* nice and elegant metric and is a particularly nice aspect of the setup of Indirect Object Identification. In general, there are two natural ways to interpret the model's outputs: the output logits, or the output log probabilities (or probabilities).
#
#
# The logits are much nicer and easier to understand, as noted above. However, the model is trained to optimize the cross-entropy loss (the average of log probability of the correct token). This means it does not directly optimize the logits, and indeed if the model adds an arbitrary constant to every logit, the log probabilities are unchanged.
# The logits are much nicer and easier to understand, as noted above. However, the model is trained to optimize the cross-entropy loss (the average of log probability of the correct token). This means it does not directly optimize the logits, and indeed if the model adds an arbitrary constant to every logit, the log probabilities are unchanged.
#
#
# But `log_probs == logits.log_softmax(dim=-1) == logits - logsumexp(logits)`, and so `log_probs(" Mary") - log_probs(" John") = logits(" Mary") - logits(" John")` - the ability to add an arbitrary constant cancels out!
# But `log_probs == logits.log_softmax(dim=-1) == logits - logsumexp(logits)`, and so `log_probs(" Mary") - log_probs(" John") = logits(" Mary") - logits(" John")` - the ability to add an arbitrary constant cancels out!
#
#
# Further, the metric helps us isolate the precise capability we care about - figuring out *which* name is the Indirect Object. There are many other components of the task - deciding whether to return an article (the) or pronoun (her) or name, realising that the sentence wants a person next at all, etc. By taking the logit difference we control for all of that.
# Further, the metric helps us isolate the precise capability we care about - figuring out *which* name is the Indirect Object. There are many other components of the task - deciding whether to return an article (the) or pronoun (her) or name, realising that the sentence wants a person next at all, etc. By taking the logit difference we control for all of that.
#
#
# Our metric is further refined, because each prompt is repeated twice, for each possible indirect object. This controls for irrelevant behaviour such as the model learning that John is a more frequent token than Mary (this actually happens! The final layernorm bias increases the John logit by 1 relative to the Mary logit)
# Our metric is further refined, because each prompt is repeated twice, for each possible indirect object. This controls for irrelevant behaviour such as the model learning that John is a more frequent token than Mary (this actually happens! The final layernorm bias increases the John logit by 1 relative to the Mary logit)
#
#
# </details>
# </details>
#
#
# <details> <summary>Ignoring LayerNorm</summary>
# <details> <summary>Ignoring LayerNorm</summary>
#
#
# LayerNorm is an analogous normalization technique to BatchNorm (that's friendlier to massive parallelization) that transformers use. Every time a transformer layer reads information from the residual stream, it applies a LayerNorm to normalize the vector at each position (translating to set the mean to 0 and scaling to set the variance to 1) and then applying a learned vector of weights and biases to scale and translate the normalized vector. This is *almost* a linear map, apart from the scaling step, because that divides by the norm of the vector and the norm is not a linear function. (The `fold_ln` flag when loading a model factors out all the linear parts).
# LayerNorm is an analogous normalization technique to BatchNorm (that's friendlier to massive parallelization) that transformers use. Every time a transformer layer reads information from the residual stream, it applies a LayerNorm to normalize the vector at each position (translating to set the mean to 0 and scaling to set the variance to 1) and then applying a learned vector of weights and biases to scale and translate the normalized vector. This is *almost* a linear map, apart from the scaling step, because that divides by the norm of the vector and the norm is not a linear function. (The `fold_ln` flag when loading a model factors out all the linear parts).
#
#
# But if we fixed the scale factor, the LayerNorm would be fully linear. And the scale of the residual stream is a global property that's a function of *all* components of the stream, while in practice there is normally just a few directions relevant to any particular component, so in practice this is an acceptable approximation. So when doing direct logit attribution we use the `apply_ln` flag on the `cache` to apply the global layernorm scaling factor to each constant. See [my clean GPT-2 implementation](https://colab.research.google.com/github/TransformerLensOrg/TransformerLens/blob/cle
# But if we fixed the scale factor, the LayerNorm would be fully linear. And the scale of the residual stream is a global property that's a function of *all* components of the stream, while in practice there is normally just a few directions relevant to any particular component, so in practice this is an acceptable approximation. So when doing direct logit attribution we use the `apply_ln` flag on the `cache` to apply the global layernorm scaling factor to each constant. See [my clean GPT-2 implementation](https://colab.research.google.com/github/TransformerLensOrg/TransformerLens/blob/clean-transformer-demo/Clean_Transformer_Demo.ipynb#scrollTo=Clean_Transformer_Implementation) for more on LayerNorm.
# </details>
# %% [markdown]
# Getting an output logit is equivalent to projecting onto a direction in the residual stream. We use `model.tokens_to_residual_directions` to map the answer tokens to that direction, and then convert this to a logit difference direction for each batch
# %%
answer_residual_directions = model.tokens_to_residual_directions(answer_tokens)
print("Answer residual directions shape:", answer_residual_directions.shape)
logit_diff_directions = (
    answer_residual_directions[:, 0] - answer_residual_directions[:, 1]
)
print("Logit difference directions shape:", logit_diff_directions.shape)
# %% [markdown]
# To verify that this works, we can apply this to the final residual stream for our cached prompts (after applying LayerNorm scaling) and verify that we get the same answer.
#
# <details> <summary>Technical details</summary>
#
# `logits = Unembed(LayerNorm(final_residual_stream))`, so we technically need to account for the centering, and then learned translation and scaling of the layernorm, not just the variance 1 scaling.
#
# The centering is accounted for with the preprocessing flag `center_writing_weights` which ensures that every weight matrix writing to the residual stream has mean zero.
#
# The learned scaling is folded into the unembedding weights `model.unembed.W_U` via `W_U_fold = layer_norm.weights[:, None] * unembed.W_U`
#
# The learned translation is folded to `model.unembed.b_U`, a bias added to the logits (note that GPT-2 is not