Views
No views yet
| d_in | 576 |
| d_sae | 4,608 |
| expansion | 8x |
| k | 16 |
| tied weights | no |
| hook layer | 48 of 80 |
| activation | TopK + ReLU |
| dtype | float32 |
x @ W_enc + b_enc -> keep top 16 values (ReLU'd) -> f @ W_dec + b_dec.| File | Size | Description |
|---|---|---|
sae_weights.pt | 20 MB | SAE weights + config (no optimizer state) |
feature_labels.jsonl | 651 KB | 4,602 autointerp labels with confidence tags |
feature_data.json | 628 KB | UMAP coords + activation stats per feature |
train_summary.json | 1 KB | Training config and final metrics |
1import torch
2import torch.nn as nn
3import torch.nn.functional as F
4
5class TopKSAE(nn.Module):
6 def __init__(self, d_in, d_sae, k):
7 super().__init__()
8 self.d_in, self.d_sae, self.k = d_in, d_sae, k
9 self.W_enc = nn.Parameter(torch.empty(d_in, d_sae))
10 self.W_dec = nn.Parameter(torch.empty(d_sae, d_in))
11 self.b_enc = nn.Parameter(torch.zeros(d_sae))
12 self.b_dec = nn.Parameter(torch.zeros(d_in))
13 self.register_buffer("feature_activations", torch.zeros(d_sae, dtype=torch.long))
14 self.register_buffer("steps_since_active", torch.zeros(d_sae, dtype=torch.long))
15
16 def forward(self, x):
17 pre_acts = x @ self.W_enc + self.b_enc
18 topk_vals, topk_idx = torch.topk(pre_acts, self.k, dim=-1)
19 topk_vals = F.relu(topk_vals)
20 f = torch.zeros_like(pre_acts)
21 f.scatter_(-1, topk_idx, topk_vals)
22 x_hat = f @ self.W_dec + self.b_dec
23 return x_hat, f, topk_idx, topk_vals
24
25# load
26ckpt = torch.load("sae_weights.pt", map_location="cpu", weights_only=False)
27cfg = ckpt["config"]
28sae = TopKSAE(cfg["d_in"], cfg["d_sae"], cfg["k"])
29sae.load_state_dict(ckpt["model_state_dict"])
30sae.eval().cuda()1from transformers import AutoModelForCausalLM, AutoTokenizer
2
3model = AutoModelForCausalLM.from_pretrained(
4 "PleIAs/Baguettotron", torch_dtype=torch.bfloat16,
5 device_map="cuda", trust_remote_code=True,
6)
7tokenizer = AutoTokenizer.from_pretrained("PleIAs/Baguettotron")
8
9# hook layer 48 to grab residual stream
10activations = {}
11def hook_fn(module, input, output):
12 # output is a tuple; first element is the hidden state
13 activations["resid"] = output[0].detach().float()
14
15handle = model.model.layers[48].register_forward_hook(hook_fn)
16
17inputs = tokenizer("The cat sat on the", return_tensors="pt").to("cuda")
18with torch.no_grad():
19 model(**inputs)
20
21x = activations["resid"] # (batch, seq, 576)
22x_hat, f, topk_idx, topk_vals = sae(x)
23
24handle.remove()feature_labels.jsonl has one JSON object per line:{"feature": 0, "interp": "", "autointerp": "military command and leadership", "confidence": "confident"}confident, tentative, or dead.feature_data.json contains per-feature metadata for building explorers or doing analysis:1{
2 "features": [
3 {"i": 0, "x": 3.08, "y": 1.94, "x3": 3.20, "y3": 1.55, "z3": 1.68,
4 "d": 0.000712, "mx": 296.19, "mn": 74.80, "c": 1178}
5 ],
6 "meta": { ... }
7}i = feature index, x/y = 2D UMAP coords, x3/y3/z3 = 3D UMAP coords, d = density, mx = max activation, mn = mean activation, c = fire count.