You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何降低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导联。请问我哪里出现了问题?


问题排查点

  1. 特征向量构建错误:训练阶段定义了2个特征,但仅填充了梯度值,第二个特征全为0,相当于只用单一特征做分类,特征区分度严重不足,是误报的核心原因之一。需补全第二个特征(比如梯度的平方、相邻梯度差值等,对照论文确定)。
  2. 标签窗口设置不合理:用样本索引±25标注QRS区域,若窗口过大,会把非QRS区域误标为正样本,导致分类器学习错误模式;需对照论文确认QRS标签的窗口范围。
  3. 测试阶段交叉验证使用错误:cross_val_predict是用于训练集内部交叉验证的函数,不能用来预测测试集。训练好模型后应直接调用neigh.predict(X)得到测试集预测标签,当前逻辑会打乱测试流程,引入额外误判。
  4. QRS后处理逻辑完全错误:
    • calculate_average_pulse_duration函数计算的是连续正样本索引的间隔,而非QRS波的持续时间,逻辑完全偏离需求。
    • detect_QRS_complex函数遍历正样本索引列表时,错误将索引值与平均间隔做比较,筛选出的索引毫无意义,直接导致后续峰值检测引入大量误报。需重新设计后处理逻辑:先将连续正样本合并为区间,计算每个区间的持续时间,再筛选符合QRS持续时间范围的区间,最后提取区间内的峰值作为QRS位置。
  5. 滤波实现不符论文要求:Pan-Tompkins带通滤波应直接设计5-12Hz的带通滤波器,而非先低通再高通串联;且每次滤波后做归一化会丢失信号相对幅值信息,削弱QRS与非QRS区域的梯度差异。
  6. 梯度归一化不当:用np.max(gradient)做归一化会拉伸不同信号段的梯度幅值到同一范围,丢失原始信号的幅值差异,而QRS波的梯度幅值本应远大于非QRS区域,这种归一化会降低特征区分度,应去掉该步骤或改用全局归一化(基于整个信号的最大梯度)。

内容的提问来源于stack exchange,提问作者karimnh

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.27 23:54:59