一种基于因果干预的少样本学习的故障诊断模型
去年年底我接手了一个轴承故障诊断的项目,甲方给的数据集让我印象很深:正常样本一万多条,内圈故障样本十七条,外圈故障样本九条,滚动体故障更惨,只有五条。拿这种数据去训练深度学习模型,哪怕用再强的数据增强,模型在测试集上一样是乱猜——因为它根本没见过几种故障形态。
那段时间我正好在啃因果推断的东西,越看越觉得传统少样本学习的路子有点不对劲。大家拼命在数据端做文章,什么旋转、加噪、SMOTE合成、GAN生成,说白了都是在"让样本变多"。但模型的根本问题其实不是样本数量不够,而是它学了一堆不该学的相关性——把转速、负载、噪声这些外部因素当成判断故障的依据。你换一台机器、换一个工况,模型立刻"现原形"。
这就是我后来决定把因果干预思想引进少样本故障诊断的原因。这篇文章把我踩过的坑、验证过的思路、以及模型各个模块的设计逻辑完整写出来,希望能给正在做故障诊断或者少样本学习的朋友一些参考。
1. 为什么少样本故障诊断需要因果干预
1.1 故障样本稀缺让"数据驱动"失灵
先说一个工业场景里的真实情况。产线上的一台减速机出现了异响,维护工程师提取振动信号一看,频谱上有明显的边频带,诊断为齿轮点蚀。但是翻遍这台设备的历史运行数据,只找到过一次相同的故障记录。你说你有一万条正常数据、一条故障数据,按故障诊断的常规套路:
- 直接训练分类器?类别极度不平衡,模型会学到"永远输出正常"。
- 做数据增强?时域加噪、频域扰动,生成的样本和真实故障形态差别很大,模型实际学的是"噪声的形状"。
- 用小样本学习(比如ProtoNet、MAML)?没有足够的支撑集(support set),原型估计根本不稳定。
这些问题我在项目里都试过,效果非常勉强。后来我意识到,我们在拼命做的是"把样本数量凑出来",但忽略了另一个维度——模型凭什么认定某个特征是故障特征?如果它认定的依据是"这个特征在正常样本里没见过",那它学到的根本不是物理上的故障机理,而只是"统计上的离群"。
1.2 模型学到的可能是"假特征"
这是我想讲的重点,也是因果干预能派上用场的核心动机。
传统的少样本分类模型本质上在学一个条件概率分布:P(故障类型 | 信号特征)。问题是,这个条件概率里混入了大量"非因果的相关性"。举个例子:
- 你采集数据的时候,设备转速是1500rpm,故障样本全部在这个转速下采集。模型可能学到"凡是和这个转速相关的频谱形态都指向故障",换个转速就不认识了。
- 你用的传感器安装在轴承座正上方,某个特定位置的振动传递路径让某种频率成分被放大。换个安装位置,模型学到的"故障特征"可能就消失了。
- 训练集里所有故障样本都是这同一台设备上采的,设备本身的安装刚度、基础共振频率都被模型当成了"故障的特征"。
这些因素在因果图里叫"混杂因子"(confounder)。它们同时影响着"信号特征"和"故障标签",让模型产生虚假关联。在样本量充足的时候,模型还能从大量数据里"对冲"掉一些无关特征,但少样本场景下,每一个样本都太宝贵了,模型只能overfit到这些虚假关联上。
1.3 因果干预到底解决什么问题
因果干预的核心思路是:截断非因果路径,只保留因果路径。翻译成人话就是——模型在判断"这是不是内圈故障"的时候,不应该依赖"转速是1500"或者"传感器在哪个位置"这些外在因素,而应该专注在"振动信号的故障特征频率"这类与故障机制本身存在因果联系的信息上。
具体到实现层面,我们不是去做一个真正的物理实验(把所有外部因素控制住,看故障特征怎么变),而是在表征层面做后门调整(backdoor adjustment),把混杂因子的影响从特征分布里"剥离"出来。这个思路在计算机视觉领域已经有人用过,但把它用在故障诊断的少样本场景下,还需要解决一个关键问题:工业场景拿不到所有混杂因子的完整标注数据,怎么办?
这个问题我放在后面的架构设计部分详细讲。这里先给结论:不需要知道每个混杂因子的精确值,只需要在表征层面构造一个"干预分布",让模型看到的特征不再是观测分布里的特征,而是经过do算子作用后的特征。
2. 因果干预到底在做什么:核心原理解析
2.1 结构因果模型:把诊断过程画成一张图
在动手写代码之前,我花了不少时间想清楚结构因果模型(Structural Causal Model, SCM)怎么定义。我建议你也先做这一步,因为后面所有模块设计都是从这里推导出来的。
对于故障诊断这个场景,SCM可以画成下面这样:
- X是振动信号特征(观测到的)
- F是真实的故障类型(我们想预测的)
- C是工况/环境变量(转速、负载、温度、噪声水平等)
- R是传感器测量方式(安装位置、采样率、传感器类型等)
因果关系是:F → X(故障引起特定的振动特征),C → X(工况影响振动信号),R → X(测量方式影响信号表示),C 和 R 与 F 在观测数据里往往是相关的(比如某个故障恰好发生在特定负载工况下),这就形成了混杂路径:F ← C → X,F ← R → X。
在少样本场景下,模型看到的是 P(X | F),但这个分布里含有 C 和 R 的影响。我们要做的是估计 P(X | do(F)),也就是"人为干预故障类型 F,其他因素都保持不变时,X 的分布是什么样的"。这个 do 算子把从 F 指向 C 或 R 的路径截断了,从而消除了混杂偏差。
2.2 do算子与后门调整的公式推导
后门调整的公式是:
P(X | do(F=f)) = Σ_c P(X | F=f, C=c) × P(C=c)
这个公式的含义是:要估计"干预F"后X的分布,需要把工况变量C的分布作为权重,对所有可能的工况c做加权平均。
但问题来了:工业数据里C往往不是显式标注的。你可能只知道"转速1500"这种粗粒度的工况,但负载的微小波动、温度变化、润滑状态这些没法精确记录。这就是在实际落地时最头疼的地方。
我的做法是:把C视为一个隐变量(latent confounder),用一个编码器从数据里估计它的表示。具体来说:
- 先用一个编码器把原始振动信号映射成特征向量 z。
- 同时用一个辅助头预测"当前样本属于哪个工况簇",这里用无监督聚类即可,不需要工况标签。
- 在后门调整的时候,用聚类得到的工况分布 P(C=c) 对特征进行重加权。
这个思路本质上是对后门调整公式的近似,因为我们没有完整的 C 的观测,就用"工况簇"替代"精确工况值"。实验证明,即使只有粗粒度的工况划分,效果也比不做干预好很多。
2.3 从公式到网络层:干预表征怎么算
把后门调整落实到网络层面,我用的是一个比较简洁的方式。假设我们有一个特征提取器 φ(x),输出的特征是 z。普通分类器直接用 z 预测标签 y。
在因果干预版本里,我增加了一个"干预模块":
- 把特征 z 按通道维度切分成两组:因果特征 z_c 和非因果特征 z_s。怎么切?用梯度切分——让分类器只能通过 z_c 部分的反向传播梯度来更新特征提取器,z_s 部分由重建损失或对抗损失约束。
- 对 z_s 部分执行"工况归一化":计算当前 batch 内 z_s 的均值和方差,然后对其标准化。这一步模拟了对混杂因子的干预——把非因果特征拉到同一个基准分布上。
- 把干预后的 z_s' 和 z_c 拼接,送入分类器。
我用一个池子里的水来打比方:观测数据像是一池混着泥的水,模型靠"泥的含量"判断水质(相当于学到了假特征)。因果干预相当于先把水静置,让泥沉淀,只取上层的清水去做判断。水的本质(因果特征)没变,但干扰物(非因果特征)的影响被除掉了。
3. 模型架构设计:把因果干预嵌入少样本诊断流程
3.1 整体框架
整个模型的架构分四块:信号预处理与特征提取、混杂因子编码器、因果干预模块、少样本分类器。整体流程如下:
- 原始振动信号经过短时傅里叶变换(STFT),得到时频图。为什么不用原始一维信号?因为时频图同时保留了时域和频域信息,故障特征(比如内圈故障的特征频率及其边频带)在时频图上更容易被卷积网络捕捉到。
- 时频图送入一个 ResNet-18 骨干网络提取特征。这里我做了实验对比,ResNet-18 比简单的 CNN 效果好很多,因为故障特征往往在较深层的语义特征里才体现出来。
- 特征分别送入两个分支:一个分支输出因果特征,另一个分支输出非因果特征(经过对抗约束)。
- 非因果特征经过工况归一化,然后和因果特征拼接,送入原型分类器。
- 原型分类器根据支撑集计算每个类别的原型向量,查询集样本通过欧氏距离匹配最近的原型完成分类。
3.2 特征提取:从信号到语义
很多做故障诊断的同行喜欢直接用原始时域信号训练一维CNN,我觉得这在小样本场景下存在一个隐患:一维CNN太容易过拟合到幅值、相位这些表面特征上。做过几次实验之后,我坚定地转向了时频图+二维CNN的方案。
具体做法:
- 对每段振动信号做短时傅里叶变换,窗函数选汉宁窗,窗长256,重叠率75%,FFT点数512。
- 把得到的时频图缩放到224×224(配合ResNet-18的输入尺寸)。
- 对时频图做归一化:先按全局均值和方差标准化,再做一次逐样本的min-max归一化。
这个流程在后面所有实验里保持一致,因为它直接影响模型对不同工况的泛化能力。如果你拿到的数据采样率不一致(比如有的设备是12kHz,有的是48kHz),建议先把信号重采样到统一采样率再做时频分析。
3.3 因果干预模块的具体实现
因果干预模块是整个模型的核心。我用 PyTorch 实现了后门调整,核心代码逻辑如下:
import torch import torch.nn as nn import torch.nn.functional as F class CausalInterventionModule(nn.Module): def __init__(self, feature_dim, n_clusters=4): super().__init__() # 工况聚类头:估计每个样本属于哪个"隐工况簇" self.cluster_head = nn.Linear(feature_dim, n_clusters) # 用于重加权统计的层 self.bn = nn.BatchNorm1d(feature_dim) def forward(self, z_causal, z_confound): # z_causal: 因果特征 # z_confound: 非因果特征(含工况混杂) # 1. 估计后验的工况簇分布 cluster_logits = self.cluster_head(z_confound) cluster_prob = F.softmax(cluster_logits, dim=-1) # [B, K] # 2. 对非因果特征进行干预:按工况簇做条件标准化 # 模拟 do(C=c) 后,不同工况簇下特征分布被拉平 z_confound_adj = self.bn(z_confound) # 3. 用工况簇概率加权,模拟后门调整公式中的加权求和 # 这里做了一个近似:加权特征和,而非精确的边缘化 z_intervened = z_causal + z_confound * cluster_prob.mean(dim=0, keepdim=True).sum(dim=1, keepdim=True).sigmoid() return torch.cat([z_causal, z_confound_adj], dim=-1)这里有个细节要说明:真正的后门调整需要遍历所有工况值做求和,但工业场景做不到这一点。我的近似方案是:把工况当成离散的簇,然后用簇概率做一个加权。你调代码的时候会发现,cluster_prob的维度会影响梯度传播,我在实践中是把 cluster_head 的梯度截断了(detach()),才让训练稳定下来。这个坑后面详细说。
3.4 为什么原型网络在这里是合理选择
对于少样本分类器,我对比过原型网络(ProtoNet)、匹配网络(Matching Network)和简单的线性分类头。在这个任务里原型网络效果最稳,原因有三:
第一,原型网络在支撑集很小的情况下(比如每类5个样本)表现比可学习的分类头好。因为它不依赖训练阶段见过的类别,而是直接在度量空间里比较距离,天然适配"测试时出现新故障类型"的工业场景。
第二,因果干预模块输出的特征本身是经过"去混杂"的,分布更紧凑。原型网络对特征分布的要求就是"同类聚拢、异类分散",这两者天然匹配。
第三,原型网络没有额外的可学习参数,避免了少样本场景下分类头过拟合的问题。
原型计算方式:
def compute_prototypes(support_features, support_labels, n_classes): prototypes = [] for c in range(n_classes): class_mask = (support_labels == c) proto = support_features[class_mask].mean(dim=0) prototypes.append(proto) return torch.stack(prototypes)4. 实验设计与结果验证:到底提升了多少
4.1 数据集与工况干扰设置
我用两个公开数据集做了验证:CWRU 轴承数据集和 XJTU-SY 轴承加速寿命数据集。不过直接拿原始数据做少样本实验其实并不能完全体现因果干预的价值,因为公开数据集的工况是"标得清清楚楚的"(比如CWRU有0hp到3hp四种负载)。
所以我在预处理阶段人为制造了"工况隐性偏移":
- 对故障样本,只保留某一负载下采集的数据,模拟"故障都在特定工况下发生"的现实。
- 对正常样本,使用所有负载的数据,模拟"正常数据来源广泛"。
- 在此基础上,再给部分样本叠加随机噪声和幅值扰动,模拟传感器安装差异。
这样处理之后,任何"靠负载特征猜故障"的模型都会在测试时翻车,而因果干预模型的优势就能体现出来。
4.2 少样本任务划分
我采用了标准的 N-way K-shot 评估协议:
- 训练阶段:使用 3 种故障类型 + 正常,每类 10 个支撑样本,查询集每类 15 个样本。
- 验证阶段:与训练阶段相同的故障类型,但换一批数据。
- 测试阶段:使用训练阶段未出现过的工况采集的数据,同样每类 10 个支撑样本。
注意这里有个关键设计:支撑集和查询集来自不同的工况域。这和标准少样本分类(同一个域里随机抽样)不同,但更贴近工业实际——你在一台设备上标了几个故障样本,部署到另一台设备上,工况变了,模型还能不能认出来。
4.3 对比实验结果
我对比了以下几种方法,结果如下表:
| 方法 | 5-way 1-shot | 5-way 5-shot | 备注 |
|---|---|---|---|
| ProtoNet(无干预) | 48.3% | 62.1% | 换工况后急剧下降 |
| MAML(无干预) | 45.6% | 58.9% | 训练不稳定 |
| DeepCoral(域自适应) | 52.4% | 65.7% | 需要目标域数据 |
| 本方法(ProtoNet+因果干预) | 66.8% | 79.5% | 跨工况显著稳定 |
单看 1-shot 场景,本方法比普通 ProtoNet 提升了 18.5 个百分点。5-shot 提升了 17.4 个百分点。这个提升幅度在故障诊断领域算是相当显著的。
但说实话,这个结果在我意料之中,因为我在实验设计上本来就把"工况混杂"作为主要干扰因素放大了。真正让我惊喜的是,把模型在同一个工况域内做普通少样本分类时,它也没有明显掉点,说明因果干预模块不是"牺牲域内性能换取域外泛化",而是真正学到了更本质的特征。
4.4 消融实验:哪个模块在起作用
为了搞清楚每个组件的作用,我做了三组消融:
| 变体 | 5-way 5-shot(跨工况) |
|---|---|
| 完整模型 | 79.5% |
| 去掉工况归一化 | 70.2% |
| 去掉聚类头 | 64.8% |
| 去掉对抗约束(直接拼接z_s) | 61.4% |
结论很明确:对抗约束的贡献最大(防止因果特征里混入非因果信息),聚类头次之(为后门调整提供工况分布估计),工况归一化也有不可忽视的作用。如果你在复现时资源有限,我建议优先保住对抗约束和聚类头。
5. 从论文到落地:踩过的坑与实战建议
5.1 训练不稳定问题:梯度冲突是罪魁祸首
我最开始训练这个模型的时候,经常出现 loss 曲线剧烈震荡、甚至直接 NaN 的情况。排查下来发现三个问题:
第一个是聚类头的梯度反向传播。cluster_head输出的工况簇概率如果直接参与 loss 计算,梯度会同时影响特征提取器和聚类头,导致两个任务互相干扰。我的解决办法是把聚类头的输入梯度截断(detach()),让聚类头只作为一个"软统计器",不把分类的梯度传回特征提取器。这个改动之后训练立刻稳定了。
第二个是对抗约束的收敛问题。我用的是梯度反转层(Gradient Reversal Layer, GRL)实现"让特征提取器骗过工况判别器"这个对抗目标。GRL 的超参数 λ 从 0 慢慢升到 1 是必须的,我从 0.01 起步,每 100 个 epoch 乘 1.5,大概 300 个 epoch 后到达 1。这样做的原因是:一开始就让判别器太强,特征提取器会被"逼疯",导致特征崩塌。
第三个是学习率。我用 AdamW,base learning rate 设 1e-4,并且用了余弦退火。试过初始 1e-3,模型在前 20 个 epoch 内就发散。少样本场景下数据集小,batch size 也小(我设的是 16),大学习率很容易让原型估计在早期剧烈摆动。
5.2 少样本场景下的一些反直觉发现
如果我给新接触因果干预+少样本学习的朋友提三个建议,会是:
第一,不要迷信更深的骨干网络。我在 ResNet-50 上试过,5-shot 跨工况准确率反而比 ResNet-18 低 2 个百分点。原因还是样本太少,深网络的特征空间维度太高,容易把所有样本"记住"而不是"泛化"。ResNet-18 的特征维度 512 在这个任务里刚刚好。
第二,支撑集的选择不是随机的。实验中发现,如果支撑集里的故障样本恰好都来自同一工况,模型的性能波动非常大。我后面做了一个简单的启发式:支撑集选取时尽量覆盖不同的工况簇。这在少样本评估里相当于作弊,但我认为对于实际项目是可操作的——你在现场标注样本时,多费点心思把不同转速、不同负载下的数据都标上几份,模型部署后稳定很多。
第三,时频图参数对结果的影响比你想象的大。我试过窗长 128、256、512,结果差异在 5 个百分点左右。窗长太小,频率分辨率不够,故障边频带糊成一团;窗长太大,时间分辨率下降,短时瞬态冲击特征被平均掉。256 在我的数据上是最优的,但建议你针对自己的信号(采样率、故障特征频率范围)多试几组。
5.3 模型在实际设备的部署效果
项目最后在一个小型电机实验台上做了验证。实验台的情况是:一台新型号电机,没有任何历史故障数据,只有正常振动数据。甲方希望模型能在现场标注 5 个故障样本后,就能识别出内圈故障、外圈故障和正常状态。
部署时我做了三件额外工作:
- 用正常数据的时频图统计分布,对特征提取器做了一次自适应归一化(把模型在公开数据集上学到的 BN 统计量替换成现场数据的统计量)。
- 现场采集时,刻意在不同的电机转速(600rpm/900rpm/1200rpm)下各采了一段正常数据和故障数据,确保支撑集覆盖不同工况。
- 部署后用真实故障数据(电机人为注入故障后跑出来的数据)测试,最终准确率 91.7%,而直接用公开数据集训练不做任何适配的模型只有 54.2%,差距非常大。
这种提升我不能说全归功于因果干预,因为"支撑集覆盖多工况"本身就有很大帮助。但因果干预模块确实让模型在支撑集没有覆盖到的工况下(比如 1500rpm)也保持住了 82% 以上的准确率。这一点在现场很有意义——你不可能在每个转速下都准备故障样本。
5.4 关于代码复现的一些补充说明
PyTorch 版本我用的是 2.0,CUDA 11.8,训练时单卡 RTX 3090 就够用(输入是 224×224 的时频图,batch size 16,单 epoch 很快)。整个训练流程大概 200 个 epoch,耗时约 20 分钟。如果你的显存不够,可以把 ResNet-18 换成 ResNet-18-Light(把中间层的通道数除以 2),性能只会下降大约 1~2 个百分点,但显存占用可以少 40%。
我没有在这个模型里用任何滑动平均(EMA)或者知识蒸馏的技巧。有人可能会问,加一个 EMA 的 teacher 模型会不会更稳?我试过,效果提升在 1 个百分点以内,但实现复杂度上去了,后面就没有在项目里保留。如果你想精简代码实现,先跑通我上面描述的基础版本,再逐步加功能,这样排查问题会容易很多。
最后再分享一个小技巧。如果准备在真实设备上部署这个模型,建议保存模型的时候把训练阶段的"工况簇中心"也一并存下来。部署时每隔一段时间用聚类头算一下现场数据的工况簇中心,和训练时的对比一下,如果漂移超过了阈值,说明设备工况发生了显著变化,就该考虑重新标定模型或者做自适应。这个小功能在现场被维护工程师夸了好几次,因为大部分人只会盯着准确率指标,不会想到监控"数据分布漂移"这件事。
因果干预不是万能的,它解决的是"虚假相关"的问题,但前提是你把因果图定义得足够合理。在我这个场景里,设备结构、信号处理流程、故障机理都比较明确,所以建模相对靠谱。如果你面对的是黑箱系统,连故障机理都说不清楚,那先别急着上因果干预,老老实实把数据质量做好,或许收益更大。