PyMC 5.1.2中如何单次调用sample_posterior_predictive生成n_samples样本
在PyMC 5中单次调用sample_posterior_predictive生成多个样本
问题描述
PyMC3中sample_posterior_predictive()的samples参数在PyMC 5.1.2中已被移除,现有代码通过for循环多次调用该函数获取预测样本,希望改为单次调用得到形状为(n_samples, 600)的结果。当前环境:matplotlib 3.7.1、numpy 1.24.2、Python 3.11.0。
解决方案
PyMC 5使用draws参数替代旧版本的samples,用于指定单次调用生成的预测样本数量。只需在调用sample_posterior_predictive()时传入draws=n_samples,即可一次性生成指定数量的样本。
修改后的关键代码
替换原有的for循环部分:
n_samples = 4 with model: # 使用draws参数指定生成4个样本 pred_samples = pm.sample_posterior_predictive([mp], var_names=["f_pred"], draws=n_samples)
提取并使用样本
生成的结果中,posterior_predictive["f_pred"]的维度为(chain=1, draw=n_samples, 600),通过以下方式提取得到(4,600)的数组:
fig, ax = plt.subplots(figsize=(12, 5)) # 提取样本并转换为numpy数组 f_preds = pred_samples.posterior_predictive["f_pred"].sel(chain=0).values # 遍历绘制所有样本 for f_pred in f_preds: ax.plot(X_new[:,0], f_pred, alpha=0.1, color='blue') # 绘制原始数据 ax.plot(X, y, "ok", ms=3, alpha=0.5, label="Observed data") plt.show()
完整修改后的代码
import matplotlib.pyplot as plt import numpy as np import pymc as pm # set the seed np.random.seed(1) # create random data n = 50 # The number of data points X = np.linspace(0, np.pi*2, n)[:, None] # The inputs to the GP, they must be arranged as a column vector y = 2*np.sin(0.25*2*np.pi*X[:,0])*X[:,0] + 2 # setup model with pm.Model() as model: ℓ = pm.Gamma("ℓ", alpha=2, beta=1) η = pm.HalfCauchy("η", beta=5) cov = η**2 * pm.gp.cov.Matern52(1, ℓ) gp = pm.gp.Marginal(cov_func=cov) σ = pm.HalfCauchy("σ", beta=5) y_ = gp.marginal_likelihood("y", X=X, y=y, sigma=σ) mp = pm.find_MAP() # new values from x=0 to x=20 X_new = np.linspace(0, 20, 600)[:, None] # add the GP conditional to the model, given the new X values with model: f_pred = gp.conditional("f_pred", X_new) # 单次调用生成n_samples个样本 n_samples = 4 with model: pred_samples = pm.sample_posterior_predictive([mp], var_names=["f_pred"], draws=n_samples) # plot result fig, ax = plt.subplots( figsize=(12, 5)) f_preds = pred_samples.posterior_predictive["f_pred"].sel(chain=0).values for f_pred in f_preds: ax.plot(X_new[:,0], f_pred, alpha=0.1, color = 'blue') # plot the data and the true latent function ax.plot(X, y, "ok", ms=3, alpha=0.5, label="Observed data") plt.show()
内容的提问来源于stack exchange,提问作者sehan2
相关产品推荐
相关产品推荐

