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
相关产品推荐
相关产品推荐

