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 of d_in = 2048), sparsity k = 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, step
  • layer_<L>/labels.json — concept labels for every feature (full dictionary)
  • sae_core.py — the BatchTopKSAE module 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 = 32 was chosen for interpretability. Ablating to k=64 raised 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)
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 stanleytheli/qwen3.6-35b-a3b-saes

Finetuned
(372)
this model