Pyro如何指定多离散父节点对应的连续子节点条件概率分布
Pyro实现离散父节点下连续子节点的贝叶斯网络建模
核心逻辑
你的三个离散父节点各有10种状态,总共有 10*10*10=1000 种组合,我们为每一种组合单独配置正态分布的均值mu和标准差sigma参数,为这两类参数设置先验后,即可通过父节点的状态索引对应分布参数,完成条件分布P(D|A,B,C)的定义。
注意你现有代码中定义的A/B/C是Dirichlet先验(对应离散类别变量的分布参数),还需要采样得到实际的离散类别值,下面的实现会补全这部分逻辑。
完整模型代码
import torch import pyro import pyro.distributions as dist from pyro.infer import MCMC, NUTS def model(data): # data为字典,包含观测的D,以及可选的A/B/C离散索引观测 # 1. 定义离散父节点的先验与采样 # 1.1 Dirichlet先验对应各状态的概率 prob_A = pyro.sample("prob_A", dist.Dirichlet(torch.ones(10))) prob_B = pyro.sample("prob_B", dist.Dirichlet(torch.ones(10))) prob_C = pyro.sample("prob_C", dist.Dirichlet(torch.ones(10))) # 1.2 采样离散类别值,若A/B/C是观测值可直接从data取索引,删掉这三行即可 with pyro.plate("obs_plate", len(data["D"])): A = pyro.sample("A", dist.Categorical(probs=prob_A), obs=data.get("A", None)) B = pyro.sample("B", dist.Categorical(probs=prob_B), obs=data.get("B", None)) C = pyro.sample("C", dist.Categorical(probs=prob_C), obs=data.get("C", None)) # 2. 定义P(D|A,B,C)的正态分布参数先验 # 10x10x10的张量,对应A/B/C各10种状态的所有组合 mu = pyro.sample("mu", dist.Normal(loc=0, scale=10).expand([10,10,10]).to_event(3)) sigma = pyro.sample("sigma", dist.HalfNormal(scale=10).expand([10,10,10]).to_event(3)) # 3. 索引对应组合的正态参数,采样子节点D with pyro.plate("obs_plate", len(data["D"])): # 按A/B/C的索引取对应的分布参数 current_mu = mu[A.long(), B.long(), C.long()] current_sigma = sigma[A.long(), B.long(), C.long()] # 观测节点D D = pyro.sample("D", dist.Normal(current_mu, current_sigma), obs=data["D"])
HMC推理示例
你可以使用Pyro内置的NUTS内核(HMC的优化实现)进行后验推断,示例代码如下:
# 模拟测试数据,替换为你的真实数据 num_samples = 1000 test_data = { "A": torch.randint(0,10,(num_samples,)), "B": torch.randint(0,10,(num_samples,)), "C": torch.randint(0,10,(num_samples,)), "D": torch.randn(num_samples) # 替换为真实的D观测值 } # 初始化NUTS内核 nuts_kernel = NUTS(model) # 初始化MCMC,设置采样数和预热步数 mcmc = MCMC( nuts_kernel, num_warmup=500, num_samples=1000, disable_progbar=False ) # 运行推断 mcmc.run(test_data) # 打印后验统计结果 mcmc.summary()
优化建议
- 可根据D的实际取值范围调整
mu和sigma的先验参数,避免先验分布与实际数据偏差过大 - 若你的观测数据量不足,可给
mu设置层次先验缩减参数空间,降低过拟合风险 - 如果A/B/C是已知的观测变量,不需要推断其分布,可直接删除
prob_A/prob_B/prob_C的Dirichlet采样逻辑,直接使用data中的A/B/C索引即可
内容的提问来源于stack exchange,提问作者Luca
相关产品推荐
相关产品推荐

