Download load_sae.py from LuckerZ/moground-saes: direct link, hf CLI and curl.
- Browser
- Download file 1.38 kB
-
https://huggingface.co/LuckerZ/moground-saes/resolve/main/load_sae.py
- Command line
-
hf download hf://LuckerZ/moground-saes/load_sae.py
-
curl -L -o load_sae.py https://huggingface.co/LuckerZ/moground-saes/resolve/main/load_sae.py
1.38 kB
| """Load a released MoGround SAE checkpoint. | |
| from load_sae import load | |
| sae = load("qwen/layer28/sae.pt") | |
| out = sae(residual) # {"x_hat", "z", "indices", "pre"} | |
| """ | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| class TopKSAE(nn.Module): | |
| """z = ReLU(TopK(W_enc (x - b_dec) + b_enc)), x_hat = W_dec z + b_dec.""" | |
| def __init__(self, d_model: int, d_sae: int, k: int): | |
| super().__init__() | |
| self.k = k | |
| self.W_dec = nn.Parameter(torch.zeros(d_sae, d_model)) | |
| self.W_enc = nn.Parameter(torch.zeros(d_model, d_sae)) | |
| self.b_enc = nn.Parameter(torch.zeros(d_sae)) | |
| self.b_dec = nn.Parameter(torch.zeros(d_model)) | |
| def forward(self, x): | |
| pre = (x - self.b_dec) @ self.W_enc + self.b_enc | |
| vals, idx = pre.topk(self.k, dim=-1) | |
| z = torch.zeros_like(pre).scatter_(-1, idx, F.relu(vals)) | |
| return {"x_hat": z @ self.W_dec + self.b_dec, "z": z, "indices": idx, "pre": pre} | |
| def load(path, device="cpu"): | |
| sd = torch.load(path, map_location=device) | |
| sd = sd.get("state_dict", sd) | |
| d_model, d_sae = sd["W_enc"].shape | |
| sae = TopKSAE(d_model, d_sae, k=int(torch.load(path, map_location="cpu").get("k", 32))) | |
| sae.load_state_dict(sd, strict=True) # strict: a renamed key must fail loudly, not silently | |
| return sae.to(device).eval().requires_grad_(False) | |