GemmaScope transcoders re-fine-tuned to instruction-tuned hidden states (gemma-2-2b, layers 0/24/25)

The setup. GemmaScope transcoders are sparse, interpretable modules that approximate each MLP layer of google/gemma-2-2b: for a layer, transcoder(x) β‰ˆ MLP_base(x) where x is the post-ln2 hidden state, and only a small number of features fire (low L0 = active features/token). They are the building block circuit-tracer uses to turn a forward pass into an interpretable feature graph. Crucially, they were trained to imitate the base MLP on the base model's hidden states.

The problem this fixes. When you circuit-trace the base↔instruct gap, you feed these transcoders the instruction-tuned model's (google/gemma-2-2b-it) hidden states instead β€” a distribution shift. We measured that this degrades reconstruction: on instruct hidden states the FVU (fraction of variance unexplained; 0 = perfect, lower is better) is higher than on base hidden states at every layer, worst at the endpoints β€” layer 0 +33%, layer 24 +32%, layer 25 +104%. In a circuit graph that lost fidelity shows up as larger "error nodes" (computation the graph cannot attribute to interpretable features).

What this artifact is. The three worst-hit transcoders β€” layers 0, 24, 25 (width_16k, from the average_l0_76-nearest per-layer selection) β€” re-fine-tuned to keep imitating the base MLP but accurately on the instruct input distribution. Each file is the full transcoder state (W_enc, W_dec, b_enc, b_dec; the JumpReLU threshold is frozen and unchanged), a drop-in replacement for the corresponding pretrained GemmaScope layer.

  • finetuned_layer_0.safetensors, finetuned_layer_24.safetensors, finetuned_layer_25.safetensors

How it was trained. Supervised, per layer, starting from the pretrained GemmaScope weights, minimizing MSE(transcoder(x), MLP_base(x)) with x = the instruct model's blocks.L.ln2.hook_normalized, over 2M instruct-distribution tokens (14 min, lr 1e-4, JumpReLU threshold frozen). A per-layer decoder-norm-weighted L1 sparsity penalty (l1_coeff = [0, 1e-3, 1e-3] for layers [0, 24, 25]) holds L0 at the base operating point β€” layer 0 was already at base L0 so it gets no penalty; layers 24/25 overshoot without one.

The result β€” feature detectors are unchanged; reconstruction and sparsity recover. The fine-tune moved the decoder/biases, not what each feature detects: per-feature cos(W_enc_orig, W_enc_ft) = 0.9997, with zero features below 0.99 β€” so the feature identities (and their Neuronpedia/collected activation examples) are intact. On 200k held-out instruct tokens, reconstruction improves and sparsity lands on the base target:

layer FVU before β†’ after base-input FVU (target) L0 before β†’ after (no-penalty β†’ +L1) base L0
0 0.108 β†’ 0.075 (βˆ’31%) 0.082 81 β†’ 81 80
24 0.281 β†’ 0.228 (βˆ’19%) 0.21 67 β†’ 56.6 56
25 0.317 β†’ 0.221 (βˆ’30%) 0.15 81 β†’ 68.2 65

The reconstruction gap the shift opened is mostly closed, and the sparsity overshoot (L24 67, L25 81) is removed β€” L0 sits on the base operating point at a ~3% FVU cost. In circuit-tracer attribution graphs, this lowers the error-node fraction ~5% on every tested prompt (0.258β†’0.239, 0.275β†’0.265, 0.273β†’0.260) β€” more of the graph carried by interpretable features.

One-sentence version. Three GemmaScope transcoders re-tuned to still reconstruct the base gemma-2-2b MLP sparsely, but on instruction-tuned hidden states β€” closing most of the reconstruction gap the distribution shift opened while leaving the features (and their sparsity) intact.

Figures

Base-vs-instruct shift (motivation), 26-layer profile:

input-shift overview

Reconstruction (FVU) and sparsity (L0) before/after the fine-tune, and the sparsity-penalty comparison (reconstruction-only vs +L1):

FVU before/after L0 before/after sparsity-penalty comparison

How to use

These patch onto the pretrained GemmaScope set at load time (layers 0/24/25 replaced). With the transcoder-adapters repo, pass the flag to any GemmaScope-loading tool β€” it accepts this HF repo id directly:

# Circuit-tracer visualizer (production base attribution / combined overlay + serve):
uv run --extra viz python -m analysis.attribution.run_base_adapter_comparison \
  --base_model google/gemma-2-2b --gemmascope_width width_16k --gemmascope_l0 average_l0_76 \
  --finetuned_transcoder_dir siddharthmb/2026.TA.gemma2_2b_gemmascope_transcoders_instruct_ft_L0-24-25 \
  --finetuned_layers 0 24 25  ... --serve

# FVU/L0 evaluation:
uv run --extra viz python -m analysis.features.transcoder_input_shift \
  --gemmascope_width width_16k --gemmascope_l0 average_l0_76 --sources base instruct --layers 0 24 25 \
  --finetuned_transcoder_dir siddharthmb/2026.TA.gemma2_2b_gemmascope_transcoders_instruct_ft_L0-24-25 --finetuned_layers 0 24 25

Or apply directly: load the GemmaScope TranscoderSet, then for each layer copy finetuned_layer_{L}.safetensors's W_enc/W_dec/b_enc/b_dec in place.

Reproduce

# 1) The fine-tune that produced these weights (job 16123724):
uv run --extra viz python -m analysis.features.finetune_transcoder_shift \
  --gemmascope_width width_16k --gemmascope_l0 average_l0_76 \
  --layers 0 24 25 --l1_coeff 0 1e-3 1e-3 \
  --train_tokens 2000000 --eval_tokens 200000 --lr 1e-4 --wandb
# 2) The base-vs-instruct shift measurement (Experiment 1, job 16113808):
uv run --extra viz python -m analysis.features.transcoder_input_shift \
  --gemmascope_width width_16k --gemmascope_l0 average_l0_76 --sources base instruct --layers all --max_tokens 500000 --wandb

Provenance

  • Fine-tuned from: google/gemma-scope-2b-pt-transcoders (width_16k, per-layer average_l0_76-nearest selection); base model google/gemma-2-2b; instruct-input source google/gemma-2-2b-it.
  • Training data: chat = siddharthmb/2026.transcoder-adapters.lmsys-chat-1m-splits (rendered with the -it chat template); the shift was also verified on web = science-of-finetuning/fineweb-1m-sample.
  • Weights & Biases: final fine-tune 9o1ebt6v; project siddharth-stanford/transcoder-feature-collection.
  • Cluster artifacts (NLP cluster, $LARGE_ARTIFACTS_DIR=/nlp/scr/siddharth):
    • fine-tune run + weights + finetune_report.json: .../transcoder-adapters/transcoder_input_shift_finetune/ft_gemmascope_width_16k_average_l0_76_L0-24-25_20260710_025814_16123724/
    • encoder-drift check: encoder_drift.json (in that dir, also included here)
    • Experiment-1 (shift) runs: .../transcoder-adapters/transcoder_input_shift/... (jobs 16113538 chat, 16113604 web, 16113808 all-layers)
    • graph comparison: .../transcoder-adapters/transcoder_finetune_graphs/L0-24-25_20260710_171004_16129263/
    • Slurm logs: logs/finetune_transcoder_shift/*16123724*, logs/transcoder_input_shift/, logs/compare_finetuned_graphs/*16129263*
    • full writeup: my_notes/07-09-26/README.md in the repo.
  • Also included here: finetune_report.json (before/after eval), encoder_drift.json, summary.md, gemmascope_config.json (the exact per-layer L0 selection).

Note on L0 selection. These use the average_l0_76-nearest per-layer selection (this project's convention), which differs from the canonical mntss/gemma-scope-transcoders selection at 18/26 layers β€” so Neuronpedia feature descriptions keyed to mntss align only where the L0 coincides; the collected activation examples (same selection) align everywhere.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support

Model tree for siddharthmb/2026.TA.gemma2_2b_gemmascope_transcoders_instruct_ft_L0-24-25

Finetuned
(564)
this model

Datasets used to train siddharthmb/2026.TA.gemma2_2b_gemmascope_transcoders_instruct_ft_L0-24-25