Lab: find the circuit — activation patching and ablation in a small transformer
Implement activation patching and mean ablation for attention heads in a small attention-only transformer (NumPy, fixed weights), and use them to find the circuit of heads that performs induction ([A][B] … [A] → [B]), verify that the circuit is complete and report.
Theory
Patching: run the model on a clean and a corrupted prompt, paste a component's activation from the clean run into the corrupted one and measure how much of the logit difference is restored (0 = nothing, 1 = everything). Mean ablation: replace the component's activation with its mean over reference prompts. A circuit is the smallest set of components that suffices: verify by ablating everything except the circuit. The model here has four heads: two perform induction (a previous-token head in layer 0, an induction head in layer 1), two are distractors that write to unused dimensions.
Sub-tasks
- logit diff and patching —
logit_diff(logits, correct, wrong)at the last position.patch_head(model, clean, corrupt, layer, head, correct, wrong)→ restored fraction (ld_patch − ld_corrupt)/(ld_clean − ld_corrupt), where the head's z is replaced by z from the clean run. - mean ablation —
mean_ablate(model, tokens, heads, refs, correct, wrong)→ logit diff when every head in heads has its z replaced by the mean z (position-wise) over refs (a list of token lists of the same length). - patching map and circuit —
patch_map(model, clean, corrupt, correct, wrong)→ {(layer, head): restored fraction};find_circuit(pmap, threshold=0.5)→ sorted list of heads ≥ threshold;verify_circuit(model, clean, circuit, refs, correct, wrong)→ fraction of the clean logit diff that remains when all heads outside the circuit are mean-ablated.
Passes when: ok >= 1
The starter code
runs in an isolated sandbox on the server"""Kretsanalys. Fyll i funktionerna. Modellen: tinymodel.TinyModel med forward(tokens, hooks, cache)."""
import numpy as np
def logit_diff(logits, correct, wrong):
"""logit[correct] − logit[wrong] på sista positionen."""
# TODO
raise NotImplementedError
def patch_head(model, clean, corrupt, layer, head, correct, wrong):
"""Kör corrupt med huvudets z ersatt av z från clean-körningen. Returnera återställd andel."""
# TODO
raise NotImplementedError
def mean_ablate(model, tokens, heads, refs, correct, wrong):
"""Ersätt z för varje huvud i heads med positionsvis medel över refs. Returnera logit diff."""
# TODO
raise NotImplementedError
def patch_map(model, clean, corrupt, correct, wrong):
"""{(layer, head): återställd andel} för alla huvuden."""
# TODO
raise NotImplementedError
def find_circuit(pmap, threshold=0.5):
"""Sorterad lista av (layer, head) med återställd andel ≥ threshold."""
# TODO
raise NotImplementedError
def verify_circuit(model, clean, circuit, refs, correct, wrong):
"""Andel av ren logit diff som återstår när alla huvuden UTANFÖR kretsen medelableras."""
# TODO
raise NotImplementedError
You write the code; tests you cannot see decide whether it holds up. Create a free account to run the lab.
Try the diagnosticCreate a free accountExpected results
The patching map gives ≈ 1.0 for (0, 0) and (1, 0) and ≈ 0 for (0, 1) and (1, 1). find_circuit → [(0, 0), (1, 0)]. verify_circuit ≥ 0.9 — the distractors can be ablated without effect. Eval: ok = 1.
Common mistakes
- Patches the residual stream instead of the head's z — then all heads get mixed together.
- Computes the restored fraction with the wrong sign or forgets to normalise by (ld_clean − ld_corrupt).
- Mean ablation: the mean should be taken position-wise over references of the same length, not over all positions.
- verify_circuit: ablate the heads outside the circuit, not the circuit.