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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 22:33:22