PyTorch实现注意力机制:少量样本故障诊断实战

PyTorch实现注意力机制:少量样本故障诊断实战 简介面向少量样本条件下的故障诊断场景这份PyTorch源码包实现了基于注意力机制的完整方案聚焦工业数据标注成本高、故障样本稀少的实际问题适合从事机械故障诊断、信号处理与深度学习交叉方向的研究者或工程师参考。资源共18个文件其中8个Python脚本覆盖一维信号注意力机制、AMSGradP优化器、1D-Meta-ACON激活函数、GAP全局池化、1D-Grad-CAM可视化以及AdaBN域自适应等关键模块10个MAT数据文件提供对应实验样本整体压缩包约10.89MB结构精简、脚本职责清晰数据与模型分离便于替换和扩展可直接运行调试。目前已有358人学习下载。除核心模型外还附带了数据保存、早停、标签平滑等工程脚本便于快速复现实验并迁移到自己的数据集上同时可作为小样本故障诊断研究的基准对理解注意力机制、域自适应与优化策略的组合使用具有实际参考价值。1. 少量样本故障诊断为什么需要注意力机制在产线上做设备状态监测的工程师大概都经历过这种处境设备正常运行大半年故障样本一只手数得过来用这些样本训练出来的故障诊断模型实验室里测试集准确率接近百分百到现场换一台设备就失效。这类问题通常被归为少量样本学习也是故障诊断从论文走向落地时最大的坎。轴承故障诊断、齿轮箱故障诊断这类任务核心是在振动信号里识别特定频带的冲击模式。传统卷积神经网络通过堆叠卷积核提取特征但卷积核的权重在整条信号上是全局共享的模型不会主动突出哪个频带对判别更有价值。注意力机制做的事情正好相反它动态计算每个通道、每个时间位置的重要性权重在样本量有限时引导模型把表达能力集中到最能区分故障类别的频带上同时抑制环境噪声和工况波动带来的干扰。下面用 PyTorch 把“注意力机制 少量样本故障诊断”的落地路径串起来先讨论选型再给一套可运行的模型和训练脚本最后落到注意力权重的提取与验证。适合正在做设备状态监测、故障诊断相关项目的工程师直接参考。2. 注意力机制选型SE通道注意力与CBAM在故障诊断中的取舍2.1 通道注意力与空间注意力在振动信号上的作用注意力机制在故障诊断中做的事情是让模型在特征提取过程中计算“哪些特征更重要”然后给它们分配更大的权重。这里的特征有两个维度通道维度和空间维度。通道维度上振动信号经过卷积层后每个卷积核输出一个特征图对应一种特征响应——有的通道对高频冲击敏感有的通道对旋转频率及其谐波敏感。通道注意力统计每个通道的全局信息计算一个重要性权重向量再逐通道乘以原始特征图。SESqueeze-and-Excitation模块是这类结构里最典型的实现先对特征图做全局平均池化再经过两层全连接和 Sigmoid 输出通道权重。空间维度上特征图的不同位置对应信号的不同时间段。空间注意力要回答的问题是“这段信号里哪几十个采样点才是故障冲击真正出现的位置”。CBAM 把通道注意力和空间注意力串接起来先做通道加权再做空间加权。它在图像分类里已经是标配处理一维振动信号时把池化和卷积相应改成 1D 版本即可。在少量样本故障诊断里注意力模块带来的归纳偏置与故障信号的物理规律是吻合的故障冲击在频带上集中、在时间上局部注意力机制恰恰是抓住这两点的最轻量手段。它不像单纯增加卷积层深度那样扩大假设空间而是把模型的拟合方向约束到与故障机理一致的特征上所以样本少时通常比同参数量的普通 CNN 更容易收敛到可泛化的解。2.2 三种注意力模块的参数量与适用场景对比故障诊断代码里最常见的注意力机制是 SE、CBAM 和多头自注意力三者的设计目标和计算代价差别很大。拿一段长度 1024 的轴承振动信号、卷积层输出 64 个通道来估算模块核心结构额外参数量约计算特点少量样本场景适用性SE全局平均池化 两层全连接约 4k64×64/16×2只做通道加权训练快高最稳妥CBAM1D 版通道注意力 一维空间注意力约 4k~8k同时关注通道与时间位置高推荐首选多头自注意力QKV 线性变换 注意力矩阵约 16k 以上注意力矩阵复杂度 O(L²)低样本少时容易过拟合多头自注意力的问题在序列长度上。输入长度 L1024 时注意力矩阵是 1024×1024也就是百万量级的元素这个自由度过大在样本只有几百条时几乎必然过拟合。SE 的参数量和计算量最小但它只在通道维做加权对故障冲击出现在哪个时间段没有建模能力。CBAM 是这两者的折中参数量比自注意力小一个量级同时覆盖了“哪个通道有用”和“哪段时间有用”两件事所以我一般把 CBAM 作为少量样本诊断模型的默认选择。还有一类坐标注意力CA在故障诊断里也偶尔见到它在通道注意力基础上把位置信息也编码进去对于需要同时感知通道和位置的任务有效但实现复杂度比 CBAM 高样本量不足时收益不明显。如果数据量不超过 2000 条不建议优先尝试。2.3 少量样本场景下的注意力机制设计要点注意力模块的细节参数对最终效果的影响比很多人大。下面这段 SE 通道注意力的定义是故障诊断代码里最常见的基础版本import torch import torch.nn as nn class ChannelAttention1d(nn.Module): def __init__(self, in_channels, reduction16): super().__init__() self.avg_pool nn.AdaptiveAvgPool1d(1) self.fc nn.Sequential( nn.Linear(in_channels, in_channels // reduction), nn.ReLU(inplaceTrue), nn.Linear(in_channels // reduction, in_channels), ) def forward(self, x): b, c, _ x.size() y self.avg_pool(x).view(b, c) weight torch.sigmoid(self.fc(y)).view(b, c, 1) return x * weight这里的reduction直接决定了注意力分支自身的参数量reduction越大中间层维度越小参数量越少。图像任务里习惯取 16但在少量样本场景下建议从 32 起步观察训练集与验证集准确率的差距后再调整。如果训练集已经收敛到接近百分百、验证集明显跟不上去说明注意力分支自身过拟合了把reduction调大如果两边都很低才考虑逐步调小恢复表达能力。空间注意力的实现也需要注意卷积核大小。CBAM 原文里空间注意力用 7×7 卷积核对应到一维信号就是 kernel_size7。少量样本条件下这个感受野偏大会把冲击位置前后的无关区段也卷进来我通常改成 5必要时降到 3。另外一个容易踩的坑是 BatchNorm 的位置注意力权重经过 Sigmoid 后与特征相乘这时不需要再接 BN否则 batch size 很小时统计量抖动反而破坏已经学好的通道比例。3. 用PyTorch搭建带CBAM的故障诊断模型从数据加载到训练3.1 振动信号的滑动窗口切分与数据集构建故障诊断原始数据一般是从传感器采集的长序列振动信号不能直接整段送进网络。常见做法是用滑动窗口把长信号切成固定长度的短样本每条样本对应一个标签。窗口长度通常取 1024 或 2048需要保证窗口内至少包含 2~3 个旋转周期的冲击序列具体由设备转速决定。切分代码如下import numpy as np from torch.utils.data import Dataset, DataLoader def sliding_windows(signal, window_size1024, stride512): windows [] for start in range(0, len(signal) - window_size 1, stride): windows.append(signal[start:start window_size]) return np.stack(windows) class FaultDataset(Dataset): def __init__(self, data, labels): self.data data.astype(np.float32) self.labels labels.astype(np.int64) def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx].reshape(1, -1), self.labels[idx]stride决定相邻窗口的重叠程度。取 512 时重叠一半数据量翻倍适合少量样本场景但要注意重叠窗口之间存在信息冗余验证集和训练集必须按窗口来源分组不能让同一条长信号切出的重叠窗口同时出现在两个集合里否则验证准确率会被高估。reshape(1, -1)把窗口变成单通道的一维信号对应nn.Conv1d的输入格式。3.2 一维CBAM模块与故障诊断网络结构CBAM 迁移到一维信号时核心改动是把空间注意力里的二维卷积换成nn.Conv1d。下面实现里空间注意力把通道维分别做均值池化和最大池化拼成两通道特征后经过一个一维卷积输出每个时间位置的权重import torch.nn.functional as F class SpatialAttention1d(nn.Module): def __init__(self, kernel_size5): super().__init__() self.conv nn.Conv1d(2, 1, kernel_size, paddingkernel_size // 2, biasFalse) self.sigmoid nn.Sigmoid() def forward(self, x): avg_out torch.mean(x, dim1, keepdimTrue) max_out, _ torch.max(x, dim1, keepdimTrue) attn torch.cat([avg_out, max_out], dim1) self.last_weight self.sigmoid(self.conv(attn)) return x * self.last_weight class CBAM1d(nn.Module): def __init__(self, in_channels, reduction32, kernel_size5): super().__init__() self.channel_attn ChannelAttention1d(in_channels, reduction) self.spatial_attn SpatialAttention1d(kernel_size) def forward(self, x): x self.channel_attn(x) x self.spatial_attn(x) return x实现里特意把last_weight保存下来后面做注意力可视化时不需要重新 forward 或挂 hook直接读取这个属性即可。均值池化保留整体能量背景最大池化突出冲击尖峰两者拼接能让空间注意力同时感知“背景强度”和“局部峰值”。特征提取网络用一个浅层一维 CNN控制参数量是少量样本场景的关键class FaultDiagnosisNet(nn.Module): def __init__(self, in_channels1, num_classes4): super().__init__() self.features nn.Sequential( nn.Conv1d(in_channels, 32, kernel_size8, stride2, padding4), nn.BatchNorm1d(32), nn.ReLU(inplaceTrue), nn.MaxPool1d(2), nn.Conv1d(32, 64, kernel_size5, padding2), nn.BatchNorm1d(64), nn.ReLU(inplaceTrue), ) self.cbam CBAM1d(in_channels64, reduction32, kernel_size5) self.classifier nn.Sequential( nn.AdaptiveAvgPool1d(1), nn.Flatten(), nn.Linear(64, num_classes), ) def forward(self, x): x self.features(x) x self.cbam(x) return self.classifier(x)输入 1024 点时第一层卷积输出 513 点MaxPool 后变成 256 点第二层卷积保持 256 点CBAM 在 256 点长度的特征图上计算空间权重。整个模型参数量约 8 万在少量样本下参数规模是可控的。第一层卷积核取 8是为了在第一个阶段就有足够大的感受野覆盖冲击响应后面的小卷积核负责局部细化。3.3 训练脚本编写与关键参数取值训练部分用标准的监督学习流程但有几个参数值得专门说明import torch.optim as optim device torch.device(cuda if torch.cuda.is_available() else cpu) model FaultDiagnosisNet(num_classes4).to(device) optimizer optim.AdamW(model.parameters(), lr1e-3, weight_decay1e-4) scheduler optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max50) criterion nn.CrossEntropyLoss() def train_one_epoch(model, loader, optimizer, criterion): model.train() total_loss 0.0 for x, y in loader: x, y x.to(device), y.to(device) optimizer.zero_grad() logits model(x) loss criterion(logits, y) loss.backward() optimizer.step() total_loss loss.item() return total_loss / len(loader) loader DataLoader(FaultDataset(train_data, train_labels), batch_size32, shuffleTrue) for epoch in range(50): loss train_one_epoch(model, loader, optimizer, criterion) scheduler.step()优化器选 AdamW 而不是 Adam因为 AdamW 把权重衰减从梯度动量中分离出来少量样本下对过拟合的抑制更干净。初始学习率 1e-3 对浅层 CNN 是合理起点如果 loss 在前 5 个 epoch 不下降降到 3e-4 再试。CosineAnnealing 配合 50 个 epoch 足够模型在几百条样本上收敛不需要训练上百轮。batch size 取 32 是折中太小则 BN 统计量不稳定太大则每个 epoch 更新次数过少。这里没有用预训练模型原因是故障诊断的输入是振动波形与图像预训练特征差异很大直接用 ImageNet 权重做迁移层反而引入无关先验。后面章节会讲少量样本下更有效的两种训练策略。4. 少量样本下的训练策略数据增强、Focal Loss 与早停4.1 数据增强时域抖动、幅值缩放与加噪样本量不够时数据增强是性价比最高的手段。故障诊断信号的数据增强必须遵守一个原则增强操作不能改变故障冲击的本质频率特征。随意对信号做时间伸缩会把冲击频率挪走模型学到的是错误的判别依据。下面三个增强函数是故障诊断代码里常见的组合def add_noise(signal, snr_db20): signal signal.astype(np.float32) sig_power np.mean(signal ** 2) noise_power sig_power / (10 ** (snr_db / 10)) noise np.random.normal(0, np.sqrt(noise_power), signal.shape).astype(np.float32) return signal noise def amplitude_scale(signal, scale_range(0.9, 1.1)): scale np.random.uniform(*scale_range) return signal * scale def time_shift(signal, max_shift50): shift np.random.randint(-max_shift, max_shift 1) return np.roll(signal, shift)加噪的snr_db是信噪比数值越小噪声越强。建议从 20dB 开始如果验证集仍然过拟合逐步降到 10dB。幅值缩放模拟的是负载波动0.9~1.1 的范围相当于正负 10% 的幅值变化超过这个范围会改变故障冲击与背景噪声的相对强度。时域平移的作用是消除窗口切分时冲击相位不一致带来的偏移敏感max_shift不能超过一个旋转周期的采样点数否则窗口内容被替换得太多。增强在训练时在线进行代码实现时把增强函数放在__getitem__里而不是预先存盘这样每个 epoch 看到的样本都是经过不同随机变换的版本。验证集不要做增强否则会掩盖真实识别能力。4.2 损失函数与标签平滑Focal Loss 缓解类别不均衡故障诊断数据集的类别分布往往不均衡正常样本远多于各类故障样本少数类只占十几条。普通的交叉熵损失会被多数类主导注意力机制也会偏向拟合样本量大的类别。Focal Loss 是处理这种情况的常用损失函数它在交叉熵基础上加了一个调制因子让模型把注意力集中在难分类的样本上class FocalLoss(nn.Module): def __init__(self, gamma2.0, alphaNone): super().__init__() self.gamma gamma self.alpha alpha def forward(self, logits, targets): ce F.cross_entropy(logits, targets, reductionnone) pt torch.exp(-ce) loss (1 - pt) ** self.gamma * ce if self.alpha is not None: alpha_t torch.tensor(self.alpha, devicelogits.device)[targets] loss alpha_t * loss return loss.mean()gamma是调制因子取 2.0 时一个样本的pt接近 1分类正确且置信度高(1 - 0.9)^2 0.01损失被压得很低pt接近 0分错时损失几乎不受影响。这样训练的重心自然转向故障样本。alpha是类别权重列表比如正常类给 0.2、稀有故障类给 0.8值在验证集上做一次粗调即可。另一件容易被忽略的事是标签平滑。少量样本下模型对训练标签的置信度过高输出层的 logit 会趋向极端。把交叉熵的目标标签从 1 换成 0.95其余类别分到 0.05 / (num_classes - 1)能显著缓解过拟合。PyTorch 的CrossEntropyLoss(label_smoothing0.1)直接支持这个参数Focal Loss 则需要在构造 target 时手动做平滑操作比较绕建议二选一不要同时堆叠。4.3 早停与模型选择不要只看训练准确率少量样本训练的另一个特点是模型在某个 epoch 后会突然从欠拟合跳到过拟合这个拐点往往只有几个 epoch 的间隔。每轮都保存验证集准确率最高的模型权重比硬性训练固定轮数可靠得多best_acc 0.0 patience 10 wait 0 for epoch in range(50): train_loss train_one_epoch(model, loader, optimizer, criterion) val_acc evaluate(model, val_loader, device) if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_model.pt) wait 0 else: wait 1 if wait patience: break这里的patience表示连续多少个 epoch 验证集准确率不提升就停止。evaluate函数在验证集上只做前向计算关闭梯度并统计类别平均准确率。注意验证集必须是同一个设备或同一工况下采样的数据如果验证集和训练集来自同一条长信号的相邻窗口信息重叠会让best_acc虚高早停也就失去了意义。判断模型是否真的有效要看验证曲线和训练曲线的分歧点。训练准确率一路接近百分百、验证准确率却停滞时优先调整的是增强强度和注意力模块的reduction参数而不是无脑加大模型。5. 提取注意力权重验证模型学到了什么5.1 把空间注意力权重暴露给外部读取CBAM 模块里保存的last_weight是训练结束后做可解释性分析的关键入口。它形状是(batch, 1, L)L 是最后的特征图长度对于本模型的 1024 点输入就是 256。推理时直接读取这个属性不用修改网络结构model.eval() x torch.from_numpy(signal).float().reshape(1, 1, -1).to(device) with torch.no_grad(): logits model(x) weight model.cbam.spatial_attn.last_weight # (1, 1, 256)sig需要是 float32从文件读入的 numpy 数组默认可能是 float64必须显式转换否则torch.from_numpy会报类型错误。model.eval()和torch.no_grad()在这里缺一不可前者关掉 Dropout 和 BN 的批统计更新后者关掉自动求导避免注意力权重被额外的前向计算污染。5.2 从特征长度映射回原始采样点256 点的权重只能对应到原始信号的大致区段要精确对齐时域波形需要用线性插值把权重上采样回 1024 点import torch.nn.functional as F weight_up F.interpolate(weight, sizex.size(-1), modelinear, align_cornersFalse) weight_np weight_up.squeeze().cpu().numpy() import matplotlib.pyplot as plt fig, axes plt.subplots(2, 1, figsize(12, 6), sharexTrue) axes[0].plot(signal, colortab:blue) axes[0].set_title(原始振动信号) axes[1].plot(weight_np, colortab:red) axes[1].set_title(空间注意力权重) plt.tight_layout() plt.savefig(attention_map.png, dpi150)modelinear只对单通道的 1D 插值有效多通道时要用modenearest或者逐通道处理。插值会抹平一些细节但用于判断注意力的峰值区间完全够用。保存成文件而不是在 Jupyter 里直接显示图片分辨率能拉高峰值位置的判断更清楚。5.3 判断模型有没有学到故障机理的三个检查点第一次画注意力图时重点看权重峰值是否落在经验冲击频率对应的时刻上。拿轴承外圈故障来说故障特征频率对应的冲击间隔是固定的如果注意力权重峰值间隔与这个频率吻合说明模型把判别依据放在了故障冲击上可靠性高。峰值落在随机位置时优先检查两件事滑动窗口的起点是否对齐了冲击序列随机拾取的窗口可能让冲击落在窗口边缘训练时用的增强里有没有把信号做大幅时间伸缩破坏频率结构。第三点要习惯性验证类别区分度分别对正常样本和故障样本画注意力图然后把两者叠在同一张图里观察。正常的注意力分布通常较平缓故障样本的注意力会集中在局部区间。如果两类样本的注意力分布没有明显差异即使分类准确率很高也要怀疑模型是否借用了转速、负载等工况信息来取巧。把小批量样本的注意力统计量做成散点图两类的分布重叠越少特征分离度越好这个验证步骤在样本量少时比多跑几个网络结构更有价值。本文还有配套的精品资源点击获取