PyMC新版本中获取采样值的语法问题求助
PyMC v4+ 获取采样值的正确语法
在PyMC v4及以上版本中,pm.sample()返回的是ArviZ InferenceData对象,而非旧版的MultiTrace,因此原有的trace.get_values()方法已被移除。以下是获取变量采样值的正确方式:
核心访问方式
直接通过trace.posterior访问后验采样结果(这是一个xarray DataSet),结合xarray方法提取数据:
1. 获取完整变量采样(xarray DataArray)
# 获取phi的所有采样(包含链、采样步维度) phi_data = trace.posterior['phi'] # 获取sigma的所有采样 sigma_data = trace.posterior['sigma']
2. 转换为NumPy数组
使用.to_numpy()将xarray对象转为NumPy数组:
phi_vals = phi_data.to_numpy() sigma_vals = sigma_data.to_numpy()
3. 提取多维度变量的子参数(如phi1、phi2)
如果phi是包含两个参数的变量(比如形状为(chain, draw, 2)),可以通过维度索引提取单个参数:
# 合并chain和draw维度为sample,再提取第0个参数(phi1) phi1_vals = trace.posterior['phi'].stack(sample=('chain', 'draw'))[:, 0].to_numpy() # 提取第1个参数(phi2) phi2_vals = trace.posterior['phi'].stack(sample=('chain', 'draw'))[:, 1].to_numpy()
4. 使用ArviZ的extract简化操作
ArviZ的az.extract()可以直接合并链和采样步,返回扁平化的采样结果:
import arviz as az # combined=True 合并所有链的采样 posterior_samples = az.extract(trace, combined=True) phi1_vals = posterior_samples['phi'].values[:, 0] phi2_vals = posterior_samples['phi'].values[:, 1] sigma_vals = posterior_samples['sigma'].values
示例对比
旧代码(PyMC v3及更早)
phi_vals = trace.get_values('phi') sigma_vals = trace.get_values('sigma') phi1_vals = phi_vals[:, 0] phi2_vals = phi_vals[:, 1]
新代码(PyMC v4+)
import arviz as az trace = pm.sample(...) posterior = az.extract(trace, combined=True) phi1_vals = posterior['phi'].values[:, 0] phi2_vals = posterior['phi'].values[:, 1] sigma_vals = posterior['sigma'].values
内容的提问来源于stack exchange,提问作者Fernando DePaolis
相关产品推荐
相关产品推荐

