医疗诊断实战:用Python构建贝叶斯网络预测感冒概率(附完整代码)

最近和一位做医疗数据分析的朋友聊天,他提到一个挺有意思的痛点:很多初级的症状筛查模型,因为把各种症状当作独立事件来处理,准确率总是不尽如人意。比如,一个人同时发烧和咳嗽,模型可能只是简单地把两个症状的权重相加,却忽略了“感冒”这个共同的潜在病因,导致预测结果偏离常识。这让我想起了概率图模型中的一个经典工具——贝叶斯网络。它不像黑箱式的深度学习模型,其内部结构清晰,能直观地表达变量间的因果关系,特别适合处理这种存在不确定性和复杂依赖关系的场景,比如医疗诊断。

这篇文章,就是为你——一位对将AI技术应用于实际问题感兴趣的开发者——准备的实战指南。我们将完全从零开始,用Python构建一个用于感冒概率预测的贝叶斯网络。我不会只给你一堆理论公式,而是会手把手带你走过从问题定义、有向无环图(DAG) 构建、条件概率表(CPT) 定义,到最终进行概率推理的完整流程。你会看到每一行代码的作用,理解每一个参数的意义,并最终获得一个可以直接运行、修改和扩展的完整项目。我们的目标很明确:掌握一种能将领域知识(比如医学常识)与数据相结合进行智能推理的实用工具。

1. 从医疗场景出发:理解贝叶斯网络为何有效

在深入代码之前,我们得先搞清楚,为什么贝叶斯网络适合医疗诊断这类问题。想象一下,一位患者来到诊所,主诉是发烧和咳嗽。医生的大脑会飞速运转:有哪些疾病可能同时引起这两种症状?是普通感冒、流感,还是肺炎?每种可能性有多大?医生在做判断时,实际上是在大脑中调用一个复杂的概率模型,这个模型综合了疾病的先验知识(比如冬季流感高发)和当前观察到的证据(症状)。

贝叶斯网络正是将这个过程形式化、计算化的工具。它的核心是两个部分:DAGCPT

  • DAG (有向无环图): 这是网络的“骨架”,用节点和有向箭头描绘了变量之间的因果关系。在我们的感冒诊断模型中,“感冒”是病因,是“因”;“发烧”和“咳嗽”是症状,是“果”。因此,箭头会从“感冒”节点分别指向“发烧”和“咳嗽”节点。这个图结构本身,就编码了我们关于这个问题的领域知识——感冒会导致发烧和咳嗽。
  • CPT (条件概率表): 这是网络的“血肉”,定量地描述了因果关系有多强。例如,它回答了“如果一个人感冒了,他发烧的概率是多少?”这样的问题。CPT将我们的知识(或从数据中学到的规律)转化为具体的数字。

注意:贝叶斯网络中的“因果”方向,通常基于先验知识或时间顺序来设定。虽然从统计上,相关关系不一定意味着因果关系,但在建模时,一个合理的因果假设能极大地提高模型的可解释性和推理效率。

传统的“症状打分”模型,相当于假设所有症状之间是独立的,这显然不符合医学事实(发烧和咳嗽常常同时出现,因为它们有共同的病因)。贝叶斯网络通过条件独立性巧妙地解决了这个问题。在给定父节点“感冒”的状态下,子节点“发烧”和“咳嗽”是条件独立的。这意味着,一旦我们知道患者是否感冒,“发烧”和“咳嗽”之间的相关性就被“解释”掉了。这个特性不仅更符合现实,也极大地简化了联合概率的计算。

2. 搭建项目环境与核心工具选择

工欲善其事,必先利其器。为了高效地构建和操作贝叶斯网络,我们将使用一个非常强大的Python库:pgmpy。它是一个专门用于概率图模型的纯Python库,支持贝叶斯网络和马尔可夫网络的构建、学习和推理,而且代码风格非常清晰。

首先,确保你的Python环境(建议3.8以上)已经就绪,然后我们通过pip安装必要的库。

# 创建并激活一个虚拟环境是个好习惯(可选)
# python -m venv bn_env
# source bn_env/bin/activate  # Linux/Mac
# bn_env\Scripts\activate  # Windows

# 安装核心库 pgmpy
pip install pgmpy

# 安装辅助库,用于后续的可视化和数值计算
pip install networkx matplotlib numpy pandas

安装完成后,可以在Python中导入它们,并简单验证一下。

import pgmpy
print(f"pgmpy版本: {pgmpy.__version__}")

接下来,我们规划一下项目的核心文件结构。一个清晰的结构能让代码更易维护:

medical_bayesian_network/
│
├── model/                  # 模型定义与构建
│   ├── __init__.py
│   └── flu_diagnosis.py    # 核心模型类
│
├── data/                   # 模拟数据或CPT定义
│   └── cpt_definitions.py
│
├── inference/              # 推理相关操作
│   └── query_examples.py
│
├── utils/                  # 工具函数(如可视化)
│   └── visualize.py
│
├── main.py                 # 主程序入口
└── requirements.txt        # 项目依赖

我们的核心工作将主要集中在 model/flu_diagnosis.pymain.py 中。pgmpy 库将承担绝大部分繁重的工作,从模型定义到概率推理。

3. 一步步构建诊断网络:定义DAG与CPT

现在进入核心环节:用代码把我们的医学知识“翻译”成贝叶斯网络。我们将创建一个 FluDiagnosisModel 类来封装所有功能。

3.1 定义变量与网络结构(DAG)

我们考虑一个稍微丰富一点的场景,包含四个变量:

  1. 季节 (Season): 影响感冒的先决因素,取值 {‘冬季’, ‘其他’}。
  2. 感冒 (Flu): 核心的诊断目标,取值 {‘是’, ‘否’}。
  3. 发烧 (Fever): 症状之一,取值 {‘是’, ‘否’}。
  4. 咳嗽 (Cough): 症状之二,取值 {‘是’, ‘否’}。

基于常识,我们定义因果关系:季节会影响感冒的概率(冬季更容易感冒),感冒又会分别引起发烧和咳嗽。因此,我们的DAG结构是:Season -> Flu -> FeverSeason -> Flu -> Cough(实际上Flu是Fever和Cough的直接父节点,Season是Flu的父节点)。

from pgmpy.models import BayesianNetwork
from pgmpy.factors.discrete import TabularCPD

class FluDiagnosisModel:
    def __init__(self):
        # 1. 定义模型结构(边)
        self.model = BayesianNetwork([
            ('Season', 'Flu'),      # 季节影响感冒概率
            ('Flu', 'Fever'),       # 感冒引起发烧
            ('Flu', 'Cough')        # 感冒引起咳嗽
        ])

这几行代码就构建了网络的骨架。BayesianNetwork 对象接收一个边(edge)的列表,每条边是一个二元组 (父节点, 子节点)

3.2 为每个节点注入知识(CPT)

CPT是贝叶斯网络的参数,它需要我们对每个节点进行定义。pgmpy 使用 TabularCPD(表格型条件概率分布)对象来表示CPT。

a) 为根节点‘Season’定义先验概率 根节点没有父节点,其CPT就是它的先验概率分布。假设我们观察到冬季感冒病例约占全年的60%。

# 变量取值顺序很重要,后续查询要一致
season_cpd = TabularCPD(
    variable='Season',          # 节点名称
    variable_card=2,            # 变量取值个数
    values=[[0.4], [0.6]],     # 概率值列表,对应['其他', '冬季']
    state_names={'Season': ['其他', '冬季']} # 明确状态名称
)

values 参数是一个列表的列表。[[0.4], [0.6]] 表示 P(Season=‘其他’)=0.4, P(Season=‘冬季’)=0.6

b) 为‘Flu’节点定义条件概率 ‘Flu’节点有一个父节点‘Season’,因此它的CPT需要定义在给定‘Season’不同取值下,‘Flu’为‘是’或‘否’的概率。

Season (父节点) P(Flu=‘是’ | Season) P(Flu=‘否’ | Season)
其他 0.05 0.95
冬季 0.15 0.85

这个表格的含义是:在非冬季,感冒的先验概率是5%;在冬季,这个概率上升到15%。用代码实现:

flu_cpd = TabularCPD(
    variable='Flu',
    variable_card=2,
    values=[[0.95, 0.85], # P(Flu='否' | Season)
            [0.05, 0.15]], # P(Flu='是’ | Season)
    evidence=['Season'],   # 父节点列表
    evidence_card=[2],     # 每个父节点的取值个数
    state_names={
        'Flu': ['否', '是'],
        'Season': ['其他', '冬季']
    }
)

注意 values 的排列方式:第一行对应 Flu='否',第二行对应 Flu='是'。列的顺序对应父节点状态的所有组合,默认按字母或数字顺序排列。这里两列分别是 (Season='其他')(Season='冬季')

c) 为症状节点‘Fever’和‘Cough’定义条件概率 这两个节点都以‘Flu’为父节点。我们假设感冒后,出现发烧的概率是80%,出现咳嗽的概率是70%;而未感冒时,由于其他原因(如劳累、轻微感染)出现发烧的概率是5%,出现咳嗽的概率是10%。

fever_cpd = TabularCPD(
    variable='Fever',
    variable_card=2,
    values=[[0.95, 0.2],  # P(Fever='否' | Flu)
            [0.05, 0.8]], # P(Fever='是’ | Flu)
    evidence=['Flu'],
    evidence_card=[2],
    state_names={
        'Fever': ['否', '是'],
        'Flu': ['否', '是']
    }
)

cough_cpd = TabularCPD(
    variable='Cough',
    variable_card=2,
    values=[[0.9, 0.3],   # P(Cough='否' | Flu)
            [0.1, 0.7]],  # P(Cough='是’ | Flu)
    evidence=['Flu'],
    evidence_card=[2],
    state_names={
        'Cough': ['否', '是'],
        'Flu': ['否', '是']
    }
)

3.3 将CPD关联到模型并检查

定义好所有CPD后,需要将它们添加到模型中,并检查模型结构的正确性。

    def build_model(self):
        # 将定义好的CPD添加到模型中
        self.model.add_cpds(season_cpd, flu_cpd, fever_cpd, cough_cpd)

        # 检查模型结构是否一致且完整
        assert self.model.check_model(), “模型构建有误,请检查CPD与结构的匹配性!”
        print("模型构建完成,结构正确。")
        return self.model

check_model() 方法会验证:1) 所有节点的CPD是否都已定义;2) CPD的维度(变量数、证据数)是否与网络结构匹配;3) 所有CPD的概率和是否为1。这个断言是保证模型可用的重要防线。

4. 进行概率推理:回答医疗诊断问题

模型建好了,它就是一个“概率计算器”。我们可以通过“提问”(输入证据)来“获得答案”(计算后验概率)。pgmpy 提供了多种推理算法,这里我们使用最常用的变量消除法(VariableElimination)。

4.1 初始化推理引擎

from pgmpy.inference import VariableElimination

class FluDiagnosisModel:
    # ... __init__, build_model 等之前的方法 ...

    def setup_inference(self):
        """初始化推理引擎"""
        self.infer = VariableElimination(self.model)
        print("概率推理引擎准备就绪。")

4.2 执行诊断查询:经典场景分析

让我们模拟几个真实的诊断场景,看看模型如何推理。

场景一:冬季,患者出现发烧和咳嗽,患感冒的概率有多大? 这是最典型的场景,我们拥有了所有信息:季节(证据1)、症状1(证据2)、症状2(证据3)。

    def diagnose_with_all_evidence(self):
        """已知季节和全部症状,诊断感冒概率"""
        # 定义证据:Season='冬季', Fever='是', Cough='是'
        evidence = {'Season': '冬季', 'Fever': '是', 'Cough': '是'}
        
        # 查询目标变量 'Flu' 的后验概率分布
        query_result = self.infer.query(variables=['Flu'], evidence=evidence)
        
        print(f"\n诊断场景:冬季,患者发烧且咳嗽。")
        print(query_result)
        
        # 以更友好的方式输出
        prob_flu = query_result.values[1] # 获取‘Flu=是’的概率值
        print(f"-> 患者患感冒的概率为:{prob_flu:.2%}")
        return prob_flu

运行这段代码,你可能会得到一个很高的概率值(例如95%以上)。模型综合了“冬季高发”的先验信息和“发烧咳嗽”的强相关症状,给出了一个强烈的阳性指示。

场景二:仅知道患者咳嗽,患感冒的概率有多大?(未提供季节和是否发烧) 这种情况也很常见,患者只主诉了咳嗽。

    def diagnose_with_cough_only(self):
        """仅已知咳嗽症状,诊断感冒概率"""
        evidence = {'Cough': '是'}
        query_result = self.infer.query(variables=['Flu'], evidence=evidence)
        
        print(f"\n诊断场景:患者咳嗽(季节、发烧未知)。")
        print(f"-> 患者患感冒的概率为:{query_result.values[1]:.2%}")
        # 同时可以看看此时季节的分布是否受影响(可选)
        season_result = self.infer.query(variables=['Season'], evidence=evidence)
        print(f"-> 此时,季节为冬季的后验概率:{season_result.values[1]:.2%}")

这个概率会比场景一低很多,因为证据较弱。有趣的是,你可能会发现,即使只提供了“咳嗽”这一证据,模型对“季节”的后验概率估计也会发生微妙变化(倾向于冬季),这体现了网络中信息的反向传播。

场景三:如果患者没有感冒,但在冬季出现了发烧,可能是什么原因? 我们可以查询在已知“未感冒”和“冬季”的情况下,“发烧”的概率。这有助于评估症状的“假阳性”率。

    def probability_fever_given_no_flu(self):
        """分析非感冒原因引起发烧的概率"""
        evidence = {'Flu': '否', 'Season': '冬季'}
        query_result = self.infer.query(variables=['Fever'], evidence=evidence)
        
        print(f"\n分析场景:冬季,患者未感冒。")
        print(f"-> 此时患者出现发烧的概率(可能由其他原因引起):{query_result.values[1]:.2%}")

4.3 理解推理结果与模型敏感性

推理输出的 query_result 是一个 DiscreteFactor 对象,其 values 属性包含了目标变量所有状态的概率数组。理解这个输出至关重要。

为了更深入地“感受”模型的行为,我们可以进行一个简单的敏感性分析:改变CPT中的某个关键概率,观察诊断结果如何变化。例如,如果我们将感冒引起咳嗽的概率 P(Cough='是' | Flu='是') 从0.7调整到0.9(即感冒后几乎必然咳嗽),那么在有咳嗽症状时,诊断感冒的概率会如何变化?

我们可以在代码中动态修改CPT并重新推理。这能帮助我们理解模型中哪些参数对结果影响最大,从而在从数据中学习参数或咨询专家时,知道该重点关注哪些信息。

    def sensitivity_analysis(self, new_cough_prob=0.9):
        """敏感性分析:改变‘咳嗽’对‘感冒’的条件概率,观察诊断结果变化"""
        print(f"\n=== 敏感性分析 ===")
        print(f"将 P(Cough='是' | Flu='是') 从 0.7 调整为 {new_cough_prob}")
        
        # 1. 创建模型副本以避免修改原模型
        from copy import deepcopy
        temp_model = deepcopy(self.model)
        
        # 2. 获取原CPD并修改值
        original_cpd = temp_model.get_cpds(node='Cough')
        new_values = original_cpd.values.copy()
        new_values[1, 1] = new_cough_prob  # 修改对应位置的值
        new_values[0, 1] = 1 - new_cough_prob # 确保概率和为1
        
        # 3. 更新CPD并重新进行推理
        temp_model.remove_cpds('Cough')
        new_cough_cpd = TabularCPD(variable='Cough', variable_card=2,
                                    values=new_values,
                                    evidence=['Flu'], evidence_card=[2],
                                    state_names={'Cough': ['否', '是'],
                                                 'Flu': ['否', '是']})
        temp_model.add_cpds(new_cough_cpd)
        
        # 4. 在新模型上执行推理
        temp_infer = VariableElimination(temp_model)
        result = temp_infer.query(variables=['Flu'],
                                   evidence={'Cough': '是', 'Season': '冬季'})
        
        print(f"调整后,冬季咳嗽患者患感冒的概率:{result.values[1]:.2%}")
        print("=== 分析结束 ===\n")

5. 超越基础:模型扩展、评估与部署思考

一个基础的模型跑通了,但要想让它真正实用,我们还需要考虑更多。

5.1 扩展模型复杂度

真实的医疗诊断涉及更多变量。我们可以轻松地扩展这个网络:

  • 更多症状: 加入“喉咙痛”、“流鼻涕”、“全身乏力”等节点。
  • 更多病因: 加入“过敏”、“支气管炎”等竞争性诊断节点,与“感冒”形成“或”关系。
  • 患者属性: 加入“年龄”、“免疫力水平”等作为父节点,影响患病的先验概率。

扩展时,关键是要仔细思考DAG的结构。新增的边代表你认为存在的直接因果关系。结构越复杂,所需的CPT参数就呈指数级增长(这就是所谓的“维度灾难”),定义所有CPT会变得非常困难。这时,通常需要从数据中学习CPT参数,甚至学习网络结构。

5.2 从数据中学习参数

我们之前是手动设定CPT(基于假设或专家知识)。如果有历史诊断数据,我们可以用 pgmpy 来学习CPT参数。假设我们有一个Pandas DataFrame df,其列对应我们的变量(Season, Flu, Fever, Cough),每一行是一个患者记录。

from pgmpy.estimators import BayesianEstimator

# 假设 `df` 是我们的数据框
# 使用贝叶斯估计(可以加入先验平滑,防止零概率问题)
estimator = BayesianEstimator(model=self.model, data=df)
learned_cpds = estimator.get_parameters(prior_type='BDeu', equivalent_sample_size=10) # BDeu先验

for cpd in learned_cpds:
    print(cpd)
    self.model.add_cpds(cpd)

BayesianEstimator 会基于数据和选择的先验分布,计算出最可能的CPT参数。这比手动设定更客观,尤其是当数据量足够大的时候。

5.3 模型验证与局限性

在将任何模型用于实际辅助决策前,验证至关重要。

  • 模拟验证: 就像我们上面做的,输入各种已知证据组合,看输出概率是否符合临床常识。
  • 历史数据回测: 如果有带真实诊断结果的数据,可以将症状作为证据输入模型,计算患病的预测概率,然后与真实标签对比,计算准确率、精确率、召回率等指标。
  • 专家评审: 将模型结构和推理案例展示给领域专家(医生),听取他们的反馈。他们可能会指出缺失的关键变量或不合理的概率设定。

必须清醒认识到这个简单模型的局限性

  1. 简化假设: 我们假设发烧和咳嗽在给定感冒下条件独立,但严重感冒可能同时加剧两种症状,存在残余关联。
  2. CPT的主观性: 手动设定的概率可能不准确。
  3. 静态模型: 无法处理症状的动态变化过程。
  4. 非排他性: 患者可能同时患有感冒和其他疾病,模型目前只考虑了单一病因。

5.4 部署为简易诊断工具

我们可以将这个模型包装成一个简单的命令行或Web工具。以下是一个极简的命令行交互示例:

def run_cli_diagnosis(model_infer):
    """简单的命令行诊断界面"""
    print("\n=== 感冒症状自查助手(基于贝叶斯网络)===")
    print("请回答以下问题(输入y/n):")
    
    try:
        season = input("当前是否是冬季?(y/n): ").strip().lower()
        season_evidence = ‘冬季’ if season == ‘y’ else ‘其他’
        
        fever = input("是否有发烧症状?(y/n): ").strip().lower()
        fever_evidence = ‘是’ if fever == ‘y’ else ‘否’
        
        cough = input("是否有咳嗽症状?(y/n): ").strip().lower()
        cough_evidence = ‘是’ if cough == ‘y’ else ‘否’
        
        evidence = {‘Season’: season_evidence,
                    ‘Fever’: fever_evidence,
                    ‘Cough’: cough_evidence}
        
        result = model_infer.query(variables=[‘Flu’], evidence=evidence)
        prob = result.values[1] * 100
        
        print(f"\n基于您的症状分析:")
        print(f"-> 患感冒的可能性约为 {prob:.1f}%")
        if prob > 70:
            print("-> 可能性较高,建议多休息、补充水分,必要时就医。")
        elif prob > 30:
            print("-> 有一定可能,请密切观察症状变化。")
        else:
            print("-> 可能性较低,但请注意身体其他不适。")
        print("(注:本结果仅为概率模型估算,不能替代专业医疗诊断)")
        
    except Exception as e:
        print(f"输入有误或推理错误:{e}")

main.py 中,我们可以这样整合所有功能:

from model.flu_diagnosis import FluDiagnosisModel

def main():
    print("开始构建医疗诊断贝叶斯网络...")
    diagnosis_model = FluDiagnosisModel()
    diagnosis_model.build_model()
    diagnosis_model.setup_inference()
    
    # 示例推理
    diagnosis_model.diagnose_with_all_evidence()
    diagnosis_model.diagnose_with_cough_only()
    diagnosis_model.probability_fever_given_no_flu()
    
    # 敏感性分析
    diagnosis_model.sensitivity_analysis(new_cough_prob=0.9)
    
    # 启动简易诊断工具
    # run_cli_diagnosis(diagnosis_model.infer)

if __name__ == "__main__":
    main()

写完这些代码并运行,你就能看到一个完整的、从逻辑到代码的贝叶斯网络应用是如何诞生的。它不仅仅是一个预测黑箱,更是一个可以解释、可以调整、可以质疑的透明推理系统。在实际项目中,下一步可能就是接入真实的、脱敏的电子病历数据,用学习算法来优化CPT,或者将模型封装成API,集成到更大的医疗辅助系统中去。这个简单的感冒诊断模型,就像一颗种子,展示了如何用概率的思维和计算工具,去刻画和解决现实世界中充满不确定性的复杂问题。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐