TinyStories sleeper Figure 2 SAEs
This repository contains 72 newly trained TopK sparse autoencoders (SAEs):
the attention input (blocks.L.ln1.hook_normalized), residual midpoint
(blocks.L.hook_resid_mid), and residual output
(blocks.L.hook_resid_post) at each of layers 0โ3, with training seeds 0โ5.
They were trained on activations from
mars-jason-25/tiny-stories-33M-TSdata-sleeper
using the balanced clean/deployment examples in
mars-jason-25/tiny_stories_instruct_sleeper_data.
Training
The recipe follows the paper appendix: 10,000 fixed 128-token training
sequences, 200 validation sequences, input width 768, SAE width 1,536,
TopK 32 after ReLU, MSE reconstruction loss, Adam with learning rate 5e-4,
batch size 4,096, 4,000 steps, and unit-norm decoder rows renormalized
every 100 steps. Each hookpoint/seed has independent SAE parameters and
optimizer state. The implementation batches the SAEs on a single GPU for
throughput. See train.py, sae_models.py, and sleeper_utils.py for the
complete data and training code.
The files are weights/sae_L{layer}_{ln1|resid_mid|resid_post}_s{seed}.pt.
Each is a PyTorch dictionary with state_dict, config, and val_fvu.
The state dictionary has W_enc, b_enc, W_dec, and b_dec and loads into
TopKSAE(d_in=768, d_sae=1536, k=32) from sae_models.py. For example:
import torch
from huggingface_hub import hf_hub_download
from sae_models import TopKSAE
path = hf_hub_download(
"dmanningcoe/fra-tinystories-fig2-saes",
"weights/sae_L0_ln1_s0.pt",
)
payload = torch.load(path, map_location="cpu", weights_only=True)
sae = TopKSAE(d_in=768, d_sae=1536, k=32)
sae.load_state_dict(payload["state_dict"])
harvest.json records the activation data sources and tensor shapes;
complete.json records the completed cells and steps; quality.json
contains held-out FVU for every checkpoint. Across the 72 checkpoints,
validation FVU ranges from 0.0251 to 0.1942. These are SAE activation
reconstruction scores, not attention-component FVU measurements.
Run python verify.py after downloading the whole repository to validate
all 72 checkpoints. The path entries in quality.json record their
original training-host locations; the files here are under weights/.