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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 09:39:05