基于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,不符合的设为0
- 边缘化与归一化:对簇的势能求和得到目标变量的边际概率,再归一化得到后验概率
内容的提问来源于stack exchange,提问作者andrea ricci
相关产品推荐
相关产品推荐

