Gemma4B_headfinal_sae
Final sparse autoencoder checkpoints for Gemma-3-4B-PT layers 0, 3, 7, 15, 23, 27, 31, 33. Each SAE was trained on 200 million retained tokens per language (Korean, Japanese, Hindi), for 600 million tokens per layer. Intermediate checkpoints are not included.
Each layer_N/ directory preserves the final checkpoint layout:
layers.N/sae.safetensors: SAE weights.layers.N/cfg.json: SAE architecture configuration.config.json: training configuration and final token counts.training_state.pt: optimizer, scheduler, token progress and dead-feature state.validation/metrics.json: final held-out reconstruction metrics.
The SAE input is the zero-indexed transformer block output, before the final model norm. Hidden size is 2,560; expansion factor is 32; width is 81,920; ReLU TopK is 50. The frozen model uses bf16 and the SAE uses fp32, with no input normalization. The global batch is 24 sequences of 1,024 tokens; AuxK coefficient is 1/32. Training excludes position 0, EOS, and layer-specific norm outliers, with independent language quotas.
All eight layers were trained independently, one GPU per layer.
The base model revision is cc012e0a6d0787b4adcc0fa2c4da74402494554d.
Training text came from HuggingFaceFW/fineweb-2; raw training text is not included here.
catalog.json lists all eight SAEs and their paths. manifest.json records checkpoint file sizes
and SHA256 hashes. The repository is public with manually approved access to model files.
After access approval, authenticate with a read token to download a layer:
from huggingface_hub import snapshot_download
path = snapshot_download(
repo_id="mjkmain/Gemma4B_headfinal_sae",
allow_patterns=["layer_23/**", "catalog.json", "manifest.json"],
token=True,
)
These are sparse autoencoders, not standalone language models. The safetensors file is the
inference artifact; training_state.pt is a PyTorch training-state checkpoint.
Model tree for mjkmain/Gemma4B_headfinal_sae
Base model
google/gemma-3-4b-pt