基于PyQt和深度学习的心脏病智能检测系统
一、研究背景与意义
心血管疾病是当今世界范围内的头号死亡原因,根据世界卫生组织的统计数据,每年约有1790万人死于心血管疾病,占全球死亡人数的32%。心脏病的早期预测和诊断对降低死亡率具有重要意义。随着人工智能和深度学习技术的快速发展,将其应用于医疗健康领域,特别是心脏病预测与诊断,已成为当前研究的热点。
本课题旨在开发一套基于深度学习算法的心脏病智能检测系统,通过整合患者的临床数据(如年龄、性别、血压、心电图结果等),构建一个能够准确预测心脏病风险的模型,并通过PyQt框架构建直观友好的用户界面,使医护人员能够便捷地使用该系统进行辅助诊断。
该研究的意义主要体现在以下几个方面:
- 辅助临床决策:为医生提供客观的风险评估依据,辅助临床决策
- 提高诊断效率:减少人工诊断的时间成本,提高医疗资源利用效率
- 普及心脏健康监测:降低专业门槛,使基层医疗机构也能获得高质量的诊断支持
- 促进人工智能医疗落地:探索AI技术在医疗领域的实际应用路径
二、国内外研究现状
国外研究现状
国外在医疗AI领域的研究已相对成熟。Stanford大学研究团队开发的CheXNet模型在肺炎诊断方面已超过放射科医师水平;Google DeepMind的AI系统在眼疾检测领域取得显著成果。在心脏病预测方面,2020年发表在《Nature》上的研究表明,基于深度学习的模型在心脏病风险预测的准确率可达85%以上。Mayo Clinic和IBM Watson Health合作开发的心脏病风险评估系统已进入临床试验阶段。
国内研究现状
国内在医疗AI领域的研究近年来快速发展。中国科学院自动化研究所与多家医院合作开发的智能医学影像分析系统已在多家医院试用;阿里健康与浙江大学医学院附属第一医院合作研发的心电图AI辅助诊断系统在临床应用中取得了良好效果。清华大学与北京协和医院合作的基于深度学习的心脏病预测模型,准确率已达80%左右。
研究现状分析
尽管国内外在医疗AI领域取得了一定进展,但仍面临以下挑战:
- 模型解释性不足,医生对AI决策过程缺乏理解
- 医疗数据获取困难,数据质量参差不齐
- AI系统与临床工作流程融合不够紧密
- 用户界面不够友好,医护人员使用门槛较高
本课题将针对上述问题,特别关注系统的实用性和用户体验,通过PyQt开发直观友好的界面,并结合深度学习模型的高准确性,提供一个完整的解决方案。
三、研究内容
本课题的主要研究内容包括:
- 数据收集与预处理
- 收集整理心脏病相关临床数据集(如UCI心脏病数据集)
- 数据清洗、标准化和特征工程
- 数据增强技术应用,解决类别不平衡问题
- 深度学习模型设计与训练
- 构建适合心脏病预测的深度神经网络模型
- 实现多种经典网络(如DNN、CNN、LSTM等)并比较性能
- 模型优化(超参数调整、正则化、早停等)
- 模型评估与验证(准确率、灵敏度、特异度、AUC等指标)
- PyQt用户界面设计与实现
- 系统总体架构设计
- 患者信息管理模块
- 数据可视化模块(特征重要性、相关性分析等)
- 预测结果展示与解释模块
- 历史记录管理模块
- 系统集成与测试
- 深度学习模型与PyQt界面的集成
- 系统功能测试与性能优化
- 用户体验测试与改进
四、研究方案与技术路线
技术路线
- 前期准备阶段(1个月)
- 文献调研
- 技术选型
- 数据集收集与分析
- 模型开发阶段(2个月)
- 数据预处理
- 特征工程
- 深度学习模型构建与训练
- 模型评估与优化
- 界面开发阶段(1.5个月)
- PyQt界面设计
- 功能模块实现
- 数据可视化实现
- 系统集成阶段(1个月)
- 模型与界面集成
- 系统测试与调优
- 用户体验优化
- 总结完善阶段(0.5个月)
- 系统功能完善
- 文档编写
- 毕业论文撰写
技术方案
- 开发环境与工具
- 编程语言:Python 3.8+
- 深度学习框架:TensorFlow 2.x/PyTorch
- GUI框架:PyQt5
- 开发工具:PyCharm, VS Code
- 版本控制:Git
- 数据处理方案
- 采用UCI心脏病数据集、Cleveland心脏病数据集等公开数据集
- 使用pandas, NumPy进行数据处理
- 采用scikit-learn进行数据预处理和特征选择
- 深度学习模型方案
- 基础模型:多层感知机(MLP)
- 进阶模型:考虑LSTM处理时序数据(如ECG数据)
- 模型评估:采用5折交叉验证
- 评价指标:准确率、精确率、召回率、F1分数、AUC
- 界面设计方案
- 采用PyQt5框架开发GUI
- 使用Matplotlib, Seaborn实现数据可视化
- 界面风格:医疗风格,简洁明了
- 主要模块:患者信息录入、数据可视化、预测结果展示
五、预期成果与创新点
预期成果
- 一套完整的基于PyQt和深度学习的心脏病智能检测系统
- 一个训练好的、准确率在85%以上的心脏病预测模型
- 详细的系统设计文档和用户手册
- 毕业论文一篇
创新点
- 模型与界面深度融合:不仅提供预测结果,还展示模型决策依据,增强可解释性
- 交互式数据可视化:通过直观的可视化方式展示患者数据特征与心脏病风险的关联
- 风险分层评估:不仅给出二分类结果(有/无心脏病),还提供风险等级评估
- 本地化部署设计:考虑医疗数据隐私保护,系统设计为可本地部署运行,无需上传数据至云端
六、研究基础与可行性分析
研究基础
- 理论基础:已完成机器学习、深度学习、数据挖掘等相关课程学习
- 技术基础:熟悉Python编程,了解PyQt框架,掌握TensorFlow/PyTorch等深度学习框架的使用
- 数据基础:可获取公开的心脏病数据集,如UCI心脏病数据集、Cleveland心脏病数据集等
可行性分析
- 技术可行性:深度学习在医疗诊断领域已有大量成功案例,PyQt是成熟的GUI开发框架
- 数据可行性:有多个公开的心脏病数据集可用于模型训练和测试
- 时间可行性:项目各阶段时间安排合理,总工期约6个月,符合毕业设计时间要求
- 硬件可行性:心脏病预测模型复杂度相对较低,普通PC即可进行开发和运行
七、预期困难与解决方案
预期困难
- 数据质量问题:公开数据集可能存在噪声、缺失值等质量问题
- 模型准确性不足:深度学习模型在小样本数据集上可能表现不佳
- 界面设计专业性:医疗软件界面设计需要符合医护人员使用习惯
- 系统实用性验证:缺乏真实医疗环境的验证渠道
解决方案
- 数据质量问题:采用多种数据清洗和填补技术;考虑多数据集融合
- 模型准确性不足:采用数据增强技术;使用迁移学习;尝试集成学习方法
- 界面设计专业性:参考现有医疗软件界面设计;查阅医疗软件设计指南
- 系统实用性验证:采用模拟医疗场景进行测试;邀请医学专业学生进行评价
八、参考文献
- Rajkomar A, Dean J, Kohane I. Machine learning in medicine[J]. New England Journal of Medicine, 2019, 380(14): 1347-1358.
- Hannun AY, Rajpurkar P, Haghpanahi M, et al. Cardiologist-level arrhythmia detection and classification in ambulatory electrocardiograms using a deep neural network[J]. Nature Medicine, 2019, 25(1): 65-69.
- Johnson KW, Soto JT, Glicksberg BS, et al. Artificial intelligence in cardiology[J]. Journal of the American College of Cardiology, 2018, 71(23): 2668-2679.
- Attia ZI, Noseworthy PA, Lopez-Jimenez F, et al. An artificial intelligence-enabled ECG algorithm for the identification of patients with atrial fibrillation during sinus rhythm: a retrospective analysis of outcome prediction[J]. The Lancet, 2019, 394(10201): 861-867.
- Riverio M, Mitsnefes M, Blake P, et al. Deep learning for the prediction of early outcomes after pediatric heart transplantation[J]. The Journal of Heart and Lung Transplantation, 2021, 40(2): 123-130.
- Summerfield M. Rapid GUI programming with Python and Qt: The definitive guide to PyQt programming[M]. Pearson Education, 2007.
- Goodfellow I, Bengio Y, Courville A. Deep learning[M]. MIT press, 2016.
- Janosi A, Steinbrunn W, Pfisterer M, et al. Heart Disease[DB/OL]. UCI Machine Learning Repository, 1988.
核心设计部分(仅供学习和参考)
基于PyQt的深度学习心脏病智能检测系统
项目架构
1. 项目结构
CardiacDiseaseDetection/
│
├── main.py # 主程序入口
├── requirements.txt # 依赖包列表
├── README.md # 项目说明
│
├── models/ # 深度学习模型
│ ├── model.py # 模型定义
│ ├── train.py # 模型训练
│ ├── predict.py # 模型预测
│ └── checkpoints/ # 保存训练好的模型
│
├── data/ # 数据
│ ├── raw/ # 原始数据
│ ├── processed/ # 处理后的数据
│ └── data_processor.py # 数据预处理
│
├── ui/ # 用户界面
│ ├── main_window.py # 主窗口
│ ├── patient_info.py # 患者信息界面
│ ├── data_visualization.py # 数据可视化界面
│ ├── prediction_view.py # 预测结果展示
│ └── resources/ # UI资源(图标等)
│
└── utils/ # 工具函数
├── logger.py # 日志
├── validators.py # 数据验证
└── file_handlers.py # 文件处理
代码实现
1. 主程序入口 (main.py)
import sys
from PyQt5.QtWidgets import QApplication
from ui.main_window import MainWindow
if __name__ == "__main__":
app = QApplication(sys.argv)
app.setStyle('Fusion')
window = MainWindow()
window.show()
sys.exit(app.exec_())
2. 主窗口界面 (ui/main_window.py)
from PyQt5.QtWidgets import (QMainWindow, QTabWidget, QAction, QMessageBox,
QFileDialog, QVBoxLayout, QWidget)
from PyQt5.QtGui import QIcon
from PyQt5.QtCore import Qt
from ui.patient_info import PatientInfoWidget
from ui.data_visualization import DataVisualizationWidget
from ui.prediction_view import PredictionWidget
from models.predict import HeartDiseasePredictor
from utils.logger import setup_logger
class MainWindow(QMainWindow):
def __init__(self):
super().__init__()
self.logger = setup_logger('main_window')
self.predictor = HeartDiseasePredictor()
self.init_ui()
def init_ui(self):
self.setWindowTitle('心脏病智能检测系统')
self.setGeometry(100, 100, 1000, 700)
# 创建菜单栏
self.create_menu_bar()
# 创建主要部件
self.tabs = QTabWidget()
self.patient_info_widget = PatientInfoWidget(self.predictor)
self.data_viz_widget = DataVisualizationWidget()
self.prediction_widget = PredictionWidget()
# 添加标签页
self.tabs.addTab(self.patient_info_widget, "患者信息")
self.tabs.addTab(self.data_viz_widget, "数据可视化")
self.tabs.addTab(self.prediction_widget, "疾病预测")
# 连接信号
self.patient_info_widget.prediction_ready.connect(self.prediction_widget.update_results)
self.patient_info_widget.data_ready.connect(self.data_viz_widget.update_charts)
# 设置中心部件
self.setCentralWidget(self.tabs)
self.logger.info('Main window initialized')
def create_menu_bar(self):
menubar = self.menuBar()
# 文件菜单
file_menu = menubar.addMenu('文件')
# 打开动作
open_action = QAction('打开数据集', self)
open_action.setShortcut('Ctrl+O')
open_action.triggered.connect(self.open_file)
file_menu.addAction(open_action)
# 保存动作
save_action = QAction('保存结果', self)
save_action.setShortcut('Ctrl+S')
save_action.triggered.connect(self.save_file)
file_menu.addAction(save_action)
# 退出动作
exit_action = QAction('退出', self)
exit_action.setShortcut('Ctrl+Q')
exit_action.triggered.connect(self.close)
file_menu.addAction(exit_action)
# 帮助菜单
help_menu = menubar.addMenu('帮助')
about_action = QAction('关于', self)
about_action.triggered.connect(self.show_about)
help_menu.addAction(about_action)
def open_file(self):
filename, _ = QFileDialog.getOpenFileName(
self, "打开文件", "", "CSV Files (*.csv);;All Files (*)"
)
if filename:
self.patient_info_widget.load_data(filename)
self.logger.info(f'Opened file: {filename}')
def save_file(self):
filename, _ = QFileDialog.getSaveFileName(
self, "保存文件", "", "CSV Files (*.csv);;Text Files (*.txt);;All Files (*)"
)
if filename:
self.prediction_widget.save_results(filename)
self.logger.info(f'Saved results to: {filename}')
def show_about(self):
QMessageBox.about(self, "关于",
"心脏病智能检测系统 v1.0\n"
"基于深度学习的心脏病风险预测系统\n"
"©2023 本科毕业设计")
3. 患者信息界面 (ui/patient_info.py)
from PyQt5.QtWidgets import (QWidget, QFormLayout, QLineEdit, QComboBox,
QPushButton, QDoubleSpinBox, QSpinBox,
QGroupBox, QVBoxLayout, QHBoxLayout,
QLabel, QScrollArea, QMessageBox)
from PyQt5.QtCore import pyqtSignal, Qt
from PyQt5.QtGui import QFont
import pandas as pd
import numpy as np
from utils.validators import validate_patient_data
class PatientInfoWidget(QWidget):
prediction_ready = pyqtSignal(dict)
data_ready = pyqtSignal(pd.DataFrame)
def __init__(self, predictor):
super().__init__()
self.predictor = predictor
self.init_ui()
def init_ui(self):
main_layout = QVBoxLayout()
# 创建表单布局
form_layout = QFormLayout()
# 姓名
self.name_edit = QLineEdit()
form_layout.addRow("姓名:", self.name_edit)
# 年龄
self.age_spin = QSpinBox()
self.age_spin.setRange(1, 120)
self.age_spin.setValue(50)
form_layout.addRow("年龄:", self.age_spin)
# 性别
self.gender_combo = QComboBox()
self.gender_combo.addItems(["男", "女"])
form_layout.addRow("性别:", self.gender_combo)
# 创建临床数据分组
clinical_group = QGroupBox("临床数据")
clinical_layout = QFormLayout()
# 休息心率
self.rest_ecg_combo = QComboBox()
self.rest_ecg_combo.addItems(["正常", "ST-T波异常", "左心室肥大"])
clinical_layout.addRow("心电图结果:", self.rest_ecg_combo)
# 最大心率
self.max_hr_spin = QSpinBox()
self.max_hr_spin.setRange(60, 220)
self.max_hr_spin.setValue(150)
clinical_layout.addRow("最大心率:", self.max_hr_spin)
# 胸痛类型
self.chest_pain_combo = QComboBox()
self.chest_pain_combo.addItems(["典型心绞痛", "非典型心绞痛", "非心绞痛", "无症状"])
clinical_layout.addRow("胸痛类型:", self.chest_pain_combo)
# 运动诱发心绞痛
self.exang_combo = QComboBox()
self.exang_combo.addItems(["是", "否"])
clinical_layout.addRow("运动诱发心绞痛:", self.exang_combo)
# 收缩压
self.trestbps_spin = QSpinBox()
self.trestbps_spin.setRange(90, 200)
self.trestbps_spin.setValue(120)
clinical_layout.addRow("收缩压(mmHg):", self.trestbps_spin)
# 胆固醇
self.chol_spin = QSpinBox()
self.chol_spin.setRange(100, 600)
self.chol_spin.setValue(200)
clinical_layout.addRow("胆固醇(mg/dl):", self.chol_spin)
# 空腹血糖
self.fbs_combo = QComboBox()
self.fbs_combo.addItems(["<120 mg/dl", ">120 mg/dl"])
clinical_layout.addRow("空腹血糖:", self.fbs_combo)
# 坡度
self.slope_combo = QComboBox()
self.slope_combo.addItems(["上坡", "平坦", "下坡"])
clinical_layout.addRow("ST段坡度:", self.slope_combo)
# 主要血管数
self.ca_spin = QSpinBox()
self.ca_spin.setRange(0, 4)
clinical_layout.addRow("主要血管数:", self.ca_spin)
# 血液流变
self.thal_combo = QComboBox()
self.thal_combo.addItems(["正常", "固定缺陷", "可逆缺陷"])
clinical_layout.addRow("血液流变:", self.thal_combo)
# ST抑制
self.oldpeak_spin = QDoubleSpinBox()
self.oldpeak_spin.setRange(0, 10)
self.oldpeak_spin.setSingleStep(0.1)
clinical_layout.addRow("ST抑制:", self.oldpeak_spin)
clinical_group.setLayout(clinical_layout)
# 添加按钮
buttons_layout = QHBoxLayout()
self.clear_btn = QPushButton("清除")
self.clear_btn.clicked.connect(self.clear_form)
self.predict_btn = QPushButton("预测")
self.predict_btn.clicked.connect(self.predict)
self.predict_btn.setStyleSheet("background-color: #4CAF50; color: white;")
buttons_layout.addWidget(self.clear_btn)
buttons_layout.addWidget(self.predict_btn)
# 将所有部件添加到主布局
main_layout.addLayout(form_layout)
main_layout.addWidget(clinical_group)
main_layout.addLayout(buttons_layout)
self.setLayout(main_layout)
def clear_form(self):
self.name_edit.clear()
self.age_spin.setValue(50)
self.gender_combo.setCurrentIndex(0)
self.rest_ecg_combo.setCurrentIndex(0)
self.max_hr_spin.setValue(150)
self.chest_pain_combo.setCurrentIndex(0)
self.exang_combo.setCurrentIndex(0)
self.trestbps_spin.setValue(120)
self.chol_spin.setValue(200)
self.fbs_combo.setCurrentIndex(0)
self.slope_combo.setCurrentIndex(0)
self.ca_spin.setValue(0)
self.thal_combo.setCurrentIndex(0)
self.oldpeak_spin.setValue(0)
def get_form_data(self):
data = {
'name': self.name_edit.text(),
'age': self.age_spin.value(),
'sex': 1 if self.gender_combo.currentText() == "男" else 0,
'cp': self.chest_pain_combo.currentIndex(),
'trestbps': self.trestbps_spin.value(),
'chol': self.chol_spin.value(),
'fbs': 1 if self.fbs_combo.currentIndex() == 1 else 0,
'restecg': self.rest_ecg_combo.currentIndex(),
'thalach': self.max_hr_spin.value(),
'exang': 1 if self.exang_combo.currentIndex() == 0 else 0,
'oldpeak': self.oldpeak_spin.value(),
'slope': self.slope_combo.currentIndex(),
'ca': self.ca_spin.value(),
'thal': self.thal_combo.currentIndex() + 1
}
return data
def predict(self):
data = self.get_form_data()
# 验证数据
if not validate_patient_data(data):
QMessageBox.warning(self, "数据验证", "请填写所有必填项!")
return
# 准备模型输入
model_input = np.array([
[data['age'], data['sex'], data['cp'], data['trestbps'],
data['chol'], data['fbs'], data['restecg'], data['thalach'],
data['exang'], data['oldpeak'], data['slope'], data['ca'], data['thal']]
])
# 进行预测
prediction = self.predictor.predict(model_input)
# 添加预测结果
data['prediction'] = prediction[0]
data['probability'] = prediction[1]
# 发送预测结果信号
self.prediction_ready.emit(data)
# 创建数据可视化用的DataFrame
df = pd.DataFrame([data])
self.data_ready.emit(df)
# 显示成功消息
QMessageBox.information(self, "预测完成", f"预测完成!\n风险评估已生成。")
def load_data(self, filename):
try:
data = pd.read_csv(filename)
if len(data) > 0:
# 使用第一行数据填充表单
row = data.iloc[0]
if 'age' in row:
self.age_spin.setValue(int(row['age']))
if 'sex' in row:
self.gender_combo.setCurrentIndex(int(row['sex']))
if 'cp' in row:
self.chest_pain_combo.setCurrentIndex(int(row['cp']))
if 'trestbps' in row:
self.trestbps_spin.setValue(int(row['trestbps']))
if 'chol' in row:
self.chol_spin.setValue(int(row['chol']))
if 'fbs' in row:
self.fbs_combo.setCurrentIndex(int(row['fbs']))
if 'restecg' in row:
self.rest_ecg_combo.setCurrentIndex(int(row['restecg']))
if 'thalach' in row:
self.max_hr_spin.setValue(int(row['thalach']))
if 'exang' in row:
self.exang_combo.setCurrentIndex(int(row['exang']))
if 'oldpeak' in row:
self.oldpeak_spin.setValue(float(row['oldpeak']))
if 'slope' in row:
self.slope_combo.setCurrentIndex(int(row['slope']))
if 'ca' in row:
self.ca_spin.setValue(int(row['ca']))
if 'thal' in row:
self.thal_combo.setCurrentIndex(int(row['thal']) - 1)
# 通知数据可视化窗口
self.data_ready.emit(data)
QMessageBox.information(self, "数据加载", "数据已成功加载!")
except Exception as e:
QMessageBox.critical(self, "数据加载错误", f"加载数据时出错: {str(e)}")
4. 数据可视化界面 (ui/data_visualization.py)
from PyQt5.QtWidgets import (QWidget, QVBoxLayout, QHBoxLayout, QComboBox,
QLabel, QPushButton, QGroupBox)
from PyQt5.QtCore import Qt
import matplotlib.pyplot as plt
from matplotlib.backends.backend_qt5agg import FigureCanvasQTAgg as FigureCanvas
from matplotlib.backends.backend_qt5agg import NavigationToolbar2QT as NavigationToolbar
import pandas as pd
import numpy as np
import seaborn as sns
class DataVisualizationWidget(QWidget):
def __init__(self):
super().__init__()
self.data = None
self.init_ui()
def init_ui(self):
main_layout = QVBoxLayout()
# 控制区域
control_layout = QHBoxLayout()
self.chart_type_combo = QComboBox()
self.chart_type_combo.addItems(["柱状图", "散点图", "箱线图", "热力图", "分布图"])
self.chart_type_combo.currentIndexChanged.connect(self.update_charts)
self.feature_combo = QComboBox()
self.feature_combo.addItems(["年龄", "胆固醇", "心率", "血压", "所有特征"])
self.feature_combo.currentIndexChanged.connect(self.update_charts)
self.refresh_btn = QPushButton("刷新")
self.refresh_btn.clicked.connect(self.update_charts)
control_layout.addWidget(QLabel("图表类型:"))
control_layout.addWidget(self.chart_type_combo)
control_layout.addWidget(QLabel("特征:"))
control_layout.addWidget(self.feature_combo)
control_layout.addWidget(self.refresh_btn)
# 图表区域
self.figure = plt.figure(figsize=(10, 8))
self.canvas = FigureCanvas(self.figure)
self.toolbar = NavigationToolbar(self.canvas, self)
chart_layout = QVBoxLayout()
chart_layout.addWidget(self.toolbar)
chart_layout.addWidget(self.canvas)
chart_group = QGroupBox("数据可视化")
chart_group.setLayout(chart_layout)
main_layout.addLayout(control_layout)
main_layout.addWidget(chart_group)
self.setLayout(main_layout)
# 初始化空白图表
self.figure.clear()
self.canvas.draw()
def update_charts(self):
if self.data is None:
# 示例数据用于初始显示
self.data = pd.DataFrame({
'age': np.random.normal(50, 10, 100),
'sex': np.random.choice([0, 1], 100),
'cp': np.random.choice([0, 1, 2, 3], 100),
'trestbps': np.random.normal(120, 20, 100),
'chol': np.random.normal(200, 40, 100),
'target': np.random.choice([0, 1], 100)
})
self.figure.clear()
chart_type = self.chart_type_combo.currentText()
feature = self.feature_combo.currentText()
if feature == "年龄":
col = 'age'
elif feature == "胆固醇":
col = 'chol'
elif feature == "心率":
col = 'thalach' if 'thalach' in self.data.columns else 'age'
elif feature == "血压":
col = 'trestbps'
if chart_type == "柱状图":
self.plot_bar_chart(col if feature != "所有特征" else None)
elif chart_type == "散点图":
self.plot_scatter(col if feature != "所有特征" else None)
elif chart_type == "箱线图":
self.plot_boxplot(col if feature != "所有特征" else None)
elif chart_type == "热力图":
self.plot_heatmap()
elif chart_type == "分布图":
self.plot_distribution(col if feature != "所有特征" else None)
self.canvas.draw()
def plot_bar_chart(self, feature=None):
ax = self.figure.add_subplot(111)
if feature:
if 'target' in self.data.columns:
# 根据目标变量分组计算平均值
grouped = self.data.groupby('target')[feature].mean().reset_index()
sns.barplot(x='target', y=feature, data=grouped, ax=ax)
ax.set_title(f'{feature} 平均值 (按诊断分组)')
ax.set_xlabel('心脏病诊断 (0=健康, 1=患病)')
else:
sns.barplot(y=feature, data=self.data, ax=ax)
ax.set_title(f'{feature} 平均值')
else:
# 显示所有数值特征的平均值
numeric_cols = self.data.select_dtypes(include=[np.number]).columns.tolist()
if 'target' in numeric_cols:
numeric_cols.remove('target')
if numeric_cols:
means = self.data[numeric_cols].mean().sort_values(ascending=False)
sns.barplot(x=means.values, y=means.index, ax=ax)
ax.set_title('特征平均值')
ax.set_xlabel('平均值')
def plot_scatter(self, feature=None):
if 'target' not in self.data.columns:
ax = self.figure.add_subplot(111)
if feature:
ax.scatter(range(len(self.data)), self.data[feature])
ax.set_title(f'{feature} 散点图')
ax.set_xlabel('样本索引')
ax.set_ylabel(feature)
else:
ax.text(0.5, 0.5, '需要目标变量进行散点图对比',
horizontalalignment='center', verticalalignment='center')
return
if feature:
ax = self.figure.add_subplot(111)
for target, group in self.data.groupby('target'):
label = '患病' if target == 1 else '健康'
ax.scatter(range(len(group)), group[feature], label=label)
ax.set_title(f'{feature} 散点图 (按诊断分组)')
ax.set_xlabel('样本索引')
ax.set_ylabel(feature)
ax.legend()
else:
# 创建特征对之间的散点图矩阵
numeric_cols = self.data.select_dtypes(include=[np.number]).columns.tolist()[:4] # 限制为4个特征
if numeric_cols:
pd.plotting.scatter_matrix(self.data[numeric_cols], figsize=(10, 10),
diagonal='kde', ax=self.figure.subplots(len(numeric_cols), len(numeric_cols)))
self.figure.suptitle('特征散点图矩阵')
def plot_boxplot(self, feature=None):
if feature:
ax = self.figure.add_subplot(111)
if 'target' in self.data.columns:
sns.boxplot(x='target', y=feature, data=self.data, ax=ax)
ax.set_title(f'{feature} 箱线图 (按诊断分组)')
ax.set_xlabel('心脏病诊断 (0=健康, 1=患病)')
else:
sns.boxplot(y=feature, data=self.data, ax=ax)
ax.set_title(f'{feature} 箱线图')
else:
# 选择数值特征进行箱线图分析
numeric_cols = self.data.select_dtypes(include=[np.number]).columns.tolist()
if 'target' in numeric_cols:
numeric_cols.remove('target')
if len(numeric_cols) > 5:
numeric_cols = numeric_cols[:5] # 限制特征数量
if numeric_cols:
melted_df = pd.melt(self.data, id_vars=['target'] if 'target' in self.data.columns else None,
value_vars=numeric_cols)
ax = self.figure.add_subplot(111)
if 'target' in self.data.columns:
sns.boxplot(x='variable', y='value', hue='target', data=melted_df, ax=ax)
ax.set_title('特征箱线图 (按诊断分组)')
ax.legend(title='诊断', labels=['健康', '患病'])
else:
sns.boxplot(x='variable', y='value', data=melted_df, ax=ax)
ax.set_title('特征箱线图')
ax.set_xlabel('特征')
ax.set_ylabel('值')
def plot_heatmap(self):
numeric_cols = self.data.select_dtypes(include=[np.number]).columns.tolist()
if len(numeric_cols) < 2:
ax = self.figure.add_subplot(111)
ax.text(0.5, 0.5, '需要至少2个数值特征进行热力图分析',
horizontalalignment='center', verticalalignment='center')
return
# 计算相关系数
corr = self.data[numeric_cols].corr()
ax = self.figure.add_subplot(111)
sns.heatmap(corr, annot=True, cmap='coolwarm', fmt='.2f', ax=ax)
ax.set_title('特征相关性热力图')
def plot_distribution(self, feature=None):
if feature:
ax = self.figure.add_subplot(111)
if 'target' in self.data.columns:
for target, group in self.data.groupby('target'):
label = '患病' if target == 1 else '健康'
sns.kdeplot(group[feature], label=label, ax=ax)
ax.set_title(f'{feature} 分布图 (按诊断分组)')
ax.legend()
else:
sns.histplot(self.data[feature], kde=True, ax=ax)
ax.set_title(f'{feature} 分布图')
ax.set_xlabel(feature)
else:
# 选择数值特征进行分布分析
numeric_cols = self.data.select_dtypes(include=[np.number]).columns.tolist()
if 'target' in numeric_cols:
numeric_cols.remove('target')
if len(numeric_cols) > 4:
numeric_cols = numeric_cols[:4] # 限制特征数量
if numeric_cols:
fig, axes = plt.subplots(2, 2, figsize=(10, 8))
axes = axes.flatten()
for i, col in enumerate(numeric_cols[:4]):
if 'target' in self.data.columns:
for target, group in self.data.groupby('target'):
label = '患病' if target == 1 else '健康'
sns.kdeplot(group[col], label=label, ax=axes[i])
axes[i].set_title(f'{col} 分布')
axes[i].legend()
else:
sns.histplot(self.data[col], kde=True, ax=axes[i])
axes[i].set_title(f'{col} 分布')
self.figure = fig
self.canvas.figure = fig
def update_data(self, data):
self.data = data
self.update_charts()
5. 预测结果界面 (ui/prediction_view.py)
from PyQt5.QtWidgets import (QWidget, QVBoxLayout, QHBoxLayout, QLabel,
QProgressBar, QGroupBox, QTextEdit, QPushButton,
QTableWidget, QTableWidgetItem, QHeaderView)
from PyQt5.QtCore import Qt
from PyQt5.QtGui import QFont, QColor
import pandas as pd
import matplotlib.pyplot as plt
from matplotlib.backends.backend_qt5agg import FigureCanvasQTAgg as FigureCanvas
import numpy as np
class PredictionWidget(QWidget):
def __init__(self):
super().__init__()
self.results = {}
self.init_ui()
def init_ui(self):
main_layout = QVBoxLayout
更多推荐


所有评论(0)