Skip to content
AI-grafen
GFrontier LabLab· about 120 min· server sandbox

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

  1. 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.
  2. 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).
  3. 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 account

Expected 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.