anubhavg97/neuron-scalpel-sae-safety-qwen
0
Neuron Scalpel: Safety SAEs for Qwen 2.5 1.5B-Instruct
Sparse Autoencoders trained on safety-relevant activations from Qwen 2.5 1.5B-Instruct, targeting layers 10 and 20 with 8x and 16x expansion.
Key Finding
SAE features reliably detect safety concepts (Cohen's d = 0.727) but have zero causal control over refusal behavior. This holds for:
- Full-replacement steering (reconstruction destroys refusal: 100% → 44%)
- Residual steering (contrastive features = random features at all alpha values)
- All tested layers (10, 14, 20), widths (x8, x16), and feature selection methods
Models
Training Data
- 1.7M activations merged from:
- BeaverTails + hh-rlhf generation-time activations (1.2M, teacher-forcing)
- WildJailbreak contrastive pairs (500K, 261K vanilla_harmful + 2.2K eval)
- Float16, memmap-backed collection
Layer Selection
Layer boundary analysis (logit lens + cosine similarity + block influence) found:
- Layer 10: safety divergence start (harmful-safe entropy gap = -1.11)
- Layer 20: peak harmful-safe differentiation (entropy gap = -2.39)
- Layer 14 (default): mid-processing confusion zone
Contrastive Feature Selection
Top features selected by Cohen's d between harmful and benign WildJailbreak activations:
- L20 x16: max |d| = 0.727, top-50 avg = 0.396
- L10 x8: max |d| = 0.394, top-50 avg = 0.254
Usage
import torch
from sae_lib import create_sae
# Load SAE
ckpt = torch.load("L20_x16_k50.pt", map_location="cpu", weights_only=True)
sd = ckpt.get("state_dict", ckpt)
sae = create_sae("topk", 1536, 24576, k=50)
sae.load_state_dict(sd, strict=False)
sae.eval()
# Encode activations (after normalizing with stats)
import numpy as np
stats = dict(np.load("layer20.npz"))
normed = (activations - stats["mean"]) / (stats["std"] + 1e-8)
_, features = sae(torch.tensor(normed, dtype=torch.float32))Citation
Part of the Neuron Scalpel project: SAE interpretability research on consumer hardware (Apple M3, 16GB).
License
MIT
