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

pgmpy中BayesianModelSampling采样触发IndexError问题咨询

问题

在使用pgmpy进行贝叶斯网络采样时,触发IndexError错误,提示“数组为1维,但使用了3个索引进行访问”,错误栈如下:

File "c:\Users\a-rotalintiy\PhD\correlations\import pandas as pd.py", line 38, in <module>
    synthetic_data = sampler.forward_sample(size=synthetic_data_size)
                     ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\a-rotalintiy\.virtualenvs\correlations-k4Gxn6_V\Lib\site-packages\pgmpy\sampling\Sampling.py", line 125, in forward_sample
    sampled[node] = sample_discrete_maps(
                    ^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\a-rotalintiy\.virtualenvs\correlations-k4Gxn6_V\Lib\site-packages\pgmpy\utils\mathext.py", line 194, in sample_discrete_maps
    samples[(weight_indices == weight_index)] = np.random.choice(
    ~~~~~~~^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
IndexError: too many indices for array: array is 1-dimensional, but 3 were indexed

运行代码如下:

import pandas as pd
from pgmpy.models import BayesianNetwork
from pgmpy.estimators import HillClimbSearch, BicScore, BayesianEstimator
from pgmpy.sampling import BayesianModelSampling

# Sample DataFrame
data = {
    'A': [1, 0, 1, 0, 1],
    'B': [0, 1, 0, 1, 0],
    'C': [1, 1, 0, 0, 1]
}
df = pd.DataFrame(data)

# Ensure the DataFrame is correctly structured
print("DataFrame shape:", df.shape)
print("DataFrame head:\n", df.head())

# Estimate the model structure
hc = HillClimbSearch(df)
scoring_method = BicScore(df)
best_model = hc.estimate(scoring_method=scoring_method)

# Print learned edges
print("Learned edges:", best_model.edges())

# Create and fit the Bayesian model
bn_model = BayesianNetwork(best_model.edges())
bn_model.fit(df, estimator=BayesianEstimator, prior_type="BDeu")

# Ensure the model is fitted correctly
for cpd in bn_model.get_cpds():
    print(f"CPD of {cpd.variable}:")
    print(cpd)

# Sample synthetic data
sampler = BayesianModelSampling(bn_model)
synthetic_data_size = len(df)  # You can adjust this size as needed
synthetic_data = sampler.forward_sample(size=synthetic_data_size)

print("Synthetic Data Sample:\n", synthetic_data)
错误原因及解决方案

核心原因

该错误是pgmpy版本与numpy版本不兼容导致的:旧版pgmpy(如0.1.20及更早)的mathext.py中,sample_discrete_maps函数的数组索引逻辑存在问题,当配合新版本numpy使用时,会触发维度不匹配的索引错误。

可行解决方案

  • 升级pgmpy到最新稳定版
    pgmpy后续版本已修复该索引问题,直接升级即可解决:

    pip install --upgrade pgmpy
    

    升级后无需修改原有代码,采样逻辑可正常运行。

  • 降级numpy至兼容版本
    若无法升级pgmpy,可安装与旧版pgmpy兼容的numpy版本(如1.23.5):

    pip install numpy==1.23.5
    
  • 临时代码补丁(不推荐)
    手动修改虚拟环境中pgmpy/utils/mathext.py的第194行,将samples[(weight_indices == weight_index)]改为samples[weight_indices == weight_index](移除多余括号),但该方式维护性差,仅适合临时应急。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 18:05:15