研究项目利用Sparse Autoencoders分析Llama 3.2的内部机制,增进对LLM工作原理的深度理解。对优化模型和Prompt有指导意义。
现代 LLM 通过将多个特征叠加到同一个神经元中来编码概念,然后通过计算一个层中所有神经元激活的线性叠加来解释它们。这种让每个神经元具有多个可解释含义且它们根据其他神经元激活的上下文而被激活的概念称为叠加(superposition)。稀疏自编码器(Sparse Autoencoders,SAE)是插入到已训练 LLM 中的模型,其目的是将激活投影到一个非常大但激活非常稀疏的隐空间。通过这样做,它们试图将这些叠加的表示解开为独立的、清晰可解释的特征,每个神经元激活都代表一个清晰的概念——这反过来会使这些神经元成为单义的。这样的机制可解释性已被证明对于理解模型行为、检测幻觉、分析信息流以进行优化等非常有价值。
本项目尝试再现由 Anthropic、OpenAI 和 Google DeepMind 在几个月前成功进行并发表的这项关于使用稀疏自编码器进行机制化 LLM 可解释性的优秀研究。该项目旨在提供一个完整的流水线,用于捕获训练数据、训练 SAE、分析学习的特征,然后通过实验验证结果。目前,该项目提供了通过运行整个项目流水线一次创建的所有代码、数据和模型,并为 Llama 3.2-3B 模型创建了一个功能性且可解释的稀疏自编码器。
显然,这样的研究项目需要大量的计算资源(意味着金钱)和时间,这些对于我的非营利性副项目来说不一定充足。因此,在我现在以 0.2 版本发布的这个项目中,处于良好、高效和可扩展的状态,但它并非最终版本,希望能随着时间推移得到更新和改进。欢迎贡献代码或反馈,或者如果发现 bug 请告诉我——感谢!
该项目主要基于以下研究论文:
Scaling Monosemanticity: Extracting Interpretable Features from Claude 3 Sonnet(Anthropic,2024 年 5 月)
Scaling and Evaluating Sparse Autoencoders(OpenAI,2024 年 6 月)
Gemma Scope: Open Sparse Autoencoders Everywhere All At Once on Gemma 2(Google DeepMind,2024 年 7 月)
以及项目当前状态使用的开源 LLM Llama 3.2:
完整的端到端流水线,从激活捕获到稀疏自编码器(SAE)训练、特征解释和验证,用纯 PyTorch 编写,依赖最少。具体来说:
使用自定义句子拆分 OpenWebText 数据集变体,从大语言模型捕获残差激活作为 SAE 训练数据集
预处理训练数据(预批处理、统计计算)以实现高效训练
支持使用捕获和预处理的激活数据进行分布式(多 GPU、单节点)大规模且高效的 SAE 训练
通过辅助损失防止和复活死亡潜在单元,以及梯度投影来稳定训练动态,从而实现 SAE 训练
通过 Weights & Biases 和控制台日志提供 SAE 训练的全面日志、可视化和检查点,包括详细日志:
通过以下方式提供可解释性分析工具进行特征提取和学习特征的语义分析:
提供纯 PyTorch 实现的 Llama 3.1/3.2 聊天和文本完成,不依赖外部库(例如 Fairscale),用于一般用途和结果验证
通过以下方式验证 SAE 对模型行为的影响并启用提取的语义特征的特征指导:
所有组件在设计和实现时都考虑了可扩展性、效率和可维护性
以下资源可用于重现项目当前状态的结果或提供训练的见解:
OpenWebText 句子数据集: 用于激活捕获的 OpenWebText 数据集自定义版本
捕获的 Llama 3.2-3B 激活: 25 百万句话的 Llama 3.2-3B 第 23 层残差激活
SAE 训练日志: Weights & Biases 可视化的训练、验证和调试指标日志
训练的 65,536 潜在单元 SAE 模型: 根据训练日志指定配置在 10 个时期后的最终训练 SAE 模型
该项目分为四个主要组件组织:
capture_activations.py - 捕获 LLM 残差激活openwebtext_sentences_dataset.py - 用于句子级处理的自定义数据集sae.py - 核心 SAE 模型实现sae_preprocessing.py - SAE 训练的数据预处理sae_training.py - 分布式 SAE 训练实现capture_top_activating_sentences.py - 识别最大化特征激活的句子interpret_top_sentences_send_batches.py - 构建和发送批次以进行解释interpret_top_sentences_retrieve_batches.py - 检索解释结果interpret_top_sentences_parse_responses.py - 解析和分析解释llama_3_inference.py - 核心推理实现llama_3_inference_text_completion_test.py - 文本完成测试llama_3_inference_chat_completion_test.py - 聊天完成测试llama_3_inference_text_completion_gradio.py - 交互式测试的 Gradio 界面# 如果尚未安装 Poetry 请安装
curl -sSL https://install.python-poetry.org | python3.12 -
# 克隆仓库
git clone https://github.com/PaulPauls/llama3_interpretability_sae
cd llama3_interpretability_sae
# 安装项目运行时使用的精确依赖
poetry install --sync
本研究以 llama_3/model_text_only.py 中自定义的 Llama 3.1/3.2 Transformer 模型实现为基础。该实现基于 Llama 模型仓库中的参考实现,但为了使其更适合本项目,我做了几项重大修改。我重写了实现,移除了对 Fairscale 库的深度依赖——原因很简单,我不熟悉这个库,更习惯直接使用 PyTorch,从而希望避免因使用不熟悉的库而产生的 bug 或性能瓶颈。同样,我也移除了多模态功能,因为在这个初始版本中研究图像可解释性会引入不必要的复杂性。修改后的 Llama 3.1/3.2 模型实现增加了新的构造函数选项,可在推理期间对指定层进行激活值捕获,或注入一个训练完成的稀疏自编码器(SAE)模型。此功能通过以下构造函数参数实现:
class Transformer(nn.Module):
def __init__(
self,
params: ModelArgs,
store_layer_activ: list[int] | None = None,
sae_layer_forward_fn: dict[int, callable] | None = None,
):
对于 llama_3/ 目录中的辅助代码,我做出了一个务实的决定:保留原始 Llama 模型仓库中的大部分辅助文件,不对其进行修改。这些辅助代码中有 95% 并未使用,仅有聊天格式化工具需要它们,而该工具高度依赖相互关联的导入。不过,由于这部分内容对研究并不关键,我决定不重写它,而只是将这些文件原样保留下来。实际的推理实现是自定义的,位于我的 llama_3_inference.py 模块中。该模块为聊天和文本补全任务提供流式处理能力,主要用于测试和验证结果。该实现支持批量推理,并提供可配置的 temperature 和 top-p 采样参数;当 temperature 设置为 0 时,会自动回退到贪心采样。
在数据捕获方面,我创建了一个自定义的 OpenWebText 数据集变体,以句子为单位处理文本,共捕获了 2,500 万个句子,每个句子的最大长度为 192 个 token。如此大规模的数据收集产生了 4TB 的原始激活值数据,压缩后为 3.2TB 的 tar.gz 归档文件。我总共从这 2,500 万个上下文中捕获了大约 7 亿个激活值,句子的平均长度为 27.3 个 token。
虽然该数据集的规模大约比 Anthropic 或 Google DeepMind 使用的数据集小一个数量级(两者都使用了约 80 亿个不同的激活值),但在我看来,它仍然为训练初始 SAE 模型提供了坚实的基础。为了弥补数据集规模较小的问题,我对 SAE 训练了 10 个 epoch,从而使处理的激活值总量实际上与 Anthropic 和 Google DeepMind 相当——关键区别在于,我的 SAE 会遇到每个激活值 10 次,而不是只遇到一次。采用这种方法纯粹是因为这是一个非营利性业余项目,受到资金限制。如果扩大规模以匹配他们的单 epoch 方案,我的 GCP 存储桶成本将从每月约 80 美元(3.2TB,含流量)增加到每月 800 美元(32TB,含流量),更不用说训练期间实例 SSD 的成本也会大幅增加。
以句子为单位处理数据是一个经过深思熟虑的决定,基于几个关键考量。句子是自然的语言单位,包含完整的思想和概念,我希望这能产生更具可解释性、语义上也更有意义的特征。这种方法避免了人为截断上下文,并防止含义跨越句子边界发生上下文渗透,同时仍能在语法完整的单位内捕获必要的上下文关系。我的想法是,这样可以更容易地将发现的特征归因于特定的语言或语义现象。选择这种方法还有一个原因:训练数据集中的激活值能够适配我随后计划用于可解释性分析的同类语言单位。
此外,我特意选择在处理句子时不添加“序列起始”(bos)token,以避免与位置相关的模式,因为我的目标是仅根据特征本身的含义来解释它们,而不受其在序列中位置的影响。句子是自然的语义单位,既能为有意义的解释提供充足的上下文,又足够具体,可以识别不同的概念,因此这种做法非常契合让 LLM 分析会激活特定潜变量的语义内容这一目标。
从技术实现角度来看,我捕获的是 Llama 3.2-3B 经过层归一化后的残差流激活值,具体取自 28 层中的第 23 层(位于模型深度接近六分之五的位置,遵循 OpenAI 的实现)。捕获过程采用基于 NCCL 的分布式实现,在单节点多 GPU 上执行推理;异步磁盘 I/O 则由单独的进程处理,以避免造成 GPU 处理瓶颈。使用 4 块 Nvidia RTX4090 GPU,整个数据捕获过程大约耗时 12 小时。
OpenAI 的实现同样使用经过层归一化后的残差流激活值,上下文长度为 64 个 token,并在 GPT-4 模型深度约六分之五的位置进行采样(不过对于较小的模型,他们会更多地在中间层附近采样)。Anthropic 的实现则有更大的差异:它从 4,000 万个上下文中使用了 80 亿个 MLP 激活向量(每个上下文 250 个 token),并将其应用于 Claude 3 Sonnet 模型中间层的残差流激活值。事后来看,对于一个参数量相对较小的 3B 模型,在下一轮实验中捕获更早层的激活值可能会带来改进。
如果需要进行复杂的批处理,我个人非常倾向于预处理数据。在本例中,挑战在于创建每批包含 1,024 个激活值的批次,同时还要处理长度不一的激活值序列,这需要处理跨批次结转,而且可能要在多进程环境中完成。考虑到批处理 bug 或 I/O 相关性能瓶颈的风险很高,我决定实现一个预处理阶段,而不是在训练期间处理这些复杂问题。
既然预处理已经不可避免,我也借此机会使用 Welford 算法计算了所有激活值张量的均值。选择该算法,主要是因为它在处理超大规模数据集时具有良好的数值稳定性和内存效率。计算得到的均值用作 SAE 模型中 b_pre 偏置项的初始值。OpenAI 的实现使用几何中位数而不是均值来初始化 b_pre,但他们只使用数据集最开始的大约 30,000 个样本进行计算,而不是完整数据集;而且 b_pre 无论如何都是可优化的,因此我认为使用均值作为初始近似已经足够。
整个预处理流水线通过多进程实现了完整的 CPU 并行化,确保能够高效处理大规模激活值数据集。这种方法通过提供干净且预先分批的数据,简化了训练过程。
稀疏自编码器实现的核心采用了简洁的编码器—解码器架构,其设计选择主要遵循 OpenAI 的方案。TopK 自编码器的完整前向传播包含两个关键偏置项:编码器和解码器共同使用的 b_pre(初始化为上一节所述预处理阶段计算得到的均值),以及编码器专用的 b_enc(随机初始化)。完整的前向传播可以描述为:
编码器:h = TopK(W_enc(x - b_pre) + b_enc)
解码器:x^ = W_dec * h (+ h_bias) + b_pre
潜在空间中的稀疏性通过 TopK 激活函数来强制实现:该函数仅保留最大的 k 个激活值,并将其余激活值设置为零。这种方法可以直接控制稀疏性,不需要像 Anthropic 的方案那样在损失函数中加入用于稀疏化的 L1 惩罚项。模型还包含一个可选的 h_bias 参数,该参数在训练期间保持禁用,但可以在训练完成后启用以进行特征引导,从而支持在训练后动态操纵潜在空间。
在数值精度方面,我选择使用 float32 dtype,因为它可以快速且精确地转换为 Llama 所需的 bfloat16 dtype。两种格式都采用 1 个符号位和 8 个指数位的结构,区别仅在于尾数位数不同(23 位与 7 位),因此二者之间的转换既快速又准确。
我的实现与 Anthropic 和 OpenAI 的方案在多个方面有所不同。Anthropic 使用带 ReLU 激活函数的单隐藏层 MLP,并通过 L1 惩罚而不是 TopK 来强制实现稀疏性。他们使用的潜在空间规模也大得多(约 100 万、约 400 万和约 3400 万个特征),不过对于所有 SAE,每个 token 平均激活的特征数量仍低于 300。OpenAI 的架构与我的实现更为相似,但他们在 GPT-4 上实验了从 2^11(2,048)到 2^24(16.7M)不等的潜在空间规模。他们的实验表明,更大的潜在空间规模通常能带来更低的损失和更好的特征可解释性,而较低的 k 值(激活的潜变量更少)则会产生更易解释的特征。
对于这个项目——使用残差流维度为 3,072 的 30 亿参数 Llama 3.2 模型——我选择了 2^16(65,536)的潜在空间维度和 64 的 k 值。这一决定旨在平衡多个因素:提供约为残差流维度 21 倍的充足特征容量;按照 OpenAI 和 Google DeepMind 论文的建议维持计算效率;以及在约 80 亿个激活值上进行训练以保证可比性的同时,将成本控制在项目预算范围内。选择 64 作为 k 值,是为了在重建能力与获得可解释特征所需的强稀疏性之间取得良好平衡。
不过,事后看来,正如我在第 6 节中所述,在未来的实验中,我会大幅增加潜在空间规模并降低 k 值,以提高特征的多样性和可解释性,同时尝试寻找效率方面的改进,从而继续满足预算限制。但作为这个项目第一次完整运行的结果,我对所选的超参数和最终结果非常满意。
稀疏自编码器的训练配置旨在平衡效率和特征可解释性。核心超参数同时反映了模型架构和训练动态:
# Set up configuration
d_model = 3072
n_latents = 2**16 # 65536
k = 64
k_aux = 2048
aux_loss_coeff = 1 / 32
dead_steps_threshold = 80_000 # ~1 epoch in training steps
sae_normalization_eps = 1e-6
batch_size = 1024
num_epochs = 10
early_stopping_patience = 10 # disabled
learning_rate = 5e-5
learning_rate_min = learning_rate / 5
optimizer_betas = (0.85, 0.9999)
optimizer_eps = 6.25e-10
dtype = torch.float32
dataloader_num_workers = 8
logs_per_epoch = 1000
train_val_split = 0.95
损失函数由重建误差产生的主要重建损失,以及用于防止潜变量失活并使其重新活跃的复杂辅助损失组成,具体形式如下:total_loss = main_loss + aux_loss_coeff * aux_loss。遵循 OpenAI 的方案,我将 aux_loss_coeff 设置为 1/32。两种损失都在归一化空间中计算,以确保所有特征无论原始尺度如何都能做出同等贡献,这有助于在整个训练过程中保持数值稳定性。
该辅助损失由 OpenAI 提出,并通过一种巧妙的机制在防止潜变量失活方面发挥着关键作用:它计算主要重建残差(输入与主要重建结果之间的差值)与一种特殊辅助重建结果之间的 MSE。这种辅助重建使用与主要重建相同的激活前潜变量,但只从近期未触发的潜变量中选取 top-(aux-k) 个激活值——这些潜变量通过共享的 stats_last_nonzero 张量进行跟踪——然后再次将它们传入解码器,得到这个“辅助重建”结果。在实际训练中,只有 top k 个潜变量用于重建;而这种机制为 k_aux = 2048 个甚至尚未激活的潜变量提供了专门的学习信号,使它们能够捕获主要潜变量遗漏的信息。这会提高失活潜变量在未来前向传播中被激活的概率,从而使失活潜变量重新活跃,并让所有潜变量保持活跃和有用。
训练过程只在计算辅助损失时考虑失活潜变量。如果一个潜变量在 dead_steps_threshold 个训练步骤内都未被激活,就会被视为失活;该阈值设置为 80,000 个步骤,在我的配置中约等于一个 epoch。批次大小为 8192 时,这相当于该潜变量在最近约 6.5 亿个激活值的重建过程中从未激活。该阈值有两个作用:一是确保主要损失在接收辅助损失信号之前拥有足够的预热时间;二是保证我们只尝试重新激活那些在观察所有不重复的训练数据激活值之后,仍一次都没有触发过的潜变量。
训练基础设施采用 NCCL 后端,在单节点多 GPU 配置下进行分布式训练。我使用 8 块 Nvidia RTX4090 GPU 训练了 10 个 epoch,每块 GPU 的批次大小为 1024(有效批次大小为 8192),在略多于 7 天的时间里处理了约 70 亿个激活值。epoch 数量的选择是为了与 Anthropic 和 Google DeepMind 实验中处理的激活值总数保持一致。所有训练进度,包括损失以及与失活潜变量相关的调试统计信息,都通过 Weights & Biases 进行了全面跟踪。
优化器参数经过精心调优,以应对训练稀疏自编码器时某些特征激活极其罕见这一特殊挑战。经过对比测试,我选择了 5e-5 作为基础学习率,因为结果表明,它能达到与更高学习率相近的优化速度,同时有望在训练后期为稀疏特征提供更好的微调能力。学习率遵循余弦退火调度,最低降至 1e-5(初始值的 1/5)。由于自编码器具有稀疏性质,AdamW 配置需要特别考虑:
beta_1 = 0.85(低于常见的 0.9;考虑到 8192 的大有效批次大小以及自编码器的稀疏性质,这能让每次更新产生更显著的影响)
beta_2 = 0.9999(适应稀疏激活模式;某些特征可能极少激活,因此需要更长时间地保留动量)
eps = 6.25e-10(在 float32 精度下提供足够的数值稳定性,同时允许进行优化罕见激活模式所需的精细参数更新)
按照 OpenAI 论文的建议,权重初始化和归一化的实现特别注重训练稳定性。编码器和解码器权重采用正交初始化(解码器为编码器的转置),以确保初始特征方向均衡且彼此独立。输入特征使用一个较小的 epsilon 项进行归一化,以增强训练的稳健性。根据 OpenAI 论文和 Bricken 等人 [2023] 的实证发现,解码器权重在初始化后以及每个训练步骤之后都会被显式归一化为单位范数,因为这样可以改善 MSE 表现。
一个关键的实现细节是通过 project_decoder_grads() 进行梯度投影。它会移除与现有字典向量平行的梯度分量,从而维持解码器权重的单位范数约束。这种投影有助于稳定训练,并防止自编码器在识别稀疏 pa