简介:本资源是一套面向生物医学工程、信号处理及机器学习初学者的ECG心电图分类实践项目,聚焦于心脏病理信号识别这一典型医疗AI任务。压缩包共825个文件,涵盖Python(20个.py)与MATLAB(11个.m)双语言实现的完整分类流程,包括数据预处理、特征提取(如RR间期、QRS波形参数)、SVM/随机森林等模型训练及评估代码;同时包含大量WFDB标准格式ECG数据文件(60个.hea、10个.dat、2个.mat)及配套工具(如gqrs、rdann、wfdb2mat等),便于直接加载MIT-BIH等公开数据库。资源大小为6.25MB,结构紧凑且模块化清晰,适合课程设计、毕业设计或Kaggle类竞赛入门实践。目前已有444人学习下载,读者可直接复现端到端分类 pipeline,掌握ECG信号去噪、波形检测、时频特征建模及跨平台算法对比等核心能力。
1. ECG分类不是调个库就完事:从原始信号到临床可解释标签的完整链路
你下载了ecg_classification-master.zip,解压后看到一堆.mat文件、main.py和README.md,但运行python main.py却卡在ImportError: No module named 'biosppy'或 MATLAB 报错Undefined function 'filtfilt'——这不是环境配置失败,而是 ECG 分类本身存在三重断层:信号预处理不统一、特征提取无医学依据、分类器输出不可解释。这个项目标题里藏着的「ECG分类」,本质是心电信号时序建模问题,不是图像分类的简单迁移。它要求你同时理解心电生理(P波、QRS复合波、T波的时域形态与节律关系)、数字信号处理(带通滤波、基线漂移校正、R峰检测)和机器学习(时序特征 vs. 深度学习端到端建模的取舍)。适合刚接触生物医学信号的 Python 工程师、需要复现论文结果的研究生,以及正在把 ECG 算法集成进嵌入式设备的固件工程师。本文不讲“如何安装 Python”,而是聚焦于:为什么同一段 ECG 数据,在 MATLAB 中用sgolayfilt滤波后送入 SVM,在 Python 中用scipy.signal.butter滤波再喂给 XGBoost,结果差异可能超过 15% 准确率?答案藏在采样率对齐、R-R 间期归一化、以及 QRS 波群截取窗口的毫秒级偏移里。
2. 用 Python + SciPy + sklearn 复现 ECG 分类最小闭环:从 .mat 加载到混淆矩阵输出
ECG 分类的起点不是模型,而是信号质量。ecg_classification-master.zip中的.mat文件通常包含data(原始电压序列)、fs(采样率)、labels(类别索引)三个字段。MATLAB 用户习惯直接load('data.mat'),但 Python 需要scipy.io.loadmat并处理结构体嵌套。常见错误是忽略 MATLAB 的列优先存储导致数组转置,或未将uint16原始数据转换为float64进行滤波运算。
2.1 加载与验证原始 ECG 信号:确认采样率与标签对齐
import scipy.io as sio import numpy as np import matplotlib.pyplot as plt # 加载 .mat 文件(注意:MATLAB v7.3+ 使用 HDF5,需 h5py;此处按传统 v7.0 处理) mat_data = sio.loadmat('data_001.mat') ecg_signal = mat_data['data'].flatten() # 强制展平为一维,避免 shape=(N,1) 导致后续滤波失败 fs = float(mat_data['fs'][0, 0]) # 提取标量采样率,避免 array([[500.0]]) labels = mat_data['labels'].flatten() # 验证信号长度与标签数量是否匹配(关键!) if len(ecg_signal) != len(labels): # 常见情况:labels 是事件标记(如 R 波位置),非逐点标签 # 此时需按滑动窗切分,例如每 256 点为一个样本 window_size = 256 n_samples = len(ecg_signal) // window_size ecg_windows = ecg_signal[:n_samples * window_size].reshape(-1, window_size) # labels 需同步截断或插值 labels = labels[:n_samples] print(f"已按 {window_size} 点窗口切分,得到 {n_samples} 个样本") else: print(f"原始信号长度 {len(ecg_signal)} 与标签数 {len(labels)} 一致") print(f"采样率: {fs} Hz, 信号类型: {ecg_signal.dtype}")提示:
scipy.io.loadmat默认将 MATLAB 结构体转为numpy.ndarray,但内部字段名可能被_替换(如data变成data_)。建议先print(mat_data.keys())查看真实键名。若.mat是 v7.3 格式(文件头含HDF5),必须改用h5py库加载,否则报NotImplementedError。
2.2 医学可信的预处理四步法:带通滤波、基线校正、R峰检测、QRS截取
ECG 分类性能对预处理极度敏感。直接使用scipy.signal.butter设计巴特沃斯滤波器,参数必须符合 AHA(美国心脏协会)标准:0.5–40 Hz 带通。低于 0.5 Hz 无法消除呼吸基线漂移,高于 40 Hz 则丢失 QRS 上升支高频信息。
from scipy.signal import butter, filtfilt, find_peaks from scipy.interpolate import interp1d def preprocess_ecg(signal, fs=500): # 1. 带通滤波:0.5-40 Hz,二阶巴特沃斯,零相位滤波(避免延迟) nyq = fs / 2 b, a = butter(N=2, Wn=[0.5/nyq, 40.0/nyq], btype='band') filtered = filtfilt(b, a, signal) # 2. 基线漂移校正:三次样条插值拟合基线(非简单移动平均) # 找出 R 波位置作为锚点(粗略估计) peaks, _ = find_peaks(filtered, height=np.mean(filtered)+0.5*np.std(filtered), distance=int(0.6*fs)) if len(peaks) < 3: # R 波太少,退化为中位数滤波 baseline = np.median(filtered) else: # 在 R 波间插入基线点(每 2 个 R 波中点) baseline_points_x = [] baseline_points_y = [] for i in range(len(peaks)-1): mid_idx = (peaks[i] + peaks[i+1]) // 2 baseline_points_x.append(mid_idx) baseline_points_y.append(np.median(filtered[mid_idx-20:mid_idx+20])) # 三次样条插值 f = interp1d(baseline_points_x, baseline_points_y, kind='cubic', fill_value="extrapolate") x_full = np.arange(len(filtered)) baseline = f(x_full) corrected = filtered - baseline # 3. R 峰检测(用于后续窗口截取) r_peaks, _ = find_peaks(corrected, height=np.percentile(corrected, 95), distance=int(0.4*fs)) # 4. 截取 QRS 波群:以 R 峰为中心,取 [-120ms, +180ms](临床标准) qrs_windows = [] for r in r_peaks: start = max(0, r - int(0.12 * fs)) end = min(len(corrected), r + int(0.18 * fs)) if end - start >= int(0.3 * fs): # 确保窗口足够长 qrs_windows.append(corrected[start:end]) return np.array(qrs_windows), r_peaks # 执行预处理 qrs_list, r_positions = preprocess_ecg(ecg_signal, fs=fs) print(f"检测到 {len(r_positions)} 个 R 波,生成 {len(qrs_list)} 个 QRS 窗口")参数说明:
distance=int(0.4*fs)表示 R 波最小间隔为 0.4 秒(对应心率 150 bpm),防止误检 T 波;height=np.percentile(..., 95)动态设定阈值,适应不同信噪比;[-120ms, +180ms]是 AHA 推荐的 QRS 宽度范围,覆盖绝大多数正常与室性早搏形态。
2.3 特征工程:时域、频域、非线性指标的临床可解释组合
深度学习流行后,很多人忽略手工特征的价值。但在资源受限设备(如可穿戴 ECG 芯片)或小样本场景(<1000 条记录),XGBoost + 手工特征仍显著优于 CNN。本项目需提取三类特征:
| 特征类型 | 具体指标 | 计算方式 | 临床意义 |
|---|---|---|---|
| 时域 | QRS 宽度、R 波振幅、PR 间期 | np.argmax(qrs_window)定位 R 峰,左右找谷点 | 宽度 >120ms 提示束支传导阻滞 |
| 频域 | 0–15Hz 能量占比、主频 | np.abs(np.fft.rfft(qrs_window)) | 心肌缺血时高频成分衰减 |
| 非线性 | 样本熵(SampEn)、Higuchi 维数 | nolds.sampen(qrs_window) | 反映心电复杂度,房颤患者显著降低 |
from nolds import sampen from scipy.fft import rfft def extract_features(qrs_window, fs=500): features = {} # 时域特征 r_idx = np.argmax(qrs_window) # QRS 宽度(FWHM:半高全宽) half_max = np.max(qrs_window) / 2 left = np.where(qrs_window[:r_idx] < half_max)[0] right = np.where(qrs_window[r_idx:] < half_max)[0] if len(left) > 0 and len(right) > 0: width_ms = (r_idx - left[-1] + right[0]) * 1000 / fs else: width_ms = 0 features['qrs_width_ms'] = width_ms features['r_amplitude'] = qrs_window[r_idx] # 频域特征:0-15Hz 能量占比 freqs = np.fft.rfftfreq(len(qrs_window), d=1/fs) fft_mag = np.abs(rfft(qrs_window)) mask = (freqs >= 0) & (freqs <= 15) features['energy_0_15hz_ratio'] = np.sum(fft_mag[mask]**2) / np.sum(fft_mag**2) # 非线性特征 features['sampen'] = sampen(qrs_window, emb_dim=2, tolerance=0.1*np.std(qrs_window)) return list(features.values()) # 提取所有 QRS 窗口的特征 X_features = np.array([extract_features(q, fs=fs) for q in qrs_list]) y_labels = labels[:len(X_features)] # 对齐标签 print(f"特征矩阵形状: {X_features.shape} (样本数 × 特征数)") print(f"特征名称: ['qrs_width_ms', 'r_amplitude', 'energy_0_15hz_ratio', 'sampen']")注意:
nolds库需pip install nolds,其sampen实现严格遵循 Richman & Moorman 2000 年论文定义。若sampen报错ValueError: emb_dim must be less than len(data)-1,说明 QRS 窗口过短(<50 点),需检查preprocess_ecg中的截取逻辑。
3. MATLAB 与 Python 特征一致性验证:用相同算法跑通两个平台
当项目标题同时出现ECG分类、python、matlab,意味着你需要跨平台复现。MATLAB 的signal工具箱和 Python 的scipy.signal在滤波器设计上存在数值差异:MATLAB 默认使用butter的zpk(零极点增益)形式,而 SciPy 默认ba(分子分母系数)形式,导致filtfilt输出有微小偏差。验证方法不是比对浮点数,而是比对R峰检测位置和QRS宽度统计分布。
3.1 MATLAB 端:导出预处理中间结果供 Python 比对
在 MATLAB 中运行以下脚本,将关键中间变量保存为.npz格式(Python 可读):
% matlab_preprocess.m load('data_001.mat'); fs = double(fs); % 1. 带通滤波(AHA标准) [b,a] = butter(2, [0.5 40]/(fs/2), 'bandpass'); filtered = filtfilt(b,a,data(:)); % 2. 基线校正(三次样条) [~,r_peaks] = findpeaks(filtered, 'MinPeakHeight', mean(filtered)+0.5*std(filtered), ... 'MinPeakDistance', round(0.4*fs)); baseline_x = zeros(length(r_peaks)-1,1); baseline_y = zeros(length(r_peaks)-1,1); for i=1:length(r_peaks)-1 mid_idx = floor((r_peaks(i)+r_peaks(i+1))/2); baseline_x(i) = mid_idx; baseline_y(i) = median(filtered(max(1,mid_idx-20):min(end,mid_idx+20))); end f = spline(baseline_x, baseline_y); baseline = ppval(f, 1:length(filtered)); corrected = filtered - baseline; % 3. QRS 截取 qrs_list = {}; for r=1:length(r_peaks) start = max(1, r_peaks(r) - round(0.12*fs)); end_idx = min(length(corrected), r_peaks(r) + round(0.18*fs)); if end_idx - start >= round(0.3*fs) qrs_list{end+1} = corrected(start:end_idx); end end % 导出为 .npz(需 Python 的 numpy 支持) % 将 cell 数组转为矩阵(补零对齐) max_len = max(cellfun(@length, qrs_list)); qrs_matrix = zeros(length(qrs_list), max_len); for i=1:length(qrs_list) qrs_matrix(i,1:length(qrs_list{i})) = qrs_list{i}; end npzwrite('matlab_qrs.npz', struct('qrs_matrix', qrs_matrix, 'r_peaks', r_peaks));3.2 Python 端:加载 MATLAB 输出并比对特征分布
import numpy as np import matplotlib.pyplot as plt # 加载 MATLAB 导出的 QRS 矩阵 matlab_data = np.load('matlab_qrs.npz') matlab_qrs = matlab_data['qrs_matrix'] matlab_r_peaks = matlab_data['r_peaks'] # 用 Python 方法处理同一段原始信号(确保 fs 相同) python_qrs, python_r_peaks = preprocess_ecg(ecg_signal, fs=fs) # 比对 R 峰位置误差(单位:毫秒) r_error_ms = np.abs(python_r_peaks[:len(matlab_r_peaks)] - matlab_r_peaks[:len(python_r_peaks)]) * 1000 / fs print(f"R 峰定位平均误差: {np.mean(r_error_ms):.2f} ± {np.std(r_error_ms):.2f} ms") # 比对 QRS 宽度分布(直方图叠加) python_widths = [] for q in python_qrs: r_idx = np.argmax(q) half_max = np.max(q) / 2 left = np.where(q[:r_idx] < half_max)[0] right = np.where(q[r_idx:] < half_max)[0] if len(left) > 0 and len(right) > 0: python_widths.append((r_idx - left[-1] + right[0]) * 1000 / fs) matlab_widths = [] for i in range(matlab_qrs.shape[0]): q = matlab_qrs[i, :] q = q[q != 0] # 去除补零 if len(q) == 0: continue r_idx = np.argmax(q) half_max = np.max(q) / 2 left = np.where(q[:r_idx] < half_max)[0] right = np.where(q[r_idx:] < half_max)[0] if len(left) > 0 and len(right) > 0: matlab_widths.append((r_idx - left[-1] + right[0]) * 1000 / fs) plt.hist(python_widths, alpha=0.5, label='Python', bins=20) plt.hist(matlab_widths, alpha=0.5, label='MATLAB', bins=20) plt.xlabel('QRS Width (ms)') plt.ylabel('Count') plt.legend() plt.title('QRS Width Distribution Comparison') plt.show()关键结论:若 R 峰误差 > 5ms 或 QRS 宽度分布 Kolmogorov-Smirnov 检验 p-value < 0.01,则说明两平台预处理流程存在系统性偏差。此时应统一使用 MATLAB 的
sgolayfilt(Savitzky-Golay 滤波)替代 Python 的butter,因其在保留 QRS 形态尖锐度上更优——这正是标题中matlab与python并列的深层原因:算法选型需服从临床需求,而非框架便利性。
4. XGBoost 二分类模型构建与评估:避开准确率陷阱的 5 项核心指标
ECG 分类常面临严重类别不平衡(如正常窦性心律 vs. 室性早搏)。ecg_classification-master.zip中的labels若为[0,1],需先检查np.bincount(labels)。若正负样本比 > 5:1,直接用accuracy_score会掩盖模型失效——99% 准确率可能只是把所有样本判为多数类。
4.1 类别平衡与模型训练:强制指定 scale_pos_weight
from sklearn.model_selection import train_test_split from sklearn.metrics import classification_report, confusion_matrix, roc_auc_score import xgboost as xgb # 检查类别分布 print("类别分布:", np.bincount(y_labels)) # 划分训练/测试集(stratify 保持比例) X_train, X_test, y_train, y_test = train_test_split( X_features, y_labels, test_size=0.3, random_state=42, stratify=y_labels ) # 计算 scale_pos_weight:负样本数 / 正样本数(XGBoost 内置平衡) pos_count = np.sum(y_train == 1) neg_count = np.sum(y_train == 0) scale_pos_weight = neg_count / pos_count if pos_count > 0 else 1 # XGBoost 参数(针对 ECG 特征优化) params = { 'objective': 'binary:logistic', 'eval_metric': 'logloss', 'scale_pos_weight': scale_pos_weight, # 关键! 'max_depth': 6, 'learning_rate': 0.1, 'subsample': 0.8, 'colsample_bytree': 0.8, 'gamma': 0.1, 'reg_alpha': 0.01, 'seed': 42 } # 训练 dtrain = xgb.DMatrix(X_train, label=y_train) dtest = xgb.DMatrix(X_test, label=y_test) model = xgb.train(params, dtrain, num_boost_round=100, evals=[(dtest, 'test')], early_stopping_rounds=10) # 预测概率 y_pred_proba = model.predict(dtest) y_pred = (y_pred_proba > 0.5).astype(int)4.2 分类评估:必须报告的 5 项指标及其临床含义
仅报告accuracy是危险的。ECG 分类需关注:
| 指标 | 计算公式 | 临床意义 | 本例典型值 |
|---|---|---|---|
| 敏感度(Recall) | TP/(TP+FN) | 漏诊率:多少真实异常被漏掉? | >95%(避免漏诊心梗) |
| 特异度(Specificity) | TN/(TN+FP) | 误诊率:多少正常人被误判为异常? | >90%(减少患者焦虑) |
| F1-score | 2×Precision×Recall/(Precision+Recall) | 精准与召回的调和平均 | >0.92 |
| ROC-AUC | ROC 曲线下面积 | 模型区分能力(0.5=随机,1.0=完美) | >0.95 |
| Youden 指数 | Sensitivity + Specificity - 1 | 综合判别效能最大值 | >0.85 |
from sklearn.metrics import recall_score, precision_score, roc_auc_score, roc_curve # 计算核心指标 sensitivity = recall_score(y_test, y_pred) # 同 TP/(TP+FN) specificity = recall_score(y_test, y_pred, pos_label=0) # TN/(TN+FP) f1 = 2 * (precision_score(y_test, y_pred) * sensitivity) / (precision_score(y_test, y_pred) + sensitivity) auc = roc_auc_score(y_test, y_pred_proba) print(f"Sensitivity (Recall): {sensitivity:.4f}") print(f"Specificity: {specificity:.4f}") print(f"F1-score: {f1:.4f}") print(f"ROC-AUC: {auc:.4f}") print(f"Youden Index: {sensitivity + specificity - 1:.4f}") # 绘制 ROC 曲线 fpr, tpr, _ = roc_curve(y_test, y_pred_proba) plt.plot(fpr, tpr, label=f'ROC Curve (AUC = {auc:.4f})') plt.plot([0,1], [0,1], 'k--', label='Random Classifier') plt.xlabel('False Positive Rate') plt.ylabel('True Positive Rate') plt.title('ROC Curve for ECG Classification') plt.legend() plt.show()提示:
recall_score(y_test, y_pred, pos_label=0)计算的是负类(正常)的召回率,即特异度。不要用classification_report的support列判断样本量——它显示的是测试集中的数量,而非原始数据集分布。
5. 模型可解释性落地:用 SHAP 值定位 ECG 分类的关键生理特征
医生不会信任一个“黑箱”模型。标题中的ecg_classification要求输出不仅准确,还要能回答:“为什么判定这个片段是室性早搏?” SHAP(SHapley Additive exPlanations)是当前最可靠的局部可解释方法,它能给出每个特征对单次预测的贡献值。
5.1 计算 SHAP 值并可视化单个样本的决策逻辑
import shap # 创建 explainer(使用 TreeExplainer 适配 XGBoost) explainer = shap.TreeExplainer(model) shap_values = explainer.shap_values(X_test) # 选择一个测试样本(例如第一个异常样本) idx = np.where(y_test == 1)[0][0] shap.plots.waterfall(explainer.expected_value, shap_values[idx], feature_names=['qrs_width_ms', 'r_amplitude', 'energy_0_15hz_ratio', 'sampen'], max_display=4)5.2 特征重要性全局分析:识别驱动分类的核心生理指标
# 全局特征重要性(基于 |SHAP| 值均值) feature_names = ['qrs_width_ms', 'r_amplitude', 'energy_0_15hz_ratio', 'sampen'] shap_abs_mean = np.abs(shap_values).mean(axis=0) feature_importance = pd.DataFrame({ 'feature': feature_names, 'shap_mean_abs': shap_abs_mean }).sort_values('shap_mean_abs', ascending=False) print("全局 SHAP 重要性(降序):") print(feature_importance) # 可视化 plt.figure(figsize=(8,4)) plt.barh(feature_importance['feature'], feature_importance['shap_mean_abs']) plt.xlabel('|SHAP Value| Mean') plt.title('Global Feature Importance (ECG Classification)') plt.gca().invert_yaxis() plt.show()临床解读示例:若
qrs_width_ms的 SHAP 值为 +0.8(正向推动异常分类),且该样本实际宽度为 160ms(>120ms),则模型依据 AHA 标准判定为束支传导阻滞;若sampen值为 -0.6(负向抑制异常分类),说明该片段复杂度高,倾向正常节律。这种解释可直接写入医疗 AI 系统的审核报告,满足 FDA 的可追溯性要求。
本文还有配套的精品资源,点击获取