稀疏自编码器解锁LLM可解释性
用直观方式讲解如何通过稀疏自编码器窥视大模型内部工作机制,对AI研究者有参考价值。
用直观方式讲解如何通过稀疏自编码器窥视大模型内部工作机制,对AI研究者有参考价值。
稀疏自编码器(Sparse Autoencoders,SAE)最近在机器学习模型可解释性领域受到广泛关注(尽管稀疏字典学习早在 1997 年就已出现)。机器学习模型和 LLM 正变得越来越强大、越来越实用,但它们仍然是黑箱,我们并不理解它们如何完成自己所具备的那些能力。如果我们能够理解它们的工作原理,显然会非常有用。
借助 SAE,我们可以开始将模型的计算过程分解为可理解的组成部分。目前已经有一些关于 SAE 的解释,而我希望从不同角度写一篇简短的文章,直观地说明它们如何工作。
神经网络中最自然的组成单元是单个神经元。遗憾的是,单个神经元并不能恰好对应单一概念。例如,语言模型中的某个神经元同时对应学术引用、英语对话、HTTP 请求和韩语文本。这种现象被称为叠加(superposition),即神经网络中的概念由多个神经元的组合表示。
出现这种现象的原因,可能是现实世界中的许多变量天然就是稀疏的。例如,某位名人的出生地可能在十亿个训练 token 中出现不到一次,但现代 LLM 仍会学到这一事实,以及关于世界的数量惊人的其他事实。叠加现象之所以出现,可能是因为训练数据中的独立事实和概念数量超过了模型中的神经元数量。
稀疏自编码器最近已成为一种广受关注的技术,用于将神经网络分解为可理解的组成部分。SAE 的灵感来自神经科学中的稀疏编码假说。有趣的是,SAE 是解释人工神经网络最具潜力的工具之一。SAE 与标准自编码器类似。
常规自编码器是一种用于先压缩输入数据、再重建输入数据的神经网络。例如,它可以接收一个 100 维向量(由 100 个数字组成的列表)作为输入,通过编码器层将输入压缩为一个 50 维向量,然后让压缩后的编码表示通过解码器,生成一个 100 维输出向量。由于压缩增加了任务难度,重建结果通常并不完美。
稀疏自编码器会将输入向量转换为一个中间向量,该向量的维度可以高于、等于或低于输入向量。应用于 LLM 时,中间向量的维度通常大于输入向量。在这种情况下,如果没有额外约束,任务将变得微不足道:SAE 可以使用单位矩阵完美重建输入,却无法向我们揭示任何有意义的信息。作为额外约束,我们会在训练损失中加入稀疏性惩罚,促使 SAE 创建稀疏的中间向量。例如,我们可以将 100 维输入扩展为一个 200 维编码表示向量,并训练 SAE,使编码表示中只有约 20 个非零元素。
我们将 SAE 应用于神经网络内部的中间激活,而神经网络可能由许多层组成。在一次前向传播过程中,每一层内部以及各层之间都会产生中间激活。例如,GPT-3 有 96 层。在前向传播期间,输入中的每个 token 都对应一个 12,288 维向量(由 12,288 个数字组成的列表),该向量会逐层传递。随着每一层对它进行处理,这个向量会累积模型用于预测下一个 token 的全部信息,但它是不透明的,我们很难理解其中包含哪些信息。
我们可以使用 SAE 来理解这种中间激活。SAE 基本上就是矩阵 -> ReLU 激活 -> 矩阵¹²。举例来说,如果 GPT-3 的 SAE 扩展因子为 4,那么输入激活是 12,288 维,SAE 的编码表示则是 49,512 维(12,288 x 4)。第一个矩阵是形状为 (12,288, 49,512) 的编码器矩阵,第二个矩阵是形状为 (49,512, 12,288) 的解码器矩阵。将 GPT 的激活与编码器相乘并应用 ReLU 后,我们会得到一个 49,512 维的 SAE 编码表示。由于 SAE 的损失函数会鼓励稀疏性,因此该表示是稀疏的。通常,我们的目标是让 SAE 表示中少于 100 个数字为非零值。将 SAE 表示与解码器相乘后,我们会得到一个 12,288 维的重建模型激活。由于稀疏性约束增加了任务难度,这个重建结果无法与原始 GPT 激活完全一致。
我们只针对模型中的一个位置训练每个 SAE。例如,我们可以针对第 26 层与第 27 层之间的中间激活训练一个 SAE。要分析 GPT-3 全部 96 层输出中包含的信息,就需要训练 96 个独立的 SAE——每一层的输出对应一个 SAE。如果还想分析每一层内部的各种中间激活,就需要数百个 SAE。这些 SAE 的训练数据来自:将各种不同的文本输入 GPT 模型,并收集每个选定位置上的中间激活。
下面给出了一个 SAE 的参考 Pytorch 实现。变量的形状标注遵循 Noam Shazeer 的建议。请注意,不同的 SAE 实现通常会采用不同的偏置项、归一化方案或初始化方案,以进一步提升性能。最常见的附加机制之一,是对解码器向量的范数施加某种约束。有关更多细节,请参阅 OpenAI、SAELens 或 dictionary_learning 等实现。
import torch
import torch.nn as nn
# D = d_model, F = dictionary_size
# e.g. if d_model = 12288 and dictionary_size = 49152
# then model_activations_D.shape = (12288,) and encoder_DF.weight.shape = (12288, 49152)
class SparseAutoEncoder(nn.Module):
"""
A one-layer autoencoder.
"""
def __init__(self, activation_dim: int, dict_size: int):
super().__init__()
self.activation_dim = activation_dim
self.dict_size = dict_size
self.encoder_DF = nn.Linear(activation_dim, dict_size, bias=True)
self.decoder_FD = nn.Linear(dict_size, activation_dim, bias=True)
def encode(self, model_activations_D: torch.Tensor) -> torch.Tensor:
return nn.ReLU()(self.encoder_DF(model_activations_D))
def decode(self, encoded_representation_F: torch.Tensor) -> torch.Tensor:
return self.decoder_FD(encoded_representation_F)
def forward_pass(self, model_activations_D: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
encoded_representation_F = self.encode(model_activations_D)
reconstructed_model_activations_D = self.decode(encoded_representation_F)
return reconstructed_model_activations_D, encoded_representation_F
标准自编码器的损失函数以输入重建的准确度为基础。为了引入稀疏性,早期 SAE 实现会在 SAE 的损失函数中加入稀疏性惩罚。最常见的惩罚形式是:计算 SAE 编码表示(而不是 SAE 权重)的 L1 损失,然后乘以一个 L1 系数。L1 系数是 SAE 训练中的关键超参数,因为它决定了实现稀疏性与保持重建准确度之间的权衡。
请注意,我们并没有直接优化可解释性。相反,可解释的 SAE 特征是优化稀疏性和重建效果时产生的副作用。
下面是一个参考损失函数。
# B = batch size, D = d_model, F = dictionary_size
def calculate_loss(autoencoder: SparseAutoEncoder, model_activations_BD: torch.Tensor, l1_coeffient: float) -> torch.Tensor:
reconstructed_model_activations_BD, encoded_representation_BF = autoencoder.forward_pass(model_activations_BD)
reconstruction_error_BD = (reconstructed_model_activations_BD - model_activations_BD).pow(2)
reconstruction_error_B = einops.reduce(reconstruction_error_BD, 'B D -> B', 'sum')
l2_loss = reconstruction_error_B.mean()
l1_loss = l1_coefficient * encoded_representation_BF.sum()
loss = l2_loss + l1_loss
return loss
2024 年 11 月 29 日更新:我认为普通的 ReLU SAE 已经相当过时了,除非将其用作基线,否则不应再使用。我更倾向于 BatchTopK SAE,因为它显著改善了稀疏度与重建准确率之间的权衡;无需调节稀疏惩罚,就能直接设定期望的稀疏度;而且训练稳定性良好。BatchTopK SAE 与 ReLU SAE 非常相似。它不使用 ReLU 和稀疏惩罚,而是直接保留最大的 k 个激活值,并将其余激活值归零。在这种情况下,超参数 k 会直接设定期望的稀疏度。这里可以看到一个 BatchTopK 实现示例。其他表现优秀的替代方案包括 TopK SAE 和 JumpReLU SAE。
理想情况下,SAE 表示中的每个活跃数值都对应某个可以理解的组成部分。举一个假想的例子,假设对于 GPT-3 而言,12,288 维向量 [1.5, 0.2, -1.2, ...] 表示“金毛寻回犬”。SAE 解码器是一个形状为 (49,512, 12,288) 的矩阵,但我们也可以将其视为由 49,512 个向量组成的集合,其中每个向量的形状都是 (1, 12,288)。如果 SAE 的第 317 个解码器向量学到了与 GPT-3 相同的“金毛寻回犬”概念,那么该解码器向量将近似等于 [1.5, 0.2, -1.2, ...]。每当 SAE 激活中的第 317 个元素非零时,一个与“金毛寻回犬”相对应的向量就会被添加到重建后的激活中,并按照第 317 个元素的大小进行缩放。用机制可解释性领域的术语,可以简洁地描述为:“解码器向量对应残差流空间中特征的线性表示。”
从向量与矩阵乘法的数学原理来看,这在直觉上很容易理解。一个向量乘以一个矩阵,本质上就是对矩阵的各行(或各列,取决于乘法顺序)进行加权求和,其中权重就是该向量的各个元素。在我们的例子中,SAE 的稀疏编码表示充当这些权重,通过选择性地激活相关的解码器向量(矩阵的行)并对其进行缩放,重建原始激活。
我们也可以说,这个编码表示为 49,512 维的 SAE 拥有 49,512 个特征。每个特征由相应的编码器向量和解码器向量组成。编码器向量的作用是检测模型内部的概念,同时尽量减少与其他概念之间的干扰;解码器向量的作用则是表示“真正的”特征方向。根据实证结果,每个特征的编码器向量与解码器向量并不相同,两者余弦相似度的中位数为 0.5。在下图中,三个红框对应同一个特征。
我们怎么知道假想的第 317 个特征表示什么?目前的做法只是查看能够最大程度激活该特征的输入,然后凭直觉判断其可解释性。每个特征所响应的输入通常是可以解释的。例如,Anthropic 在 Claude Sonnet 上训练了 SAE,并发现了彼此独立的 SAE 特征,分别会被与金门大桥、神经科学和热门旅游景点相关的文本及图像激活。另一些特征所响应的概念则不那么显而易见,例如,在 Pythia 上训练的一个 SAE 中存在这样一个特征:它会“在修饰句子主语的关系从句或介词短语的最后一个词元上”激活。
由于 SAE 解码器向量的形状与 LLM 的中间激活相匹配,我们只需将解码器向量添加到模型激活中,就能实施因果干预。通过将解码器向量乘以一个缩放因子,我们可以调节干预强度。当 Anthropic 的研究人员将表示金门大桥的 SAE 解码器向量添加到 Claude 的激活中时,Claude 会被迫在每次回答中提到金门大桥。
下面是一个使用假想的第 3173 个特征实施因果干预的参考实现。与“金门大桥 Claude”类似,这个非常简单的干预会迫使我们的 GPT-3 模型在每次回答中都提到金毛寻回犬。
def perform_intervention(model_activations_D: torch.Tensor, decoder_FD: torch.Tensor, scale: float) -> torch.Tensor:
intervention_vector_D = decoder_FD[317, :]
scaled_intervention_vector_D = intervention_vector_D * scale
modified_model_activations_D = model_activations_D + scaled_intervention_vector_D
return modified_model_activations_D
使用 SAE 的主要挑战之一在于评估。我们训练稀疏自编码器是为了理解语言模型,但在自然语言中,我们没有可测量的底层真实答案。目前,我们的评估是主观的,基本上相当于:“我们查看了一系列特征的激活输入,然后凭直觉判断这些特征的可解释性。”这是可解释性领域的一项重大局限。
研究人员发现了一些似乎与特征可解释性相关的常用代理指标,其中最常用的是 L0 和 Loss Recovered。L0 是 SAE 编码后的中间表示中非零元素的平均数量。Loss Recovered 的计算方法是,用重建后的激活替换 GPT 的原始激活,然后测量重建不完美所导致的额外损失。这两个指标之间通常存在权衡,因为 SAE 可能会选择牺牲重建准确率,以获得更高稀疏度的解。
一种常见的 SAE 比较方式,是将这两个变量绘制成图,并考察二者之间的权衡⁴。许多新的 SAE 方法,例如 DeepMind 的 Gated SAE 和 OpenAI 的 TopK SAE,都会修改稀疏惩罚,以改善这一权衡。下图来自 Google DeepMind 的 Gated SAE 论文。表示 Gated SAE 的红线更接近图的左上角,这意味着它在这一权衡上的表现更好。
SAE 的测量难题包含多个层次。我们的代理指标是 L0 和 Loss Recovered。然而,我们在训练时并不会使用这些指标,因为 L0 不可微,而且在 SAE 训练过程中计算 Loss Recovered 的计算成本很高⁵。相反,我们的训练损失由 L1 惩罚和内部激活的重建准确率决定,而不是由重建结果对下游损失的影响决定。
我们的训练损失函数并不直接对应这些代理指标,而这些代理指标本身又只是我们对特征可解释性进行主观评估的代理。这里还存在另一层错位,因为我们对可解释性的主观评估,也只是我们真正目标——“这个模型如何工作”——的代理。LLM 内部的一些重要概念可能并不容易解释;如果盲目优化可解释性,我们可能会忽略这些概念。
如需更详细地了解 SAE 评估方法,以及一种使用棋盘游戏模型 SAE 的评估方案,请参阅我的博客文章《使用棋盘游戏模型评估稀疏自编码器》。
可解释性领域还有很长的路要走,但 SAE 代表着真正的进步。它们催生了一些有趣的新应用,例如以无监督方式寻找类似“金门大桥”引导向量的引导向量。SAE 也让我们更容易在语言模型中发现回路,而这些回路可能被用于消除模型内部不必要的偏见。
尽管 SAE 的目标仅仅是在激活中识别模式,但它们依然能够发现可解释的特征,这表明它们确实揭示了一些有意义的东西。这也证明了 LLM 学到的内容具有实际意义,而不只是记住了表层统计规律。
它们也代表着 Anthropic 等公司所追求的一个早期里程碑,即“为机器学习模型打造 MRI”。目前,它们还无法提供完美的理解,但可能有助于检测不良行为。SAE 以及 SAE 评估所面临的挑战并非无法克服,并且仍是大量正在进行的研究所关注的主题。
如需进一步学习稀疏自编码器,我推荐 Callum McDougal 的 Colab 笔记本。
致谢:感谢 Justis Mills、Can Rager、Oscar Obeso 和 Slava Chalnev 为本文提供的宝贵反馈。
ReLU 激活函数就是 y = max(0, x)。也就是说,任何负数输入都会被设为 0。↩
ReLU 激活函数就是 y = max(0, x)。也就是说,任何负数输入都会被设为 0。↩
通常还会在多个位置设置偏置项,包括编码器层和解码器层。↩
通常还会在多个位置设置偏置项,包括编码器层和解码器层。↩
注意,该函数会在单一层上进行干预,且 SAE 应该在与模型激活相同的位置进行训练。例如,如果干预在第 6 层和第 7 层之间进行,那么 SAE 应该在第 6 层和第 7 层之间的模型激活上进行训练。干预也可以同时在多个层上进行。↩
注意,该函数会在单一层上进行干预,且 SAE 应该在与模型激活相同的位置进行训练。例如,如果干预在第 6 层和第 7 层之间进行,那么 SAE 应该在第 6 层和第 7 层之间的模型激活上进行训练。干预也可以同时在多个层上进行。↩
值得注意的是,这只是一个代理指标,改进这种权衡可能并不总是更好的。如最近的 OpenAI TopK SAE 论文所述,一个无限宽的 SAE 可以用 L0 为 1 实现完美的 Loss Recovered,但这会完全没有意义。↩
值得注意的是,这只是一个代理指标,改进这种权衡可能并不总是更好的。如最近的 OpenAI TopK SAE 论文所述,一个无限宽的 SAE 可以用 L0 为 1 实现完美的 Loss Recovered,但这会完全没有意义。↩
Apollo Research 最近发布了一篇论文,使用了旨在生成相同输出分布(而不是重建单一层激活)的损失函数。这种方法效果更好,但计算成本也更高。↩
Apollo Research 最近发布了一篇论文,使用了旨在生成相同输出分布(而不是重建单一层激活)的损失函数。这种方法效果更好,但计算成本也更高。↩