【机器学习实战】基于临床指标的心脏病预测(5 种分类模型对比完整代码)
文章标签
#Python 机器学习 #心脏病预测 #数据分析 #逻辑回归 #XGBoost #决策树 #医疗数据挖掘
一、项目简介
心血管疾病是全球致死率最高的疾病之一,传统心脏病诊断依赖医生经验、多项器械检查,成本高、耗时久。本文基于患者临床体检数据集,使用逻辑回归、决策树、随机森林、XGBoost、SVM 支持向量机5 种经典分类算法构建心脏病二分类预测模型,完整覆盖数据加载→清洗异常值→EDA 可视化探索→特征工程→多模型训练→混淆矩阵 / ROC/AUC 全方位评估全流程。
数据集包含 1319 条患者临床记录,特征涵盖年龄、性别、心率、血压、血糖、肌酸激酶同工酶、肌钙蛋白等关键心梗指标,标签分为positive(患病)/negative(健康),非常适合医疗二分类入门学习。
环境依赖
python
运行
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import matplotlib as mpl
import scipy.stats as stats
from scipy.stats import chi2_contingency
from sklearn.preprocessing import StandardScaler
from sklearn.model_selection import train_test_split
from sklearn.linear_model import LogisticRegression
from sklearn.tree import DecisionTreeClassifier
from sklearn.ensemble import RandomForestClassifier
from sklearn.svm import SVC
from xgboost import XGBClassifier
from sklearn.metrics import classification_report, confusion_matrix, roc_curve, auc
import warnings
warnings.filterwarnings('ignore')
# 解决中文、负号显示
mpl.rcParams['font.family']='SimHei'
plt.rcParams['axes.unicode_minus']=False
二、数据集加载与基础查看
2.1 读取数据 + 重命名中文列
python
运行
# 读取医疗数据集
df=pd.read_csv("Medicaldataset.csv")
# 自定义中文列名
columns=['年龄','性别','心率','收缩压','舒张压','血糖','肌酸激酶同工酶','肌钙蛋白','结果']
df.columns=columns
# 查看前5行
print(df.head())
数据集共 1319 行、9 列,无缺失值,目标列结果为二分类文本标签:positive = 心脏病阳性,negative = 阴性。
2.2 基础信息与描述统计
python
运行
# 数据类型、缺失值检测
print(df.info())
# 数值特征统计分布
print(df.describe())
数据特征:
- 样本总量 1319 条,无缺失、无重复值;
- 性别 0 = 女性,1 = 男性,男性占比 66%;
- 年龄均值 56 岁,50-75 岁人群占比最高;
- 心率、血压、肌酸激酶存在极端异常极值,需要清洗。
三、数据预处理(医疗数据核心步骤)
医疗数据集存在大量录入错误、生理逻辑矛盾数据,必须严格清洗:
3.1 剔除极端生理异常值
python
运行
# 删除心率>1000、收缩压<50、舒张压>140极端错误数据
df=df.loc[(df['心率']<1000)& (df['收缩压']>50)&(df['舒张压']<140)]
print("清洗后数据量:",df.shape[0]) # 剩余1314条
3.2 修正血压逻辑错误(舒张压>收缩压)
正常人体收缩压一定大于舒张压,对录入颠倒的数据交换两列数值:
python
运行
# 筛选异常血压记录
wrong_data=df[df['舒张压']>df['收缩压']]
# 交换收缩压、舒张压
df.loc[wrong_data.index,['收缩压','舒张压']]=df.loc[wrong_data.index,['舒张压','收缩压']].values.copy()
# 校验修正结果
print(df.loc[150:154])
四、探索性数据分析 EDA 可视化
4.1 人群基础分布(年龄 + 性别)
python
运行
plt.figure(figsize=(12,8))
# 年龄直方图+核密度曲线
plt.subplot(211)
plt.hist(df['年龄'],bins=20)
df['年龄'].plot(kind='kde',secondary_y=True)
plt.title('患者年龄分布')
# 性别饼图
plt.subplot(212)
gender_data=df['性别'].value_counts()
plt.pie(gender_data,autopct='%.2f%%',labels=['男','女'],shadow=True,explode=(0.05,0))
plt.title('样本性别比例')
plt.tight_layout()
plt.show()
结论:样本中老年群体居多,男性样本远多于女性。
4.2 各指标与心脏病诊断关系箱线图
python
运行
plt.figure(figsize=(20,16))
# 1.年龄与患病
plt.subplot(3,3,1)
pos_age=df[df['结果']=='positive']['年龄']
neg_age=df[df['结果']=='negative']['年龄']
plt.boxplot([pos_age,neg_age], labels=['阳性','阴性'])
plt.title('年龄与心脏病诊断关系')
# 2.心率
plt.subplot(332)
pos_heart=df[df['结果']=='positive']['心率']
neg_heart=df[df['结果']=='negative']['心率']
plt.boxplot([pos_heart,neg_heart],labels=['阳性','阴性'])
plt.title('心率与心脏病诊断关系')
# 3.阳性患者性别占比
plt.subplot(333)
pos_data=df[df['结果']=='positive']
gender_data=pos_data['性别'].value_counts()
plt.pie(gender_data,labels=['男','女'],autopct='%.2f%%')
plt.title('确诊心脏病患者性别比例')
# 4.收缩压、舒张压、血糖、肌酸激酶、肌钙蛋白分组箱线图
ax4=plt.subplot(334)
df.boxplot(column='收缩压',by='结果',ax=ax4)
ax5=plt.subplot(335)
df.boxplot(column='舒张压',by='结果',ax=ax5)
ax6=plt.subplot(336)
df.boxplot(column='血糖',by='结果',ax=ax6)
ax7=plt.subplot(337)
df.boxplot(column='肌酸激酶同工酶',by='结果',ax=ax7)
ax8=plt.subplot(338)
df.boxplot(column='肌钙蛋白',by='结果',ax=ax8)
plt.show()
关键医学结论:
- 阳性患者肌钙蛋白、肌酸激酶同工酶数值显著高于阴性,是心梗核心标志物;
- 心脏病阳性人群心率、血压整体高于健康人群;
- 男性心脏病发病样本数量更多。
五、特征工程与数据集划分
5.1 标签二值化
将文本标签转换为 0/1 数值:1 = 患病 positive,0 = 健康 negative
python
运行
df['Result_Binary']=(df['结果']=='positive').astype(int)
5.2 筛选核心预测特征
根据医学常识与可视化分布,选取区分度最高 4 个特征:性别、年龄、肌酸激酶同工酶、肌钙蛋白
python
运行
features=['性别','年龄','肌酸激酶同工酶','肌钙蛋白']
X=df[features].copy()
y=df['Result_Binary']
5.3 连续特征标准化
年龄、肌酸激酶、肌钙蛋白为连续变量,使用 Z-score 标准化,消除量纲影响(SVM、逻辑回归依赖标准化)
python
运行
continuous_vars=['年龄','肌酸激酶同工酶','肌钙蛋白']
X[continuous_vars]=X[continuous_vars].astype(float)
scaler=StandardScaler()
X.loc[:,continuous_vars]=scaler.fit_transform(X[continuous_vars])
5.4 划分训练集 / 测试集(8:2 分层抽样)
分层抽样保证训练、测试集正负样本比例和原数据一致:
python
运行
X_train,X_test,y_train,y_test=train_test_split(X,y,test_size=0.2,random_state=15,stratify=y)
print(f"训练集大小:{X_train.shape}, 测试集大小:{X_test.shape}")
六、通用模型评估函数(复用所有模型)
统一输出分类报告、混淆矩阵热力图、ROC 曲线 + AUC 值,减少重复代码:
python
运行
def evaluate_model(model,model_name,X_train,X_test,y_train,y_test):
# 训练模型
model.fit(X_train,y_train)
# 预测标签
y_pred=model.predict(X_test)
print(f"====={model_name} 模型评估=====")
print(classification_report(y_test,y_pred))
# 绘制混淆矩阵
cm=confusion_matrix(y_test,y_pred)
plt.figure(figsize=(8,6))
plt.imshow(cm, cmap='hot', interpolation='nearest')
plt.colorbar()
for i in range(2):
for j in range(2):
plt.text(j, i,f'{cm[i, j]:.2f}',ha='center',va='center',color='black')
plt.title(f'{model_name}混淆矩阵')
plt.show()
# 计算ROC与AUC
if hasattr(model,"predict_proba"):
y_prob=model.predict_proba(X_test)[:,1]
else:
y_prob=model.decision_function(X_test)
fpr, tpr,_= roc_curve(y_test, y_prob)
roc_auc=auc(fpr,tpr)
plt.figure(figsize=(8,6))
plt.plot(fpr,tpr,color='darkorange',lw=2,label=f'ROC曲线(面积={roc_auc:.2f})')
plt.plot([0,1],[0,1],color='navy',lw=2,linestyle='--')
plt.xlabel('假阳率(FPR)')
plt.ylabel('真阳率(TPR)')
plt.title(f'{model_name}ROC曲线')
plt.legend(loc="lower right")
plt.show()
return model,roc_auc
七、5 种模型训练与性能结果
7.1 逻辑回归(线性基准模型)
python
运行
lr_model=LogisticRegression(random_state=15, class_weight='balanced')
lr_model,lr_auc=evaluate_model(lr_model,"逻辑回归",X_train,X_test,y_train,y_test)
结果:准确率 0.83,AUC=0.92;线性模型,可解释性强,但无法捕捉特征非线性关系。
7.2 决策树(单棵树,高可解释)
python
运行
dt_model=DecisionTreeClassifier(random_state=15, class_weight='balanced')
dt_model,dt_auc=evaluate_model(dt_model,"决策树",X_train,X_test,y_train,y_test)
结果:准确率 0.98,AUC=0.98;单树拟合能力极强,无需标准化,可直接输出特征判断规则。
7.3 随机森林(集成树,泛化能力强)
python
运行
rf_model=RandomForestClassifier(random_state=15, class_weight='balanced')
rf_model,rf_auc=evaluate_model(rf_model,"随机森林",X_train,X_test,y_train,y_test)
结果:准确率 0.97,AUC=0.99;多棵树集成,缓解单棵决策树过拟合,稳定性最优。
7.4 XGBoost(梯度提升树,工业级高效模型)
python
运行
xgb_model = XGBClassifier(random_state=15, scale_pos_weight=sum(y_train==0)/sum(y_train==1))
xgb_model, xgb_auc = evaluate_model(xgb_model,"XGBoost",X_train,X_test,y_train,y_test)
结果:准确率 0.97,AUC=0.99;梯度提升算法,对医疗不平衡样本适配性优秀。
7.5 SVM 支持向量机(高维分类)
python
运行
svm_model=SVC(random_state=15,class_weight='balanced',probability=True)
svm_model, svm_auc=evaluate_model(svm_model,"支持向量机",X_train,X_test,y_train,y_test)
结果:准确率 0.77,AUC=0.92;本数据集表现最差,对特征噪声敏感,依赖精细调参。
八、模型性能汇总对比表
表格
| 模型 | 测试集准确率 | AUC 值 | 优缺点总结 |
|---|---|---|---|
| 逻辑回归 | 0.83 | 0.92 | 可解释性强,适合基线对比,非线性拟合弱 |
| 决策树 | 0.98 | 0.98 | 精度最高,规则直观,易过拟合 |
| 随机森林 | 0.97 | 0.99 | 综合最优,泛化能力强,稳定性高 |
| XGBoost | 0.97 | 0.99 | 梯度提升,医疗不平衡数据表现优秀 |
| SVM | 0.77 | 0.92 | 效果最差,对数据分布敏感,调参成本高 |
最优模型结论
- 若追求最高预测精度:选择决策树,准确率 98%,可直接提取临床判断规则;
- 若追求泛化能力、线上落地稳定性:选择随机森林 / XGBoost,AUC 达到 0.99,抗噪声能力更强;
- 逻辑回归适合做基线模型,用于特征重要性分析;SVM 在此场景不推荐优先使用。
九、项目总结与拓展方向
1. 项目收获
- 掌握医疗表格数据完整处理流程:异常值清洗、生理逻辑纠错、EDA 临床可视化;
- 熟练 5 种主流分类模型实现,统一评估框架(分类报告、混淆矩阵、ROC-AUC);
- 理解医疗预测核心指标:肌钙蛋白、肌酸激酶同工酶是心脏病强区分特征。
2. 后续优化拓展
- 超参数网格搜索:GridSearchCV 调优树模型、SVM,进一步提升精度;
- 可解释 AI:SHAP 值分析各临床指标对心脏病预测的贡献度;
- 类别不平衡优化:SMOTE 过采样处理极端不均衡医疗样本;
- 系统部署:使用 Flask/Streamlit 搭建网页心脏病风险预测工具;
- 多模型融合:Stacking 融合随机森林 + XGBoost,进一步提升 AUC。
十、完整源码获取
文中所有代码可直接复制运行,仅需替换Medicaldataset.csv文件路径即可复现全部图表与模型结果。数据集为公开心血管临床数据集,适合机器学习初学者医疗方向练手。
博文配套优化建议(发布时使用)
- 配图:把箱线图、混淆矩阵、ROC 曲线截图插入文章对应段落;
- 目录:添加 markdown 自动目录,提升阅读体验;
- 评论区预留:数据集下载链接、答疑;
- 置顶标签:# 机器学习实战 #医疗数据分析 #XGBoost 心脏病预测
更多推荐
所有评论(0)