最近有个做NLP的师弟来问我,说自己读可解释性论文时看到一堆Probing实验,结论写得一个比一个花:什么“BERT内部编码了句法树”“词向量里藏着性别偏见”,但他始终没想明白一个问题——线性探针能把一个属性解出来,是不是就代表模型在做预测时真的用了这个属性?这问题问到点子上了。说实话,我自己刚接触Probing时也天真地以为探针是“读心术”,后来踩了不少坑才意识到:探针读到的信息,跟模型实际用到的信息,中间隔着一条很宽的河。这篇文章就把这条河讲清楚,也顺手聊聊线性探针实验从设计到落地那些容易翻车的地方。
1. 先搞清楚线性探针到底在干什么
1.1 探针实验的基本玩法
如果你还没接触过探针,我先给个最简单的定义:线性探针是在某个模型中间层的隐藏表示上,训练一个线性分类器或者线性回归头,用来判断这个表示里是否“藏”着某个属性。例如在词向量中判断词的时态、情感极性,在图像特征中判断物体的颜色、形状,在语音特征中判断说话人身份。
它的原理很朴素:如果某一层表示的某个方向上存在与属性相关的线性模式,那么一个简单的线性分类器就能利用这个方向把不同类别分开。换句话说,线性可分性意味着这些信息在表示中是“摆在明面上”的。我们经常在论文里看到这样的结论:某层隐藏状态经过逻辑回归后,情感分类准确率能到90%,于是认为该层编码了情感信息。
为什么大家都用线性探针而不是多层感知机探针?这是很多新手容易忽略的点。非线性探针的拟合能力太强,它可以任意重排特征组合,把一个只在非常复杂的非线性路径上存在的信息也硬“解”出来。这样一来,即使探针分数很高,我们也不清楚信息是以什么形式存储的——是简单线性方向,还是经过了多次特征交互?线性探针限制在一个超平面上,如果它能成功,至少说明这些信息在表示空间里具备线性可分的结构。这种结构在后续模型中更容易被直接消费,也更接近我们直觉上的“编码”。
1.2 线性探针为什么是“线性”的
打个比方,一个仓库堆满了货物,你能在仓库侧面开一扇小窗户看到里面有啤酒瓶,但不能因此断定仓库里的传送带分拣货物时真的把啤酒瓶当作判断依据。你可能只是从窗户缝里看到了堆在外面的几个瓶子而已。线性探针就是这个窗户,它只能看到窗口对准的方向上的东西,而且看到的仅仅是“有没有”,不是“用不用”。
在表示空间里,每个词或者每个样本都被映射成一个高维向量。线性探针做的事情,就是在这个高维空间里画一个超平面,把不同类别的点分开。如果分布在超平面两侧的点能很好地对应到不同标签,那说明这个表示空间里存在一个线性方向,与标签是相关的。但注意关键词是“相关”,不是“因果”。
实操中,探针通常会做成一个不带隐藏层的逻辑回归,或者一个单层线性层加Softmax。如果需要在PyTorch里快速验证,我一般直接用torch.nn.Linear(hidden_size, num_labels),配合交叉熵损失训练几十个epoch。另一种更省事的做法是把中间层特征导出来,用sklearn.linear_model.LogisticRegression直接拟合,省去自己写训练循环。
这里有一个容易被忽视的细节:在训练线性探针时,要不要冻结骨干模型?答案是必须冻结。探针实验的目的是分析已经训练好的模型内部表示,如果骨干模型还在更新,探针组合骨干一起训练,最后得到的可能是一套新的联合表示,这就失去了“探测”的意义。所有探针实验都应该在torch.no_grad()下提取特征,再单独训练探针。
2. 核心误区:探针检测到的信息,不等于模型真正依赖的信息
2.1 一个直观反例:颜色特征与形状分类
假设你在训练一个区分猫和狗的图像分类模型。训练数据里所有猫的照片都是红色背景,狗的照片都是蓝色背景。模型很可能学到一个捷径:根据背景颜色分类,而不是真正的猫狗形状特征。这时你在最后一层卷积特征上训练一个线性探针,它能非常轻松地区分红蓝背景,进而区分猫和狗,准确率可能接近100%。
那么这个探针读到的“颜色信息”,能说模型“用”了吗?确实能,模型在训练集分布内就是靠颜色在做判断。但是如果我们换一组背景色反转的数据,模型马上崩盘。这个例子说明,探针检测到了“模型在数据上可以利用的特征”,但并没有揭示模型是否学会了我们真正关心的泛化性特征。更麻烦的是,探针检测到的信息可能根本不在模型的决策路径上。比如一个情感分类模型可能依赖关键词“好”“差”,而探针也能从中间表示中解码出这些关键词的存在,但模型真正用于情感判断的可能是更深层的语义结构。探针读到的关键词信息只是相关性,不是因果性。
我在实际项目中遇到过类似情况。有一次想验证对话模型是否“记住”了用户性别,于是在最后一层特征上训练性别探针,准确率很高。但后来做干预实验,把性别相关的特征方向进行翻转,模型的对话输出几乎没有变化。这说明性别信息虽然存在于表示里,但模型在做回复生成时并没有把它作为关键依赖。这种情况下,如果只汇报探针结果,就会给读者造成严重误导。
2.2 信息存储与信息使用的“两条路”
要理解这个问题,得先接受一个事实:现代神经网络的中间表示是一个“信息杂货铺”。训练任务会促使模型保留很多与任务相关的、有区分力的信息,因为它不知道下游哪些细节有用。于是每一层都塞满了各种属性:句法、语义、词序、词性,甚至输入文本的长度、标点习惯都可能被编码。这些信息不一定都会被后续模块利用,有些可能只是作为中间计算的副产物被保留下来。
举一个生活化的类比:一个人的日记里写了很多想法、情绪、对未来的计划,但你读出这些内容,只能说明它们“存在于日记中”,不能直接断定这个人的下一个行动完全受这些文字驱动。可能他行动时只看了其中一条备忘录,剩下的都是背景噪音。模型内部也是类似,表示空间中的“可用信息”和“实际利用信息”是两回事。
探针测量的主要是前者——表示中是否编码了某个属性。后者需要用更严格的因果干预来判断,比如激活替换、消融某个方向、logit lens,或者用扰动测试观察模型输出是否显著改变。这些方法各有局限,但结合起来才能逼近“模型是否真的使用”这个结论。
2.3 探针容量陷阱:线性与非线性之间的灰色地带
线性探针只能测试线性可分的属性,非线性属性会被它漏掉。但这并不意味着非线性属性不存在。反过来,当表示维度非常高时,线性探针也可能过拟合,在训练集上表现得非常好,但在测试集上一塌糊涂。更隐蔽的问题是,高维空间中的线性分类器有时可以“近似”某些非线性决策边界,从而造成一种“信息可以被线性解码”的假象。
有一个典型现象:在BERT的768维表示上,你用2000个样本训练线性探针,即便标签是随机分配的,训练集准确率也能轻松超过90%。这就是高维空间过拟合的威力。所以做探针实验时,不能只看训练集分数,也不能只看测试集分数,还要和随机标签的基线做对比。
我自己常用的一个基线方法是:把标签随机打乱后重新训练一个同样结构的线性探针,重复多次取平均准确率。如果真实标签的准确率只是比随机基线高几个点,那这个探针结论基本没有说服力。更严谨的做法是做置换检验,计算p值,但工程实践中随机标签基线已经能过滤掉大部分噪声。
3. 手把手跑一个线性探针实验(附PyTorch代码)
3.1 实验设计思路
实验目的很明确:验证BERT某个隐藏层中是否包含情感信息,同时用控制变量的方法证明这个结论不是靠过拟合或数据泄露得到的。用情感分类数据集,取每一句的[CLS]向量作为句子表示,训练一个线性逻辑回归探针。之所以用[CLS],是因为BERT在预训练时把这个位置的输出设计成聚合整个序列的信息,很多下游分类任务都直接用它。
实际操作步骤:
- 加载预训练模型和分词器。
- 对输入句子做分词、padding、truncation,统一长度。
- 前向传播时设置
output_hidden_states=True,取出某一层的隐藏状态。 - 取
[CLS]位置的向量作为整句特征。 - 将所有特征拼接成矩阵,划分训练集和测试集。
- 用逻辑回归训练探针,评估准确率。
整个过程不需要更新BERT参数,所以用torch.no_grad()包住即可,显存占用也小。
3.2 完整代码与逐步解读
下面这段代码可以直接在Jupyter Notebook里跑通,依赖torch、transformers、datasets、scikit-learn。
import numpy as np import torch from transformers import AutoTokenizer, AutoModel from sklearn.linear_model import LogisticRegression from sklearn.model_selection import train_test_split from sklearn.metrics import accuracy_score from datasets import load_dataset model_name = "bert-base-uncased" tokenizer = AutoTokenizer.from_pretrained(model_name) model = AutoModel.from_pretrained(model_name).to("cuda") model.eval() # 加载一个小型情感数据集,这里以 imdb 的前400条为例 dataset = load_dataset("imdb", split="train[:400]") texts = dataset["text"] labels = dataset["label"] def get_cls_embedding(texts, layer=12): feats = [] with torch.no_grad(): for text in texts: encoded = tokenizer( text, return_tensors="pt", truncation=True, max_length=64, padding="max_length", ).to("cuda") out = model(**encoded, output_hidden_states=True) # hidden_states 是一个元组,包含每一层的输出 hidden = out.hidden_states[layer] # [B, L, D] # 取 [CLS] token 对应的向量 feats.append(hidden[:, 0, :].squeeze(0).cpu().numpy()) return np.array(feats) X = get_cls_embedding(texts) y = np.array(labels) X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.2, random_state=42, stratify=y ) probe = LogisticRegression(max_iter=1000, C=1.0) probe.fit(X_train, y_train) y_pred = probe.predict(X_test) print("linear probe acc:", accuracy_score(y_test, y_pred))这里有两个细节值得展开说。第一,max_length=64不是拍脑袋定的,情感分类的句子一般不会太长,64足够覆盖大部分样本,同时能减少padding带来的干扰。如果你处理长文档,可以调大到512,但要注意BERT有位置编码上限。第二,C=1.0是逻辑回归的正则化强度,值越小正则越强。如果样本量少或维度高,建议把C调小到0.1甚至0.01,能有效缓解过拟合。
跑出来的准确率通常在90%以上,因为IMDB情感分类任务比较简单,BERT表示里的情感信息本身就非常显著。但你可能会问:这个结果能说明“BERT内部编码了情感”吗?如果只看这个单一实验,我倾向于说“BERT的最后一层表示中存在可与情感标签线性对齐的方向”,但不会直接说“模型做预测时依赖这些方向”。后者需要用因果探针或者特征干预才能下结论。
3.3 探针的四种常见变体
除了上面这种最简单的逻辑回归,做研究时还会用到其他变体。
- 线性层加交叉熵:在PyTorch里构造一个无隐藏层的分类头,和骨干模型一起做梯度评估不更新,但探针本身可以训练。这种方式适合大批量数据,而且可以方便地加入正则化。
- 最小描述长度探针(MDL Probe):通过在线编码方式计算用少量比特就能编码标签,得分越低说明信息越容易被压缩提取。MDL的好处是天然防止过拟合,因为它惩罚了过于复杂的探针。
- 迭代剪枝探针:反复训练探针后,把注意力权重接近零的特征剪掉,再重新训练,直到准确率下降。它可以在一定程度上去除掉冗余特征,告诉你哪些维度真正贡献了分类能力。
- CKA线性探针:用中心核对齐计算表示与标签之间的相似度,适合比较不同层或不同模型之间的表示差异。
实际写代码时,迭代剪枝会稍微麻烦一点,但很值得尝试。我常用它来回答“模型到底用了哪几个维度”的问题。做法是在逻辑回归权重上取绝对值,把权重最小的维度删掉,重训,重复多次,准确率曲线下降越慢,说明信息分布越冗余;如果删了少数维度准确率就崩了,说明关键维度很集中。
3.4 让结果更可信的控制变量设计
单次探针实验很容易被审稿人或者同事质疑,所以我在正式实验里一定会做三组对照。
第一组是随机标签基线。把训练和测试标签同时打乱,重新训练探针,得到一组随机分数。如果真实标签的探针分数只是略高于随机基线,就别再讨论“编码了”什么了。第二组是随机初始化模型对比。用同一个架构但权重完全随机初始化,提取特征训练探针。这能排除“模型结构本身导致输出特征容易被线性分离”的可能性。第三组是跨领域验证。在训练集上拟合探针,在另一个来源不同的测试集上评估。如果准确率大幅下降,说明探针可能学到了数据集的表面特征,而不是通用属性。
这三组对照不需要花很多代码,但对结论的可信度提升是决定性的。现在很多探针论文被吐槽,就是因为只报了一个漂亮ACC,没有做任何基线控制。
4. 常见翻车现场与排查技巧
4.1 特征泄露:探针学到了不该学的东西
文本探针最经典的特征泄露是句子长度。假设在做情感分类,负面评论普遍偏长,模型表示里可能包含了长度信息,线性探针能通过长度方向来分类,准确率还很高。但那不是情感信息,而是长度信息。类似的情况还有很多,比如某些数据集里某个类别的样本都集中在特定主题,探针实际上学的是主题分类,而不是你声称的属性。
怎么排查?把探针的输入特征降到二维或者三维可视化一下,看类别之间是不是有明显距离。更定量的方法是看探针权重向量与可能的混淆属性方向向量之间的余弦相似度。如果相似度很高,说明探针确实在用混淆属性。预防手段则是做特征白化或控制协变量,让探针无法依赖这些混杂特征。
4.2 标签泄漏与数据划分错误
数据划分错误在探针实验里特别隐蔽。最常见的翻车是:同一个原始句子的多种数据增强版本同时出现在训练集和测试集。比如你用回译做了数据增强,删除重复时会保留多个相似句子,如果随机划分,模型就会在训练时“见过”测试句子的某种变体,探针准确率虚高。解决方法是按原始句子ID分组划分,保证任何组内句子只出现在训练集或测试集之一。
另一个容易踩坑的地方是特征提取和探针训练共用同一个模型状态。比如你提取特征时开启了dropout,那每次前向传播得到的特征都不同,探针会不稳定。务必在特征提取阶段设置model.eval(),并且固定种子。
4.3 探针太强或太弱怎么判断
训练集准确率接近100%,测试集准确率却比随机基线高不了多少,这是典型的探针过拟合。常见原因有三个:表示维度太高、样本太少、正则化太弱。第一步加L2正则化,把C调小;如果还不行,就对特征做PCA降维,保留能解释95%方差的前几十个主成分;再不行就增加样本量,或者换成MDL探针。
相反的情况是训练集准确率也很低,说明这个属性在当前层确实不是线性可分的。但这不代表表示里没有这个信息。建议先用非线性探针复测,比如一个两层MLP,如果非线性探针能到高分,那说明信息以非线性方式存在。然后再考虑是不是需要换一层,因为不同层的信息抽象程度不同。
这里放一个简单的排查速查表:
| 现象 | 可能原因 | 处理方式 |
|---|---|---|
| 训练ACC高,测试ACC低 | 探针过拟合 | 减小C,加PCA,增加数据 |
| 训练ACC低,测试ACC也低 | 信息线性不可分 | 换非线性探针或换层 |
| 真实标签与随机标签ACC接近 | 探针无区分能力 | 检查特征提取,改变层位置 |
| 换测试集后ACC大跌 | 数据分布相关,探针学到表面特征 | 做跨域验证,控制混杂变量 |
| 多次运行结果波动大 | 未固定种子或特征提取未关dropout | 固定种子,设置model.eval() |
4.4 探针实验的黄金组合
我的习惯是,线性探针只做初筛,发现某个属性ACC高时,绝不去直接写“模型编码了这个属性”,而是补三样东西:随机化特征扰动、因果干预、跨域测试。随机化特征扰动是把特征某一维或多维打乱,看ACC变化,如果变化很小说明探针对依赖的特征不敏感;因果干预是在模型前向传播时替换或者消除对应方向上的激活值,观察预测结果是否改变;跨域测试就是换一个数据集验证探针是否抓住了通用规律。
这三样都做下来,才敢说“模型在推理时使用这个信息”。虽然在论文里全做完很费时间,但如果目标是真正理解模型行为,这个成本是必须的。我在实际项目中靠着这套流程,至少避免了两三次“假阳性结论”带来的误导,少走了很多弯路。
最后分享一个小技巧,做探针实验时一定要保存特征矩阵和标签,不要只保存ACC。这样后续做消融、可视化、换探针模型时都不用重新跑一遍骨干模型的前向传播,省时省力。我自己吃过亏,原本只存了ACC,后来审稿人要求补充特征可视化,只能把几千条数据重新过了一遍模型,浪费了不少时间。特征矩阵通常不会太大,存成.npy文件就好。