直面噪声问题!深度残差收缩网络的Python编程复现
在复杂工业环境下,旋转机械(如轴承、齿轮箱)的故障振动信号常被背景噪声掩盖。传统的深度学习模型(如CNN、ResNet)在处理高信噪比数据时具有较好的特征提取能力,但在强噪声干扰下,其诊断精度往往受到影响。
1. 深度残差收缩网络的概况
2020年,《IEEE Transactions on Industrial Informatics》上的“Deep residual shrinkage networks for fault diagnosis”这篇论文提出了深度残差收缩网络(Deep Residual Shrinkage Network, DRSN)。该网络是残差网络(ResNet)的一种改进版本,其核心思想是在残差网络内嵌入可学习的软阈值收缩模块。
该模型通过引入注意力机制,使网络能够在特征图层面自适应地确定阈值,从而对噪声相关的特征响应进行压制,保留与故障特征相关的分量。本文复现了其中的 DRSN-CW(Channel-wise thresholds)变体,即实现逐通道阈值的自适应调节。

2. 核心原理:软阈值化与自适应学习
2.1 软阈值函数 (Soft Thresholding)
软阈值函数是信号去噪处理中的经典算子。与ReLU激活函数(仅截断负值)不同,软阈值函数对正负特征进行对称压缩。其数学表达式为:y=sgn(x)⋅max(|x|-τ,0),其中τ为阈值。当特征绝对值小于τ时,输出为0;当大于τ时,输出向0收缩。这一机制可用于剔除特征图中幅值较低的噪声分量。
2.2 自动学习阈值的子网络
DRSN的特征在于其阈值τ是由网络自主学习得到的,无需人工干预。在RSBU-CW(残差收缩单元)中,阈值的生成逻辑如下:
①全局特征汇聚:对输入特征图进行全局平均池化(GAP),获取各通道的平均响应强度。
②注意力机制映射:通过两层全连接(FC)层处理池化后的向量,第二层使用 Sigmoid 激活函数输出一个范围在(0, 1)之间的缩放因子α。
③动态阈值计算:阈值由缩放因子与特征绝对值的均值相乘得到,即τ=α×average(|x|)。这种设计确保了阈值始终处于合理的范围内,避免将所有特征全部软阈值化为零。
3. 实验设置与复现逻辑
3.1 数据集与预处理
实验基于CWRU(凯斯西储大学)轴承数据集。
①故障分类:包含正常状态、内圈故障、外圈故障、滚动体故障,结合不同损伤直径共计10类样本。
②信号切片:使用长度为1024的滑动窗口对原始加速度序列进行非重叠采样。
③数据增强:训练过程中引入了随机循环平移(Random Roll)、局部冲击干扰和动态噪声混合,模拟现场工况的相位偏移和瞬态冲击。

3.2 强噪声环境模拟
为了测试模型的稳健性,在验证集和测试集中加入了-8dB的高斯白噪声(AWGN)。在此信噪比下,噪声功率显著高于信号功率。

3.3 训练策略
① 优化器:Adam(初始学习率1e-3)。
② 正则化:采用L2权重衰减(系数1e-4)防止过拟合。
③ 回调机制:集成学习率衰减(ReduceLROnPlateau)和早停法(EarlyStopping)以优化训练过程。
4. 复现代码结构
以下为基于TensorFlow实现的DRSN-CW核心架构逻辑:
"""
本程序基于 TensorFlow 框架复现了深度残差收缩网络(DRSN),用于旋转机械(如轴承)
的振动信号故障诊断。代码具体实现了论文中提出的 "DRSN-CW" (Channel-wise thresholds)
变体,即具有逐通道阈值的残差收缩单元。
主要功能:
1. 数据处理:CWRU 数据集的加载、切片及加噪处理(高斯白噪声)。
2. 模型构建:实现 RSBU-CW 模块,通过注意力机制自适应学习软阈值,剔除噪声特征。
3. 训练评估:端到端的模型训练、验证及在低信噪比环境下的鲁棒性测试。
[参考文献]
Zhao M, Zhong S, Fu X, Tang B, Pecht M. Deep residual shrinkage networks for fault
diagnosis. IEEE Transactions on Industrial Informatics, 2020, 16(7): 4681-4690.
===============================================================================
"""
import os
import sys
import logging
import numpy as np
import scipy.io as sio
import tensorflow as tf
from tensorflow.keras import layers, Model, regularizers
from sklearn.model_selection import train_test_split as split_data
# =============================================================================
# 运行环境配置
# =============================================================================
logging.basicConfig(level=logging.INFO, format='[%(asctime)s] %(levelname)s: %(message)s')
class GPUConfig:
"""
运行环境配置管理器。
负责底层计算资源的分配、依赖库状态校验以及计算图优化策略的设定。
"""
@staticmethod
def init_tf():
"""
配置深度学习计算后端:
1. 屏蔽 TensorFlow 冗余的调试信息。
2. 针对物理 GPU 设备启用显存动态增长(Dynamic Memory Allocation),防止资源预占溢出。
"""
os.environ['TF_CPP_MIN_LOG_LEVEL'] = '2'
physical_gpu_list = tf.config.list_physical_devices('GPU')
if physical_gpu_list:
try:
for gpu_device in physical_gpu_list:
tf.config.experimental.set_memory_growth(gpu_device, True)
logging.info("成功检测到 GPU 设备 {0} 台,显存动态增长模式已激活。".format(len(physical_gpu_list)))
except RuntimeError as hardware_error:
logging.warning("后端配置发生异常: %s", hardware_error)
else:
logging.info("计算环境未检测到可用 GPU,任务将回退至通用 CPU 执行。")
# 执行后台环境初始化
GPUConfig.init_tf()
# =============================================================================
# 数据获取与时序处理模块
# =============================================================================
class CWRULoader:
"""
振动信号处理器类。
负责 CWRU 原始加速度序列的解析、非重叠滑动窗口采样以及特征矩阵的构建。
"""
def __init__(self, dataset_root, window_size=1024):
"""
构造函数。
:param dataset_root: 原始 .mat 文件存储的根路径。
:param window_size: 单个时序样本覆盖的时间步长。
"""
self.base_directory = os.path.abspath(dataset_root)
self.sample_length = window_size
self.sampling_interval = window_size
def _parse_mat_content(self, target_file):
"""
从 MATLAB 容器文件中检索驱动端加速度计数据(DE_time)。
"""
try:
storage = sio.loadmat(target_file)
for identifier in storage.keys():
if 'DE_time' in identifier:
return storage[identifier].flatten()
except Exception as parse_error:
logging.debug("文件 %s 解析失败: %s", target_file, parse_error)
return None
return None
def load_data(self, category_dictionary):
"""
根据指定的类别索引与文件对应关系,构建完整的训练/测试数据集。
:param category_dictionary: 映射字典 {标签编号: [相关文件名]}。
:return: (特征向量张量, 标签索引张量)。
"""
feature_collection, label_collection = [], []
is_data_found = False
for class_idx, name_list in category_dictionary.items():
for filename in name_list:
full_path = os.path.join(self.base_directory, "{0}.mat".format(filename))
if not os.path.exists(full_path):
continue
vibration_series = self._parse_mat_content(full_path)
if vibration_series is None:
continue
is_data_found = True
# 执行固定窗口步长的时序切片
for pointer in range(0, len(vibration_series) - self.sample_length + 1, self.sampling_interval):
sub_sequence = vibration_series[pointer : pointer + self.sample_length]
feature_collection.append(sub_sequence)
label_collection.append(class_idx)
if not is_data_found:
raise FileNotFoundError("在指定路径下未能定位到符合规则的故障诊断数据。")
return np.array(feature_collection, dtype='float32'), np.array(label_collection, dtype='int32')
def add_awgn(signal_input, snr_value):
"""
基于信噪比(SNR)控制的高斯白噪声注入算法。
该函数通过计算信号平均功率,反向推导所需的噪声标准差,实现对信号质量的精确降级。
"""
signal_input = np.array(signal_input)
random_engine = np.random.default_rng()
# 确定目标信噪比强度
target_snr = snr_value if not isinstance(snr_value, (list, tuple)) \
else random_engine.uniform(snr_value[0], snr_value[1])
# 噪声功率计算逻辑:P_noise = P_signal / 10^(SNR/10)
signal_power = np.mean(np.square(signal_input), axis=1, keepdims=True)
noise_variance = signal_power / (10 ** (target_snr / 10.0))
noise_component = random_engine.normal(0, np.sqrt(noise_variance), signal_input.shape)
return (signal_input + noise_component).astype('float32')
# =============================================================================
# 深度残差收缩网络 (DRSN) 架构组件
# =============================================================================
class SoftThresholdOperator(layers.Layer):
"""
深度残差收缩网络核心组件:非线性软阈值化层。
算子定义:y = sign(x) * max(|x| - τ, 0),其中 τ 为学习到的正数阈值。
"""
def __init__(self, **kwargs):
super(SoftThresholdOperator, self).__init__(**kwargs)
def call(self, inputs):
"""
前向运算:执行逐元素的特征收缩。
inputs 包含 [特征映射图, 阈值向量]。
"""
x_conv, tau = inputs
# 调整阈值张量维度以适配特征图通道
expanded_tau = tf.expand_dims(tau, axis=1)
return tf.sign(x_conv) * tf.maximum(tf.abs(x_conv) - expanded_tau, 0.0)
class RSBU_CW(layers.Layer):
"""
深度残差收缩网络基本单元 (RSBU-CW)。
集成了全局平均池化、注意力权重学习以及自适应阈值生成的残差块。
"""
def __init__(self, filters, kernel_size, strides=1, **kwargs):
super(RSBU_CW, self).__init__(**kwargs)
self.num_kernels = filters
self.step_size = strides
self.width = kernel_size
self.weight_decay = regularizers.l2(1e-4)
# 恒等路径
self.shortcut = None
# 卷积变换主路径
self.bn_alpha = layers.BatchNormalization()
self.relu_alpha = layers.Activation('relu')
self.conv_alpha = layers.Conv1D(filters, kernel_size, strides=strides, padding='same',
kernel_initializer='he_normal', kernel_regularizer=self.weight_decay)
self.bn_beta = layers.BatchNormalization()
self.relu_beta = layers.Activation('relu')
self.conv_beta = layers.Conv1D(filters, kernel_size, strides=1, padding='same',
kernel_initializer='he_normal', kernel_regularizer=self.weight_decay)
# 子网络:计算逐通道的阈值
self.gap = layers.GlobalAveragePooling1D()
self.fc1 = layers.Dense(filters, kernel_initializer='he_normal')
self.bn_gamma = layers.BatchNormalization()
self.relu_gamma = layers.Activation('relu')
self.fc2 = layers.Dense(filters, activation='sigmoid')
self.threshold_op = SoftThresholdOperator()
def build(self, input_dim):
"""
根据输入维度判断是否需要对恒等路径进行线性变换。
"""
if self.step_size != 1 or input_dim[-1] != self.num_kernels:
self.shortcut = tf.keras.Sequential([
layers.Conv1D(self.num_kernels, 1, strides=self.step_size, padding='same'),
])
super(RSBU_CW, self).build(input_dim)
def call(self, layer_inputs):
"""
深度残差收缩逻辑流:
1. 卷积提取初步特征。
2. 全局池化压缩特征统计量。
3. 全连接层感知通道权重并映射为动态阈值。
4. 执行阈值去噪并累加残差信号。
"""
identity = layer_inputs
if self.shortcut:
identity = self.shortcut(layer_inputs)
# 两次卷积变换
x_conv = self.bn_alpha(layer_inputs)
x_conv = self.relu_alpha(x_conv)
x_conv = self.conv_alpha(x_conv)
x_conv = self.bn_beta(x_conv)
x_conv = self.relu_beta(x_conv)
x_conv = self.conv_beta(x_conv)
# 特征绝对值感知
x_abs = tf.abs(x_conv)
abs_mean = self.gap(x_abs)
# 阈值计算子路径
z = self.fc1(abs_mean)
z = self.bn_gamma(z)
z = self.relu_gamma(z)
alpha = self.fc2(z)
# 生成动态阈值:权重 * 绝对值均值
tau = tf.multiply(alpha, abs_mean)
# 收缩处理与残差融合
denoised_output = self.threshold_op([x_conv, tau])
return layers.Add()([denoised_output, identity])
class DRSN_CW(Model):
"""
深度残差收缩网络完整模型架构。
输入:一维振动信号。
输出:故障分类概率分布。
"""
def __init__(self, num_classes):
super(DRSN_CW, self).__init__(name="Bearing_Fault_DRSN")
# 初始特征感知层
self.conv1 = layers.Conv1D(32, 15, strides=2, padding='same', kernel_initializer='he_normal')
self.bn1 = layers.BatchNormalization()
self.relu1 = layers.Activation('relu')
# 顺序堆叠收缩残差单元,逐层抽象高级特征
self.rsbu_blocks = [
RSBU_CW(32, 5, strides=2),
RSBU_CW(32, 5, strides=1),
RSBU_CW(64, 5, strides=2),
RSBU_CW(64, 5, strides=1),
RSBU_CW(128, 5, strides=2),
RSBU_CW(128, 5, strides=1)
]
# 全局池化与分类头
self.post_norm = layers.BatchNormalization()
self.post_relu = layers.Activation('relu')
self.gap_layer = layers.GlobalAveragePooling1D()
self.classifier = layers.Dense(num_classes, activation='softmax')
def call(self, network_input):
"""
端到端前向推理过程。
"""
x = self.conv1(network_input)
x = self.bn1(x)
x = self.relu1(x)
for block in self.rsbu_blocks:
x = block(x)
x = self.post_norm(x)
x = self.post_relu(x)
x = self.gap_layer(x)
return self.classifier(x)
# =============================================================================
# 训练与性能评估工作流
# =============================================================================
def train_and_test(dataset_path, seq_len=1024):
"""
执行故障诊断生命周期管理:加载、训练、在线增强及抗噪评估。
"""
# 构建故障类别映射体系
label_map = {
0: ['Normal_0', 'Normal_1', 'Normal_2', 'Normal_3'],
1: ['IR007_0', 'IR007_1', 'IR007_2', 'IR007_3'],
2: ['IR014_0', 'IR014_1', 'IR014_2', 'IR014_3'],
3: ['IR021_0', 'IR021_1', 'IR021_2', 'IR021_3'],
4: ['B007_0', 'B007_1', 'B007_2', 'B007_3'],
5: ['B014_0', 'B014_1', 'B014_2', 'B014_3'],
6: ['B021_0', 'B021_1', 'B021_2', 'B021_3'],
7: ['OR007@6_0', 'OR007@6_1', 'OR007@6_2', 'OR007@6_3'],
8: ['OR014@6_0', 'OR014@6_1', 'OR014@6_2', 'OR014@6_3'],
9: ['OR021@6_0', 'OR021@6_1', 'OR021@6_2', 'OR021@6_3']
}
data_engine = CWRULoader(dataset_root=dataset_path, window_size=seq_len)
try:
signals, labels = data_engine.load_data(label_map)
except Exception as data_err:
logging.error("数据集生成失败: %s", data_err)
return
# 训练、验证、测试集的切分
train_x_pre, temp_x, train_y_pre, temp_y = split_data(
signals, labels, test_size=0.3, random_state=42
)
val_x_pre, test_x_pre, val_y_pre, test_y_pre = split_data(
temp_x, temp_y, test_size=0.5, random_state=42
)
# 特征标准化处理:基于训练集分布
mu, sigma = np.mean(train_x_pre), np.std(train_x_pre)
def normalize(obs):
return ((obs - mu) / sigma).reshape(-1, seq_len, 1)
train_set_x = normalize(train_x_pre)
val_set_x = normalize(val_x_pre)
test_set_x = normalize(test_x_pre)
# 标签转码
num_classes = len(label_map)
train_set_y = tf.keras.utils.to_categorical(train_y_pre, num_classes).astype('float32')
val_set_y = tf.keras.utils.to_categorical(val_y_pre, num_classes).astype('float32')
test_set_y = tf.keras.utils.to_categorical(test_y_pre, num_classes).astype('float32')
# 测试环境模拟:注入 -8dB 强背景噪声
val_x_awgn = add_awgn(val_set_x, snr_value=-8)
test_x_awgn = add_awgn(test_set_x, snr_value=-8)
def augment_batch(feat_batch, label_batch):
"""
在线数据增强:通过随机相位平移、稀疏脉冲叠加及动态噪声混合提升模型的泛化能力。
"""
rand_gen = np.random.default_rng()
augmented_x = feat_batch.copy()
batch_n, steps_n, _ = augmented_x.shape
# 随机循环移位
for sample_idx in range(batch_n):
offset = rand_gen.integers(0, steps_n)
augmented_x[sample_idx, :, 0] = np.roll(augmented_x[sample_idx, :, 0], offset)
# 瞬态冲击干扰
if rand_gen.random() > 0.9:
for sample_idx in range(batch_n):
if rand_gen.random() > 0.5:
num_spikes = rand_gen.integers(1, 3)
positions = rand_gen.integers(0, steps_n, num_spikes)
spike_mag = np.std(augmented_x[sample_idx]) * rand_gen.uniform(1.5, 2.5)
augmented_x[sample_idx, positions, 0] += spike_mag * rand_gen.choice([-1, 1], size=num_spikes)
# 混合信噪比噪声注入(50% 概率触发)
if rand_gen.random() > 0.5:
augmented_x = add_awgn(augmented_x, snr_value=(-8, 8))
return augmented_x.astype(np.float32), label_batch.astype(np.float32)
def _tensor_spec_binding(f_tensor, l_tensor):
f_tensor.set_shape([None, seq_len, 1])
l_tensor.set_shape([None, num_classes])
return f_tensor, l_tensor
# 构建高性能 tf.data 数据流
training_pipeline = tf.data.Dataset.from_tensor_slices((train_set_x.astype('float32'), train_set_y))
training_pipeline = training_pipeline.shuffle(len(train_set_x)).batch(64)
training_pipeline = training_pipeline.map(
lambda x, y: tf.numpy_function(augment_batch, [x, y], [tf.float32, tf.float32]),
num_parallel_calls=tf.data.AUTOTUNE
).map(_tensor_spec_binding).prefetch(tf.data.AUTOTUNE)
# 编译深度残差收缩网络模型
model_instance = DRSN_CW(num_classes=num_classes)
model_instance.compile(
optimizer=tf.keras.optimizers.Adam(learning_rate=1e-3),
loss=tf.keras.losses.CategoricalCrossentropy(label_smoothing=0.0),
metrics=['accuracy']
)
logging.info("基于深度残差收缩网络的诊断系统初始化就绪。分类规模: %d, 窗口跨度: %d", num_classes, seq_len)
# 训练监控回调
optimization_callbacks = [
tf.keras.callbacks.ReduceLROnPlateau(monitor='val_loss', factor=0.5, patience=7, min_lr=1e-6, verbose=1),
tf.keras.callbacks.EarlyStopping(monitor='val_loss', patience=20, restore_best_weights=True)
]
# 启动神经网络迭代
model_instance.fit(
training_pipeline,
epochs=100,
validation_data=(val_x_awgn, val_set_y),
callbacks=optimization_callbacks,
verbose=2
)
# 抗噪性能验证
final_loss, final_acc = model_instance.evaluate(test_x_awgn, test_set_y, verbose=0)
print("\n" + "="*50)
print("评估报告: 深度残差收缩网络 (DRSN)")
print("测试环境信噪比: -8dB SNR")
print("鲁棒识别准确率: {0:.2f}%".format(final_acc * 100))
print("="*50)
# =============================================================================
# 主程序入口
# =============================================================================
if __name__ == "__main__":
# 配置默认的数据搜索域
DATA_PATH = os.path.join(os.getcwd(), 'data_path')
if not os.path.exists(DATA_PATH):
logging.warning("默认数据存放路径无效: %s", DATA_PATH)
user_input_path = input("请手动指定 CWRU 数据集 (.mat) 的存储路径: ").strip()
if user_input_path:
DATA_PATH = user_input_path
else:
logging.critical("路径输入缺失,系统终止。")
sys.exit(0)
# 启动故障诊断流水线
train_and_test(DATA_PATH, seq_len=1024)
5. 实验结果分析
根据复现程序的运行日志(如图所示),在叠加了-8dB强噪声的测试环境下:
①训练表现:经过约100个Epoch的迭代,模型在训练集上的准确率趋于99%以上。
②测试结果:在-8dB SNR条件下,测试集的故障识别准确率达到了95.16%。
③稳定性:验证集损失值在训练后期保持在0.28左右,学习率经过数次衰减(从 1e-3降至1.56e-5)后,模型收敛状态稳定。

论文原文:
更多推荐
所有评论(0)