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:
Reconstruction (FVU) and sparsity (L0) before/after the fine-tune, and the sparsity-penalty comparison (reconstruction-only vs +L1):
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-layeraverage_l0_76-nearest selection); base modelgoogle/gemma-2-2b; instruct-input sourcegoogle/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.mdin the repo.
- fine-tune run + weights +
- 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.
Model tree for siddharthmb/2026.TA.gemma2_2b_gemmascope_transcoders_instruct_ft_L0-24-25
Base model
google/gemma-2-2b


