如何降低ECG R峰检测中KNN分类器的高误报率?
QRS检测算法复现中的误报问题排查
我使用Python 3.11.7(Jupyter Notebook环境)复现一篇QRS检测论文的算法:预处理采用Pan-Tompkins 5-12Hz带通滤波、梯度计算与梯度曲线提取,之后用K近邻分类器(K=3,欧氏距离)检测ECG信号中的QRS复合波。但在MIT/BIH数据集的100号记录上训练并测试时,出现大量误报,测试结果如下:
2273 reference annotations, 2824 test annotations True Positives (matched samples): 2273 False Positives (unmatched test samples): 551 False Negatives (unmatched reference samples): 0
训练阶段代码
import wfdb from wfdb import processing import matplotlib.pyplot as plt import numpy as np import scipy.signal as sg from sklearn.neighbors import KNeighborsClassifier from sklearn.model_selection import cross_val_predict def calculate_gradient(signal): """ Calculate the gradient of a signal using central finite differences. Parameters: signal (numpy.ndarray): Input signal. Returns: numpy.ndarray: Gradient of the signal. """ gradient = np.zeros_like(signal) n = len(signal) for i in range(1, n - 1): gradient[i] = (signal[i + 1] - signal[i - 1]) / 2.0 gradient[0] = (signal[1] - signal[0]) gradient[n - 1] = (signal[n - 1] - signal[n - 2]) gradient /= np.max(gradient) return gradient def bandpass_filter(ecg, fs): Wn = 12*2/fs N = 3 a, b = sg.butter(N, Wn, btype='lowpass') ecg_l = sg.filtfilt(a, b, ecg) ecg_l = ecg_l/np.max(np.abs(ecg_l)) Wn = 5*2/fs N = 3 a, b = sg.butter(N, Wn, btype='highpass') ecg_h = sg.filtfilt(a, b, ecg_l, padlen=3*(max(len(a), len(b))-1)) ecg_h = ecg_h/np.max(np.abs(ecg_h)) return ecg_h y1_0 = bandpass_filter(data, Fs) y2_0 = calculate_gradient(y1_0) m = len(y2_0) # Number of training instances n = 2 # Number of features feature_vector = np.zeros((m, n)) label_vector = np.zeros(m) feature_vector[:, 0] = y2_0 # Setting labels for QRS/non-QRS regions for sample_index in ann_ref: window_start = max(0, sample_index - 25) window_end = min(len(data), sample_index + 25) label_vector[window_start:window_end] = 1 neigh = KNeighborsClassifier(n_neighbors=3, p=2, metric='minkowski') neigh.fit(feature_vector, label_vector)
测试阶段代码
y1_0 = bandpass_filter(data0, Fs) y2_0 = calculate_gradient(y1_0) m = len(y2_0) # Number of training instances n = 2 # Number of features X = np.zeros((m, n)) y = np.zeros(m) X[:, 0] = y2_0 for sample_index in ann_ref: window_start = max(0, sample_index - 25) window_end = min(len(data0), sample_index + 25) y[window_start:window_end] = 1 predicted_labels0 = cross_val_predict(neigh, X[:, 0].reshape(-1, 1), y, cv=5) # Performing fivefold cross-validation # Function to calculate the average pulse duration def calculate_average_pulse_duration(predicted_labels): peak_durations = np.diff(np.where(predicted_labels == 1)[0]) return np.mean(peak_durations) # Function to detect QRS-complex based on the average pulse duration def detect_QRS_complex(train_of_ones, average_pulse_duration): QRS_indices = [] for i, train_duration in enumerate(train_of_ones): if train_duration > 3 * average_pulse_duration: QRS_indices.append(train_duration) return QRS_indices train_of_ones_0 = [] for label, index in zip(predicted_labels0, range(len(predicted_labels0))): if label == 1: train_of_ones_0.append(index) average_pulse_duration_0 = calculate_average_pulse_duration(predicted_labels0) QRS_indices_0 = detect_QRS_complex(train_of_ones_0, average_pulse_duration_0) QRS_indices_0 = np.array(QRS_indices_0, dtype=int) samples = 100 index0 = [] for i in QRS_indices_0: start_index = max(0, i - samples) end_index = min(len(y2_0), i + samples + 1) signal_within_margin = y2_0[start_index:end_index] peaks = np.max(signal_within_margin) peaks_idx = np.argmax(signal_within_margin) index0.append(peaks_idx + start_index) index0 = np.unique(index0) index0 = np.sort(index0) index0 = np.array(index0) comparitor = processing.compare_annotations(ann_ref_indices, index0, int(0.1*Fs)) comparitor.print_summary()
注:我通过后缀0(如y2_0、train_of_ones_0、index0)指代ECG记录的第一导联,因算法需适配MLII和VI导联。请问我哪里出现了问题?
问题排查点
- 特征向量构建错误:训练阶段定义了2个特征,但仅填充了梯度值,第二个特征全为0,相当于只用单一特征做分类,特征区分度严重不足,是误报的核心原因之一。需补全第二个特征(比如梯度的平方、相邻梯度差值等,对照论文确定)。
- 标签窗口设置不合理:用
样本索引±25标注QRS区域,若窗口过大,会把非QRS区域误标为正样本,导致分类器学习错误模式;需对照论文确认QRS标签的窗口范围。 - 测试阶段交叉验证使用错误:
cross_val_predict是用于训练集内部交叉验证的函数,不能用来预测测试集。训练好模型后应直接调用neigh.predict(X)得到测试集预测标签,当前逻辑会打乱测试流程,引入额外误判。 - QRS后处理逻辑完全错误:
calculate_average_pulse_duration函数计算的是连续正样本索引的间隔,而非QRS波的持续时间,逻辑完全偏离需求。detect_QRS_complex函数遍历正样本索引列表时,错误将索引值与平均间隔做比较,筛选出的索引毫无意义,直接导致后续峰值检测引入大量误报。需重新设计后处理逻辑:先将连续正样本合并为区间,计算每个区间的持续时间,再筛选符合QRS持续时间范围的区间,最后提取区间内的峰值作为QRS位置。
- 滤波实现不符论文要求:Pan-Tompkins带通滤波应直接设计5-12Hz的带通滤波器,而非先低通再高通串联;且每次滤波后做归一化会丢失信号相对幅值信息,削弱QRS与非QRS区域的梯度差异。
- 梯度归一化不当:用
np.max(gradient)做归一化会拉伸不同信号段的梯度幅值到同一范围,丢失原始信号的幅值差异,而QRS波的梯度幅值本应远大于非QRS区域,这种归一化会降低特征区分度,应去掉该步骤或改用全局归一化(基于整个信号的最大梯度)。
内容的提问来源于stack exchange,提问作者karimnh
相关产品推荐
相关产品推荐

