BatchTopK SAEs for Qwen3.6-35B-A3B
BatchTopK sparse autoencoders (Bussmann et al.) trained on the residual-stream
activations of Qwen/Qwen3.6-35B-A3B
at layers [10, 20, 30], with automatically generated feature labels.
Specs
- Sites: residual stream, layers [10, 20, 30] (~0.25 / 0.5 / 0.75 depth)
- Dictionary:
d_sae = 65536(32x expansion ofd_in = 2048), sparsityk = 32 - Training:
250M tokens of FineWeb-Edu + LMSYS-Chat-1M, bf16 activations, AuxK dead-feature revival, then a low-LR anneal (LR→0) that added **+0.02 explained variance** - Inference: per-token thresholding via a stored EMA
threshold
Files (per layer)
layer_<L>/sae.pt— checkpoint:state_dict,cfg,norm_scale,steplayer_<L>/labels.json— concept labels for every feature (full dictionary)sae_core.py— theBatchTopKSAEmodule to load them
Feature labels & eval
Labels were generated by Qwen2.5-32B-Instruct from activation-magnitude-ranked examples (each shown with a 0–10 strength), for every (non-dead) feature in the dictionary.
Label quality was measured on a held-out ~900 features/layer sample (labeling the full dictionary per-feature would be prohibitively expensive); those means (with 95% CIs) are reported below and carry over to the full set, which uses the identical method:
detection_bal_acc— balanced accuracy of a detection quiz (does the label predict whether the feature fires; 0.5 = chance).rank_corr— Spearman correlation between predicted-strength ranks and true activation-magnitude ranks (does the label predict how strongly).
| layer | detection bal-acc | rank corr | eval features | labeled features |
|---|---|---|---|---|
| 10 | 0.761 ± 0.013 | 0.51 ± 0.021 | 858 | 65015 |
| 20 | 0.766 ± 0.013 | 0.476 ± 0.023 | 902 | 64940 |
| 30 | 0.75 ± 0.013 | 0.473 ± 0.022 | 842 | 63202 |
Detection (0.76) is well above chance and roughly constant across depth; ranking
(0.49) is weaker and mildly higher in early layers — the labels capture presence
better than intensity. Metrics are held-out means ± 95% CI over the eval sample (method is identical for the full set). detection bal-acc is flat across depth (0.76); the labeled count is the number of features labeled per layer.
Design notes
k = 32was chosen for interpretability. Ablating tok=64raised reconstruction but lowered label quality (detection −0.05, rank −0.15) — denser codes are less monosemantic.- A from-scratch JumpReLU baseline did not beat these BatchTopK SAEs at matched sparsity (with comparable training budget).
Usage
import torch
from sae_core import BatchTopKSAE, SAEConfig
ckpt = torch.load("layer_20/sae.pt", map_location="cpu")
sae = BatchTopKSAE(SAEConfig(**ckpt["cfg"]))
sae.load_state_dict(ckpt["state_dict"])
sae.eval()
# resid: [..., 2048] residual-stream activations from the same layer
feats = sae.encode(resid * ckpt["norm_scale"], use_threshold=True) # sparse [..., 65536]
recon = sae.decode(feats)
Model tree for stanleytheli/qwen3.6-35b-a3b-saes
Base model
Qwen/Qwen3.6-35B-A3B