moground-saes / load_sae.py
LuckerZ's picture
Vision-language Top-K SAEs for Qwen2.5-VL-3B and LLaVA-NeXT-8B
b8c93e9 verified
Raw History Blame Contribute Delete
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)