[ PROMPT_NODE_22640 ]
Mechanistic Interpretability Saelens API 参考
[ SKILL_DOCUMENTATION ]
# SAELens API 参考
## SAE 类
代表稀疏自编码器的核心类。
### 加载预训练的 SAE
python
from sae_lens import SAE
# 从官方发布加载
sae, cfg_dict, sparsity = SAE.from_pretrained(
release="gpt2-small-res-jb",
sae_id="blocks.8.hook_resid_pre",
device="cuda"
)
# 从 HuggingFace 加载
sae, cfg_dict, sparsity = SAE.from_pretrained(
release="username/repo-name",
sae_id="path/to/sae",
device="cuda"
)
# 从本地磁盘加载
sae = SAE.load_from_disk("/path/to/sae", device="cuda")
### SAE 属性
| 属性 | 形状 | 描述 |
|-----------|-------|-------------|
| `W_enc` | [d_in, d_sae] | 编码器权重 |
| `W_dec` | [d_sae, d_in] | 解码器权重 |
| `b_enc` | [d_sae] | 编码器偏置 |
| `b_dec` | [d_in] | 解码器偏置 |
| `cfg` | SAEConfig | 配置对象 |
### 核心方法
#### encode()
python
# 将激活编码为稀疏特征
features = sae.encode(activations)
# 输入: [batch, pos, d_in]
# 输出: [batch, pos, d_sae]
#### decode()
python
# 从特征重构激活
reconstructed = sae.decode(features)
# 输入: [batch, pos, d_sae]
# 输出: [batch, pos, d_in]
#### forward()
python
# 完整前向传播 (编码 + 解码)
reconstructed = sae(activations)
# 返回重构后的激活
#### save_model()
python
sae.save_model("/path/to/save")
---
## SAEConfig
用于 SAE 架构和训练上下文的配置类。
### 关键参数
| 参数 | 类型 | 描述 |
|-----------|------|-------------|
| `d_in` | int | 输入维度 (模型的 d_model) |
| `d_sae` | int | SAE 隐藏层维度 |
| `architecture` | str | "standard", "gated", "jumprelu", "topk" |
| `activation_fn_str` | str | 激活函数名称 |
| `model_name` | str | 源模型名称 |
| `hook_name` | str | 模型中的钩子点 |
| `normalize_activations` | str | 归一化方法 |
| `dtype` | str | 数据类型 |
| `device` | str | 设备 |
### 访问配置
python
print(sae.cfg.d_in) # GPT-2 small 为 768
print(sae.cfg.d_sae) # 例如 24576 (32 倍扩展)
print(sae.cfg.hook_name) # 例如 "blocks.8.hook_resid_pre"
---
## LanguageModelSAERunnerConfig
用于训练 SAE 的综合配置。
### 配置示例
python
from sae_lens import LanguageModelSAERunnerConfig
cfg = LanguageModelSAERunnerConfig(
# 模型和钩子
model_name="gpt2-small",
hook_name="blocks.8.hook_resid_pre",
hook