使用az.from_pyjags触发ValueError:多维变量导致解包失败
解决ArviZ
az.from_pyjags()处理高维变量时的ValueError问题 我在调用az.from_pyjags(tr)时触发了如下错误:
--------------------------------------------------------------------------- ValueError Traceback (most recent call last) /home/anaconda/workspace/group_code/long_rt/simulation1/jags_test.ipynb Cell 12' in <cell line: 1>() ----> 1 az.from_pyjags(tr) File ~/anaconda3/envs/mcmc/lib/python3.10/site-packages/arviz/data/io_pyjags.py:374, in from_pyjags(posterior, prior, log_likelihood, coords, dims, save_warmup, warmup_iterations) 313 def from_pyjags( 314 posterior: tp.Optional[tp.Mapping[str, np.ndarray]] = None, 315 prior: tp.Optional[tp.Mapping[str, np.ndarray]] = None, (...) 320 warmup_iterations: int = 0, 321 ) -> InferenceData: 322 """ 323 Convert PyJAGS posterior samples to an ArviZ inference data object. 324 (...) 364 InferenceData 365 """ 366 return PyJAGSConverter( 367 posterior=posterior, 368 prior=prior, 369 log_likelihood=log_likelihood, 370 dims=dims, 371 coords=coords, 372 save_warmup=save_warmup, 373 warmup_iterations=warmup_iterations, --> 374 ).to_inference_data() File ~/anaconda3/envs/mcmc/lib/python3.10/site-packages/arviz/data/io_pyjags.py:107, in PyJAGSConverter.to_inference_data(self) 103 save_warmup = self.save_warmup and self.warmup_iterations > 0 104 # self.posterior is not None 106 idata_dict = { --> 107 "posterior": self.posterior_to_xarray(), 108 "prior": self.prior_to_xarray(), 109 "log_likelihood": self.log_likelihood_to_xarray(), 110 "save_warmup": save_warmup, 111 } 113 return InferenceData(**idata_dict) File ~/anaconda3/envs/mcmc/lib/python3.10/site-packages/arviz/data/io_pyjags.py:83, in PyJAGSConverter.posterior_to_xarray(self) 80 if self.posterior is None: 81 return None --> 83 return self._pyjags_samples_to_xarray(self.posterior) File ~/anaconda3/envs/mcmc/lib/python3.10/site-packages/arviz/data/io_pyjags.py:62, in PyJAGSConverter._pyjags_samples_to_xarray(self, pyjags_samples) 59 def _pyjags_samples_to_xarray( 60 self, pyjags_samples: tp.Mapping[str, np.ndarray] 61 ) -> tp.Tuple[xarray.Dataset, xarray.Dataset]: --> 62 data, data_warmup = get_draws( 63 pyjags_samples=pyjags_samples, 64 warmup_iterations=self.warmup_iterations, 65 warmup=self.save_warmup, 66 ) 68 return ( 69 dict_to_dataset(data, library=self.pyjags, coords=self.coords, dims=self.dims), 70 dict_to_dataset( (...) 75 ), 76 ) File ~/anaconda3/envs/mcmc/lib/python3.10/site-packages/arviz/data/io_pyjags.py:165, in get_draws(pyjags_samples, variables, warmup, warmup_iterations) 161 data_warmup = _convert_pyjags_dict_to_arviz_dict( 162 samples=warmup_samples, variable_names=variables 163 ) 164 else: --> 165 data = _convert_pyjags_dict_to_arviz_dict(samples=pyjags_samples, variable_names=variables) 167 return data, data_warmup File ~/anaconda3/envs/mcmc/lib/python3.10/site-packages/arviz/data/io_pyjags.py:238, in _convert_pyjags_dict_to_arviz_dict(samples, variable_names) 236 for variable_name, chains in samples.items(): 237 if variable_name in variable_names: --> 238 parameter_dimension, _, _ = chains.shape 239 if parameter_dimension == 1: 240 variable_name_to_samples_map[variable_name] = chains[0, :, :].transpose() ValueError: too many values to unpack (expected 3)
经排查,该错误由模型中变量维度超过2维触发:ArviZ的PyJAGS转换代码硬假设了采样结果的数组维度为3,但高维变量会生成维度数更多的数组,导致解包失败。
临时修复:修改ArviZ源码
找到ArviZ安装目录下的data/io_pyjags.py文件,定位到_convert_pyjags_dict_to_arviz_dict函数(约第236-240行),替换原有维度处理逻辑:
原代码
parameter_dimension, _, _ = chains.shape if parameter_dimension == 1: variable_name_to_samples_map[variable_name] = chains[0, :, :].transpose()
替换为兼容高维的代码
import numpy as np shape_len = len(chains.shape) parameter_dimension = chains.shape[0] if parameter_dimension == 1: if shape_len == 3: variable_name_to_samples_map[variable_name] = chains[0, :, :].transpose() else: # 高维变量:将链和迭代维度移到最前面 variable_name_to_samples_map[variable_name] = np.moveaxis(chains[0, ...], -1, 0) else: # 多参数维度变量同样调整维度顺序 variable_name_to_samples_map[variable_name] = np.moveaxis(chains, -1, 0)
这段代码的核心是:
- 不再固定假设数组维度数,动态判断长度
- 使用
np.moveaxis替代硬编码转置,将链(倒数第二维)和迭代数(最后一维)调整为前两维,适配ArviZ的预期格式 - 兼容标量、一维及任意高维变量
长期方案:手动转换采样格式
不想修改源码的话,可以先手动调整PyJAGS采样结果的维度顺序,再传入az.from_pyjags:
import numpy as np import arviz as az def fix_pyjags_samples(tr): fixed_tr = {} for var_name, chains in tr.items(): # 将原格式(param_dim, ..., n_chains, n_draws)转为(n_chains, n_draws, param_dim, ...) fixed_tr[var_name] = np.moveaxis(chains, (-2, -1), (0, 1)) return fixed_tr # 修复后再生成InferenceData fixed_tr = fix_pyjags_samples(tr) idata = az.from_pyjags(fixed_tr)
内容的提问来源于stack exchange,提问作者qipeng chen
相关产品推荐
相关产品推荐

