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/.

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 dmanningcoe/fra-tinystories-fig2-saes

Finetuned
(1)
this model