You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于Junction Tree算法的贝叶斯网络信念传播Python实现求助

联合树信念传播实现修正与完整方案

原代码的几个关键问题

  • 节点拼写错误:多处将flu写成flue,导致节点名称不匹配
  • 势能与cluster对应错误:potential_1对应flu+cough,但原代码里绑定到了flu+fatigue的cluster
  • 边的节点列表顺序不统一:比如["cough", "flu"]和cluster的["flu", "cough"]不一致,会导致后续消息传递出错

完整实现代码

下面是修正并完善后的代码,包含联合树的消息传递、证据引入和后验概率计算功能:

# 定义贝叶斯网络结构
network = {
    'flu': {'prob': 0.05},
    'cough': {'prob': 0.1, 'parents': ['flu']},
    'fever': {'prob': 0.3, 'parents': ['flu']},
    'fatigue': {'prob': 0.2, 'parents': ['flu']}
}

# 定义条件概率表(CPT)
flu_cpt = {"flu": {0: 0.95, 1: 0.05}}
cough_cpt = {"flu": {0: {0: 0.9, 1: 0.1}, 1: {0: 0.2, 1: 0.8}}}  # P(cough|flu)
fever_cpt = {"flu": {0: {0: 0.7, 1: 0.3}, 1: {0: 0.5, 1: 0.5}}}  # P(fever|flu)
fatigue_cpt = {"flu": {0: {0: 0.8, 1: 0.2}, 1: {0: 0.6, 1: 0.4}}}  # P(fatigue|flu)

# 构建联合树的簇与势能函数:每个簇的势能是对应变量的联合概率
# 簇(flu, cough)的势能 = P(flu) * P(cough|flu)
cluster_flu_cough_potential = {
    (0, 0): flu_cpt["flu"][0] * cough_cpt["flu"][0][0],
    (0, 1): flu_cpt["flu"][0] * cough_cpt["flu"][0][1],
    (1, 0): flu_cpt["flu"][1] * cough_cpt["flu"][1][0],
    (1, 1): flu_cpt["flu"][1] * cough_cpt["flu"][1][1]
}
# 簇(flu, fever)的势能 = P(flu) * P(fever|flu)
cluster_flu_fever_potential = {
    (0, 0): flu_cpt["flu"][0] * fever_cpt["flu"][0][0],
    (0, 1): flu_cpt["flu"][0] * fever_cpt["flu"][0][1],
    (1, 0): flu_cpt["flu"][1] * fever_cpt["flu"][1][0],
    (1, 1): flu_cpt["flu"][1] * fever_cpt["flu"][1][1]
}
# 簇(flu, fatigue)的势能 = P(flu) * P(fatigue|flu)
cluster_flu_fatigue_potential = {
    (0, 0): flu_cpt["flu"][0] * fatigue_cpt["flu"][0][0],
    (0, 1): flu_cpt["flu"][0] * fatigue_cpt["flu"][0][1],
    (1, 0): flu_cpt["flu"][1] * fatigue_cpt["flu"][1][0],
    (1, 1): flu_cpt["flu"][1] * fatigue_cpt["flu"][1][1]
}

class JunctionTree:
    def __init__(self, clusters, edges):
        # clusters格式: [(变量列表, 势能字典), ...]
        self.clusters = clusters
        # edges格式: [(簇变量列表1, 簇变量列表2), ...]
        self.edges = edges
        # 为每个簇维护邻接表,方便消息传递
        self.adj = {tuple(cluster[0]): [] for cluster in clusters}
        for edge in edges:
            c1 = tuple(edge[0])
            c2 = tuple(edge[1])
            self.adj[c1].append(c2)
            self.adj[c2].append(c1)
        # 存储簇的当前势能(初始为输入的势能,后续会更新)
        self.potentials = {tuple(cluster[0]): cluster[1].copy() for cluster in clusters}
    
    def _get_separator(self, c1, c2):
        # 计算两个簇的分隔符(交集)
        return tuple(set(c1) & set(c2))
    
    def _marginalize(self, potential, variables_to_keep):
        # 边缘化势能函数,保留指定变量,对其他变量求和
        cluster_vars = next(k for k in self.potentials if self.potentials[k] is potential)
        var_indices = [cluster_vars.index(var) for var in variables_to_keep]
        new_potential = {}
        
        for assignment, prob in potential.items():
            # 提取需要保留的变量的取值
            key = tuple(assignment[i] for i in var_indices)
            if key not in new_potential:
                new_potential[key] = 0.0
            new_potential[key] += prob
        return new_potential
    
    def _normalize(self, potential):
        # 归一化势能函数,使其求和为1
        total = sum(potential.values())
        return {k: v/total for k, v in potential.items()}
    
    def _absorb_message(self, from_cluster, to_cluster):
        # 从from_cluster向to_cluster传递消息
        from_vars = from_cluster
        to_vars = to_cluster
        separator = self._get_separator(from_vars, to_vars)
        
        # 对from_cluster的势能边缘化到分隔符,得到消息
        message = self._marginalize(self.potentials[from_vars], separator)
        
        # 将消息乘到to_cluster的势能上
        separator_indices = [to_vars.index(var) for var in separator]
        new_potential = {}
        for assignment, prob in self.potentials[to_vars].items():
            sep_key = tuple(assignment[i] for i in separator_indices)
            new_potential[assignment] = prob * message[sep_key]
        
        self.potentials[to_vars] = new_potential
    
    def _propagate(self, root=None):
        # 消息传递:先向下传递,再向上传递(树的遍历)
        if root is None:
            root = tuple(self.clusters[0][0])
        
        # 深度优先遍历,先处理子节点
        visited = set()
        stack = [(root, None)]
        
        # 向下传递(收集消息)
        while stack:
            node, parent = stack.pop()
            if node in visited:
                continue
            visited.add(node)
            for neighbor in self.adj[node]:
                if neighbor != parent and neighbor not in visited:
                    stack.append((neighbor, node))
        
        # 向上传递(发送消息)
        stack = [(root, None)]
        visited = set()
        while stack:
            node, parent = stack.pop()
            if node in visited:
                continue
            visited.add(node)
            for neighbor in self.adj[node]:
                if neighbor != parent:
                    self._absorb_message(node, neighbor)
                    stack.append((neighbor, node))
    
    def apply_evidence(self, evidence):
        # 引入证据:将观测到的变量取值固定,更新对应簇的势能
        # evidence格式: {变量名: 取值, ...}
        for cluster_vars in self.potentials:
            # 检查当前簇是否包含证据变量
            evidence_in_cluster = [(var, val) for var, val in evidence.items() if var in cluster_vars]
            if not evidence_in_cluster:
                continue
            # 过滤势能,只保留符合证据的赋值
            new_potential = {}
            var_indices = [cluster_vars.index(var) for var, _ in evidence_in_cluster]
            target_vals = [val for _, val in evidence_in_cluster]
            for assignment, prob in self.potentials[cluster_vars].items():
                match = True
                for idx, val in zip(var_indices, target_vals):
                    if assignment[idx] != val:
                        match = False
                        break
                if match:
                    new_potential[assignment] = prob
            self.potentials[cluster_vars] = new_potential
    
    def query(self, query_var):
        # 查询P(query_var | evidence),先确保已经完成消息传递
        # 找到包含查询变量的簇
        target_cluster = None
        for cluster_vars in self.potentials:
            if query_var in cluster_vars:
                target_cluster = cluster_vars
                break
        if not target_cluster:
            raise ValueError(f"查询变量{query_var}不在联合树的任何簇中")
        
        # 边缘化目标簇的势能,只保留查询变量
        marginal = self._marginalize(self.potentials[target_cluster], [query_var])
        # 归一化得到后验概率
        return self._normalize(marginal)

# 构建联合树
junction_tree_clusters = [
    (["flu", "cough"], cluster_flu_cough_potential),
    (["flu", "fever"], cluster_flu_fever_potential),
    (["flu", "fatigue"], cluster_flu_fatigue_potential),
]
junction_tree_edges = [
    (["flu", "cough"], ["flu", "fever"]),
    (["flu", "fever"], ["flu", "fatigue"]),
]
jt = JunctionTree(junction_tree_clusters, junction_tree_edges)

# ------------------- 示例使用 -------------------
# 示例1:查询先验概率P(flu)
print("先验概率P(flu):")
jt_prior = JunctionTree(junction_tree_clusters, junction_tree_edges)
prior_result = jt_prior.query("flu")
print(f"flu=0: {prior_result[(0):.4f]}, flu=1: {prior_result[(1):.4f]}")

# 示例2:查询条件概率P(flu | cough=1)
print("\n条件概率P(flu | cough=1):")
jt_evidence = JunctionTree(junction_tree_clusters, junction_tree_edges)
jt_evidence.apply_evidence({"cough": 1})
jt_evidence._propagate()
posterior_result = jt_evidence.query("flu")
print(f"flu=0: {posterior_result[(0):.4f]}, flu=1: {posterior_result[(1):.4f]}")

# 示例3:查询条件概率P(cough | fever=1, fatigue=1)
print("\n条件概率P(cough | fever=1, fatigue=1):")
jt_multi_evidence = JunctionTree(junction_tree_clusters, junction_tree_edges)
jt_multi_evidence.apply_evidence({"fever": 1, "fatigue": 1})
jt_multi_evidence._propagate()
cough_posterior = jt_multi_evidence.query("cough")
print(f"cough=0: {cough_posterior[(0):.4f]}, cough=1: {cough_posterior[(1):.4f]}")

关键功能说明

  1. 消息传递:通过深度优先遍历实现联合树的双向消息传递,确保所有簇都能接收到来自其他部分的证据信息
  2. 证据引入:过滤簇的势能函数,只保留符合观测证据的变量赋值,相当于将证据的概率设为1,不符合的设为0
  3. 边缘化与归一化:对簇的势能求和得到目标变量的边际概率,再归一化得到后验概率

内容的提问来源于stack exchange,提问作者andrea ricci

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.23 12:59:55