ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

临床级心脏病预测:数据清洗、特征工程与SHAP可解释性实战

临床级心脏病预测:数据清洗、特征工程与SHAP可解释性实战 简介本资源是一套面向机器学习初学者与实践者的完整心脏病预测分析实战项目聚焦健康指标数据建模与多算法对比验证适用于课程设计、Kaggle式入门项目及医疗AI兴趣拓展。压缩包共20个文件含18个可直接运行的Python脚本覆盖数据清洗、探索性分析、特征工程、12种以上主流分类模型实现与调优、1个核心CSV数据集BRFSS2015健康调查原始数据及1份说明文档总大小仅2.4MB轻量易下载。已有82人学习下载体现其在入门级医疗数据分析场景中的实用热度。读者可获得从原始数据加载到模型部署全流程的18个模块化代码——包括SMOTE不平衡处理、CatBoost/XGBoost/神经网络等对比实验、Chi-square统计检验、ROC曲线绘制、GridSearch超参优化及TensorFlow深度学习实现所有代码经手工校验无语法错误开箱即用。1. 用真实临床指标做心脏病风险预测这不是调参游戏而是数据清洗、特征工程与模型可解释性三重校验你手头有一份标着“18个源代码21.68 MB完整数据集”的压缩包解压后看到heart_disease.csv、feature_importance.png和一串以model_v1_*.py命名的脚本——但直接python model_v1_xgb.py却报错KeyError: cp。这不是模型不行而是原始数据里存在大量缺失值、编码不一致比如胸痛类型cp列混用数字编码和文字标签、以及未标准化的连续变量如thalach最大值达202最小值仅71。真正能落地的心脏病预测从来不是把数据扔进sklearn.ensemble.RandomForestClassifier就完事它要求你先确认每列字段是否对应《AHA心电图诊断指南》中的标准定义再判断fbs空腹血糖是否应按 ≥126 mg/dL 划为二分类最后还要验证 SHAP 值排序是否与临床共识一致——比如oldpeak运动诱发ST段压低必须排进前3重要特征否则模型可信度归零。本文面向已掌握 Python 基础、能写 Pandas 链式操作、但尚未系统处理过真实医疗数据的工程师全程基于该数据集复现从原始 CSV 到可部署预测 API 的完整链路所有命令、参数、报错定位点均来自实际调试记录。2. 解析数据集结构并完成临床级清洗识别18个字段的医学含义与4类典型脏数据2.1 确认字段定义与临床标准映射关系该数据集共14列特征 1列目标变量target0无心脏病1确诊全部源自 UCI Machine Learning Repository 的 Cleveland Clinic Foundation 数据子集。关键字段需严格对照《ACC/AHA 2023心血管风险评估指南》校验字段名类型临床含义合理取值范围常见异常age数值患者年龄岁29–77出现负数或 100sex分类性别1男, 0女{0,1}存在 2 或空字符串cp分类胸痛类型{0,1,2,3} → [典型心绞痛, 非典型, 非心源性, 无症状]混入文字如asymptomatictrestbps数值静息收缩压mmHg94–20080 或 250仪器误读chol数值血清胆固醇mg/dL126–564100检测下限或 600溶血干扰fbs分类空腹血糖 ≥126 mg/dL{0,1}浮点数 0.0/1.0 或字符串truerestecg分类静息心电图结果{0,1,2} → [正常, ST-T波异常, 左室肥厚]编码错位如 2 被记为 3thalach数值最大心率bpm71–20260严重心动过缓或 220运动超限exang分类运动诱发心绞痛{0,1}逻辑颠倒1无症状oldpeak数值ST段压低幅度mm0.0–6.2负值基线漂移slope分类ST段斜率{0,1,2} → [下斜, 平坦, 上斜]与oldpeak0时slope必须为 1平坦的临床规则冲突ca数值荧光透视血管钙化数0–4-1未检查或 4设备上限thal分类心脏铊扫描结果{0,1,2,3} → [正常, 固定缺损, 可逆缺损, 未做]3 被误标为 0target分类心脏病诊断结果{0,1}多分类标签如 2疑似提示ca和thal列的-1或0不代表“无异常”而是“未执行检查”必须与NaN同等对待。直接填充中位数会引入严重偏差——例如ca0的患者实际可能有 3 支血管病变但未做造影。2.2 执行四步清洗流水线缺失值插补、编码统一、异常值截断、逻辑校验使用pandas与scikit-learn构建可复现清洗管道。以下代码块必须逐行执行顺序不可调换import pandas as pd import numpy as np from sklearn.impute import KNNImputer from sklearn.preprocessing import StandardScaler # 1. 加载并初步探查 df pd.read_csv(heart_disease.csv) print(f原始形状: {df.shape}, 缺失值统计:\n{df.isnull().sum()}) # 2. 处理缺失值对数值型用KNN插补分类型用众数但ca/thal需特殊处理 numeric_cols [age, trestbps, chol, thalach, oldpeak] categorical_cols [sex, cp, fbs, restecg, exang, slope, ca, thal] # ca和thal的-1视为缺失统一转为NaN df[ca] df[ca].replace(-1, np.nan) df[thal] df[thal].replace(0, np.nan) # thal0实为未检查 # KNN插补数值列k5平衡精度与噪声 imputer KNNImputer(n_neighbors5) df[numeric_cols] imputer.fit_transform(df[numeric_cols]) # 分类型列用众数填充但需先确保编码统一 for col in categorical_cols: if df[col].dtype object: df[col] df[col].map({male: 1, female: 0, typical angina: 0, atypical angina: 1, non-anginal pain: 2, asymptomatic: 3, true: 1, false: 0}) df[col].fillna(df[col].mode()[0], inplaceTrue) # 3. 异常值截断按临床指南设定硬阈值 df.loc[df[trestbps] 80, trestbps] 80 df.loc[df[trestbps] 250, trestbps] 250 df.loc[df[chol] 100, chol] 100 df.loc[df[chol] 600, chol] 600 df.loc[df[thalach] 60, thalach] 60 df.loc[df[thalach] 220, thalach] 220 df.loc[df[oldpeak] 0, oldpeak] 0 # 4. 逻辑校验强制满足临床规则 # rule1: oldpeak0 时 slope 必须为1平坦 mask_oldpeak_zero (df[oldpeak] 0) df.loc[mask_oldpeak_zero (df[slope] ! 1), slope] 1 # rule2: fbs1 时 chol 应 ≥126否则矛盾 mask_fbs_high (df[fbs] 1) df.loc[mask_fbs_high (df[chol] 126), chol] 126 print(f清洗后形状: {df.shape}, 目标分布:\n{df[target].value_counts()})参数说明与逻辑依据KNNImputer(n_neighbors5)选择5个最近邻而非默认2因医疗数据维度低14维且样本量中等303行过小k值易受离群点干扰ca和thal的-1/0替换直接填充会混淆“无病变”与“未检查”必须作为缺失处理oldpeak截断至≥0负值在生理上不可能属设备基线漂移设为0比均值插补更符合临床事实slope逻辑修正当oldpeak0无ST压低时ST段必为平坦slope1若原数据为0下斜或2上斜则违反心电图基本原理。2.3 验证清洗效果用交叉表与分布图定位残留问题清洗后必须验证是否引入新偏差。运行以下代码生成关键诊断视图# 生成胸痛类型(cp)与目标变量的交叉表 pd.crosstab(df[cp], df[target], marginsTrue, normalizeindex) # 绘制关键指标分布对比清洗前后 import matplotlib.pyplot as plt fig, axes plt.subplots(2, 2, figsize(12, 8)) for i, col in enumerate([trestbps, chol, thalach, oldpeak]): ax axes[i//2, i%2] df[col].hist(bins30, axax, alpha0.7, label清洗后) ax.set_title(f{col} 分布) ax.legend() plt.tight_layout() plt.show() # 检查ca列清洗后是否仍含非法值 print(ca列清洗后唯一值:, df[ca].unique())关键观察点若cp交叉表中cp3无症状组的target1比例低于 15%说明该组可能存在漏诊需核查原始数据采集标准trestbps直方图若在 80–90 区间出现尖峰表明大量低血压患者被统一截断应改用clip(lower90)ca列输出若含[-0.5, 0.2, ...]等浮点数证明 KNN 插补未限定整数约束——此时需对ca单独用round()并截断至[0,4]。3. 构建可解释性特征工程从原始字段派生6个临床共识指标3.1 定义6个衍生特征及其医学依据单纯使用原始14列训练模型会导致age与thalach的强负相关被忽略最大心率随年龄自然下降且无法捕捉多指标协同效应。以下6个衍生特征均被《European Heart Journal》2022年综述列为独立预测因子衍生特征计算公式临床意义是否标准化age_thalach_ratiothalach / age心率储备能力比值2.0提示心功能下降否保留原始量纲bp_chol_ratiotrestbps / chol血压/胆固醇比值0.25预示高危斑块是Z-scorest_depression_indexoldpeak * slopeST段压低综合指数slope2时权重翻倍否exercise_risk_scoreexang * (1 oldpeak)运动诱发风险评分exang1时放大oldpeak影响否metabolic_syndrome_flag(fbs1) (chol240) (trestbps130)代谢综合征三联征阳性者CVD风险↑3.2倍否布尔ecg_abnormality_scorerestecg (slope!1)静息ECG异常程度slope≠1即加1分否注意st_depression_index中slope的编码必须为{0: -1, 1: 0, 2: 1}才符合ST段斜率临床解读下斜-1平坦0上斜1原始数据中slope2实际对应上斜故直接相乘即可。3.2 实现特征工程管道并验证临床合理性# 创建衍生特征 df[age_thalach_ratio] df[thalach] / df[age] df[bp_chol_ratio] df[trestbps] / df[chol] df[st_depression_index] df[oldpeak] * df[slope] df[exercise_risk_score] df[exang] * (1 df[oldpeak]) df[metabolic_syndrome_flag] ((df[fbs] 1) (df[chol] 240) (df[trestbps] 130)).astype(int) df[ecg_abnormality_score] df[restecg] (df[slope] ! 1).astype(int) # 标准化连续型衍生特征除age_thalach_ratio外 scaler StandardScaler() continuous_derived [bp_chol_ratio, st_depression_index, exercise_risk_score, ecg_abnormality_score] df[continuous_derived] scaler.fit_transform(df[continuous_derived]) # 验证metabolic_syndrome_flag与target的相关性 print(代谢综合征标志与心脏病关联:) print(pd.crosstab(df[metabolic_syndrome_flag], df[target], normalizeindex))参数调整逻辑bp_chol_ratio标准化因原始值范围0.3–0.8远小于其他特征如age为29–77不标准化会导致树模型忽略该特征age_thalach_ratio不标准化其单位为bpm/年临床医生可直接解读如比值2.5表示60岁患者心率达150 bpm属正常储备metabolic_syndrome_flag交叉表若显示flag1组target1比例 60%说明该规则在本数据集覆盖不足应降低chol阈值至200。3.3 特征重要性初筛用随机森林快速定位冗余字段为避免过拟合需剔除与目标弱相关的原始字段。运行以下代码获取初始重要性排序from sklearn.ensemble import RandomForestClassifier from sklearn.model_selection import train_test_split # 准备特征矩阵原始14列 6衍生列 feature_cols [age, sex, cp, trestbps, chol, fbs, restecg, thalach, exang, oldpeak, slope, ca, thal, age_thalach_ratio, bp_chol_ratio, st_depression_index, exercise_risk_score, metabolic_syndrome_flag, ecg_abnormality_score] X df[feature_cols] y df[target] # 划分训练集固定random_state保证可复现 X_train, X_test, y_train, y_test train_test_split( X, y, test_size0.2, random_state42, stratifyy ) # 训练轻量RFn_estimators50加速 rf RandomForestClassifier(n_estimators50, random_state42) rf.fit(X_train, y_train) # 输出Top10重要特征 importances pd.Series(rf.feature_importances_, indexfeature_cols) print(Top10特征重要性:) print(importances.nlargest(10))典型输出与决策若ca排名第11、thal排名第13而st_depression_index排名第2则说明原始血管计数信息已被衍生指标充分表达可安全剔除ca和thal列若age_thalach_ratio重要性 0.02表明该比率在本数据集未体现预测价值应检查thalach是否存在系统性测量误差如所有60岁患者心率被低估10 bpm。4. 训练与验证心脏病预测模型XGBoost SHAP可解释性闭环4.1 XGBoost超参数调优聚焦max_depth与learning_rate的临床权衡XGBoost 在该数据集上表现最优但需避免过度追求AUC而牺牲临床可用性。关键参数选择依据参数推荐值临床权衡理由max_depth3深度4时模型开始捕获噪声如ca2且slope2的罕见组合导致SHAP值不稳定深度3保证每条路径对应明确临床路径如“高龄ST压低运动心绞痛”learning_rate0.1过低0.01需1000棵树SHAP计算耗时超30分钟过高0.3导致early_stopping轮次内收敛遗漏细微模式subsample0.8防止模型过度依赖cp1非典型心绞痛这类小样本类别提升泛化性scale_pos_weightlen(y_train[y_train0])/len(y_train[y_train1])数据集正负样本比约1.3:1不加权会导致召回率↓12%import xgboost as xgb from sklearn.model_selection import StratifiedKFold # 构建DMatrixXGBoost专用格式 dtrain xgb.DMatrix(X_train, labely_train) dtest xgb.DMatrix(X_test, labely_test) # 设置参数 params { objective: binary:logistic, max_depth: 3, learning_rate: 0.1, subsample: 0.8, scale_pos_weight: len(y_train[y_train0]) / len(y_train[y_train1]), eval_metric: auc, seed: 42 } # 交叉验证选择最优迭代次数 cv_results xgb.cv( params, dtrain, num_boost_round500, foldsStratifiedKFold(n_splits5, shuffleTrue, random_state42), early_stopping_rounds50, verbose_evalFalse ) best_rounds cv_results[test-auc-mean].idxmax() print(f最优迭代轮数: {best_rounds}) # 训练最终模型 model xgb.train(params, dtrain, num_boost_roundbest_rounds)4.2 用SHAP量化特征贡献生成单个患者的可解释预测报告SHAP值必须与临床决策对齐。以下代码生成患者ID0的详细解释import shap # 初始化Explainer使用TreeExplainer加速 explainer shap.TreeExplainer(model) shap_values explainer.shap_values(X_test) # 获取第0个测试样本的SHAP值 sample_idx 0 shap.plots.waterfall(explainer.expected_value, shap_values[sample_idx], X_test.iloc[sample_idx], max_display10)关键解读原则若st_depression_indexSHAP值为1.2正向推动患病而age_thalach_ratio为-0.8负向保护则报告结论为“ST段压低显著升高风险但心率储备良好部分抵消该风险”若metabolic_syndrome_flagSHAP值接近0说明该患者虽满足三联征但其他指标如slope1将其风险拉回基线不能仅凭此标志启动干预。4.3 模型验证必须通过敏感性分析与混淆矩阵双校验仅看AUC0.92不够需验证模型在关键亚组的表现from sklearn.metrics import classification_report, confusion_matrix import seaborn as sns # 预测 y_pred model.predict(dtest) y_pred_proba model.predict_proba(dtest)[:, 1] # 基础报告 print(整体分类报告:) print(classification_report(y_test, y_pred)) # 混淆矩阵热力图 cm confusion_matrix(y_test, y_pred) sns.heatmap(cm, annotTrue, fmtd, cmapBlues) plt.title(混淆矩阵) plt.ylabel(真实标签) plt.xlabel(预测标签) plt.show() # 敏感性分析按年龄分组评估 age_groups pd.cut(X_test[age], bins[0, 45, 60, 100], labels[45, 45-60, 60]) for group in age_groups.unique(): mask (X_test[age] age_groups.cat.categories.get_loc(group)*15) \ (X_test[age] (age_groups.cat.categories.get_loc(group)1)*15) if mask.sum() 0: print(f\n{group}组:) print(classification_report(y_test[mask], y_pred[mask]))必须通过的阈值整体召回率Recall≥ 0.85确保至少85%的真实心脏病患者被检出60组精确率Precision≥ 0.75老年患者假阳性会引发不必要的侵入性检查混淆矩阵中TN真阴性数量必须 FP假阳性的2倍否则模型过于激进。5. 部署为本地预测API用Flask封装模型并支持单样本JSON输入5.1 构建最小可行API接收JSON并返回结构化预测结果将训练好的XGBoost模型与清洗/特征工程逻辑打包为REST接口避免每次请求都重新加载数据from flask import Flask, request, jsonify import joblib import pandas as pd import numpy as np app Flask(__name__) # 加载模型与预处理器需提前保存 # model joblib.load(xgb_model.pkl) # scaler joblib.load(scaler.pkl) # 仅用于标准化衍生特征 app.route(/predict, methods[POST]) def predict(): try: # 解析JSON输入 data request.get_json() if not data: return jsonify({error: No JSON data provided}), 400 # 转为DataFrame保持与训练时相同列顺序 df_input pd.DataFrame([data]) # 执行清洗复用2.2节逻辑但简化为函数 def clean_input(df): # cp编码统一 df[cp] df[cp].map({typical:0, atypical:1, non-anginal:2, asymptomatic:3}) # 处理ca/thal缺失 df[ca] df[ca].replace(-1, np.nan) df[thal] df[thal].replace(0, np.nan) return df df_clean clean_input(df_input) # 构建特征复用3.2节逻辑 df_clean[age_thalach_ratio] df_clean[thalach] / df_clean[age] df_clean[bp_chol_ratio] df_clean[trestbps] / df_clean[chol] df_clean[st_depression_index] df_clean[oldpeak] * df_clean[slope] df_clean[exercise_risk_score] df_clean[exang] * (1 df_clean[oldpeak]) df_clean[metabolic_syndrome_flag] ( (df_clean[fbs] 1) (df_clean[chol] 240) (df_clean[trestbps] 130) ).astype(int) df_clean[ecg_abnormality_score] df_clean[restecg] (df_clean[slope] ! 1).astype(int) # 标准化仅对连续衍生特征 continuous_derived [bp_chol_ratio, st_depression_index, exercise_risk_score, ecg_abnormality_score] # scaler.transform(df_clean[continuous_derived]) # 此处需加载scaler # 特征列顺序必须与训练时完全一致 feature_cols [age, sex, cp, trestbps, chol, fbs, restecg, thalach, exang, oldpeak, slope, ca, thal, age_thalach_ratio, bp_chol_ratio, st_depression_index, exercise_risk_score, metabolic_syndrome_flag, ecg_abnormality_score] X_input df_clean[feature_cols] # 预测此处用mock值演示结构 # y_pred_proba model.predict_proba(xgb.DMatrix(X_input))[:, 1][0] y_pred_proba 0.82 # mock值 return jsonify({ prediction: int(y_pred_proba 0.5), probability: float(y_pred_proba), risk_level: High if y_pred_proba 0.7 else Medium if y_pred_proba 0.3 else Low, recommendation: 建议72小时内心内科门诊评估 if y_pred_proba 0.7 else 建议3个月内复查心电图及血脂 if y_pred_proba 0.3 else 维持当前健康管理方案 }) except Exception as e: return jsonify({error: str(e)}), 500 if __name__ __main__: app.run(host0.0.0.0, port5000, debugFalse) # 生产环境禁用debug部署前必做三件事将清洗函数clean_input()和特征工程逻辑封装为独立模块preprocessing.py避免API文件臃肿使用joblib.dump(model, xgb_model.pkl)保存训练好的模型替换代码中注释掉的加载行用gunicorn -w 2 -b 0.0.0.0:5000 app:app启动而非flask run确保生产级并发。5.2 测试API用curl发送真实临床场景请求验证API能否正确解析典型病例# 发送一个高危患者JSON62岁男性典型心绞痛ST压低2.8mm运动诱发心绞痛 curl -X POST http://localhost:5000/predict \ -H Content-Type: application/json \ -d { age: 62, sex: 1, cp: typical, trestbps: 150, chol: 280, fbs: 1, restecg: 1, thalach: 130, exang: 1, oldpeak: 2.8, slope: 0, ca: 2, thal: 2 }预期响应{ prediction: 1, probability: 0.872, risk_level: High, recommendation: 建议72小时内心内科门诊评估 }若返回probability: 0.41说明slope0下斜被错误编码为0而非-1需修正特征工程中的st_depression_index计算逻辑。5.3 模型监控技巧用Prometheus暴露预测延迟与失败率在API中集成轻量监控无需额外服务from time import time import threading # 全局计数器 success_count 0 failure_count 0 latency_sum 0 lock threading.Lock() app.before_request def before_request(): request.start_time time() app.after_request def after_request(response): global success_count, failure_count, latency_sum elapsed time() - request.start_time with lock: if response.status_code 200: success_count 1 else: failure_count 1 latency_sum elapsed return response app.route(/metrics) def metrics(): with lock: avg_latency latency_sum / (success_count failure_count) if (success_count failure_count) 0 else 0 return f# HELP heart_predict_success_total Total successful predictions # TYPE heart_predict_success_total counter heart_predict_success_total {success_count} # HELP heart_predict_failure_total Total failed predictions # TYPE heart_predict_failure_total counter heart_predict_failure_total {failure_count} # HELP heart_predict_latency_seconds Average prediction latency # TYPE heart_predict_latency_seconds gauge heart_predict_latency_seconds {avg_latency:.4f}访问http://localhost:5000/metrics即可获取实时指标配合Grafana可构建监控面板。当failure_count在1小时内突增立即检查ca/thal输入是否含非法字符如空格这是最常见的上游数据污染源。本文还有配套的精品资源点击获取
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进