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

使用PyMC3实现含3个分类变量的基础贝叶斯网络并推理

解决简化贝叶斯网络的搭建与后验推理问题

我来帮你搞定这个仅含3个分类变量的贝叶斯网络搭建,以及固定D3=0后推断D1、D2后验分布的问题!咱们先把代码补全并梳理清楚每一步的逻辑:

第一步:补全概率数组并明确变量关系

首先你给出的d3_prob没写完,它需要是一个2×3×2的数组:第一维度对应D1的2种取值(0/1),第二维度对应D2的3种取值(0/1/2),第三维度对应D3的2种取值(0/1),每个元素是对应条件下的概率值。我先给你补一个完整示例:

import pymc3 as pm
import numpy as np
import arviz as az

# 先验概率:D1有2个取值,D2有3个取值
d1_prob = np.array([0.3, 0.7])  # P(D1=0)=0.3, P(D1=1)=0.7
d2_prob = np.array([0.6, 0.3, 0.1])  # P(D2=0)=0.6, P(D2=1)=0.3, P(D2=2)=0.1
# 条件概率P(D3|D1,D2):维度[D1取值数, D2取值数, D3取值数]
d3_prob = np.array([
    [[0.1, 0.9],  # D1=0, D2=0时,P(D3=0)=0.1, P(D3=1)=0.9
     [0.3, 0.7],  # D1=0, D2=1时,P(D3=0)=0.3, P(D3=1)=0.7
     [0.4, 0.6]], # D1=0, D2=2时,P(D3=0)=0.4, P(D3=1)=0.6
    [[0.8, 0.2],  # D1=1, D2=0时,P(D3=0)=0.8, P(D3=1)=0.2
     [0.5, 0.5],  # D1=1, D2=1时,P(D3=0)=0.5, P(D3=1)=0.5
     [0.2, 0.8]]  # D1=1, D2=2时,P(D3=0)=0.2, P(D3=1)=0.8
])

第二步:搭建贝叶斯网络模型

在PyMC3中,分类变量用Categorical分布定义,固定D3=0只需要给D3设置observed=0即可:

with pm.Model() as bayes_net:
    # 定义先验变量D1和D2
    D1 = pm.Categorical('D1', p=d1_prob)
    D2 = pm.Categorical('D2', p=d2_prob)
    
    # 定义条件变量D3:根据D1和D2的取值索引获取对应的条件概率
    # 这里用pm.math.subtensor来索引多维概率数组
    d3_cond_prob = pm.math.subtensor(d3_prob, D1, D2)
    D3 = pm.Categorical('D3', p=d3_cond_prob, observed=0)  # 固定D3=0
    
    # 采样:离散变量推荐用Metropolis采样器,NUTS对离散变量支持不佳
    trace = pm.sample(5000, tune=2000, chains=2, cores=2, step=pm.Metropolis())

第三步:分析后验分布

采样完成后,我们可以统计D1和D2各取值的后验概率,或者用ArviZ可视化:

# 查看后验采样的摘要统计
az.summary(trace, kind='stats')

# 统计D1各取值的后验概率
d1_posterior = np.mean(trace['D1'] == 0), np.mean(trace['D1'] == 1)
print(f"D1的后验概率:P(D1=0)={d1_posterior[0]:.3f}, P(D1=1)={d1_posterior[1]:.3f}")

# 统计D2各取值的后验概率
d2_posterior = np.mean(trace['D2'] == 0), np.mean(trace['D2'] == 1), np.mean(trace['D2'] == 2)
print(f"D2的后验概率:P(D2=0)={d2_posterior[0]:.3f}, P(D2=1)={d2_posterior[1]:.3f}, P(D2=2)={d2_posterior[2]:.3f}")

# 可视化后验分布
az.plot_posterior(trace)

关键注意事项

  • 索引匹配:PyMC3的Categorical变量取值从0开始,所以你的概率数组维度和索引必须严格对应,否则会取错条件概率。
  • 采样器选择:因为D1、D2都是离散变量,NUTS采样器(PyMC3默认)对离散变量支持不好,所以一定要指定用Metropolis采样器。
  • 样本量:采样时要设置足够的tune(热身样本)和采样次数,确保后验收敛(可以用az.plot_trace(trace)查看收敛情况)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 08:59:03