174 lines
7.2 KiB
Python
174 lines
7.2 KiB
Python
import numpy as np
|
|
import numpyro
|
|
import jax
|
|
import jax.numpy as jnp
|
|
import numpyro.distributions as dist
|
|
from numpyro.infer import MCMC, NUTS
|
|
import matplotlib.pyplot as plt
|
|
from scipy import stats
|
|
from scipy.stats import gaussian_kde
|
|
from functools import partial # 引入 partial 来固定函数参数
|
|
|
|
# =========================================================================
|
|
# 配置中文字体
|
|
# =========================================================================
|
|
try:
|
|
# Windows 系统优先尝试这些字体
|
|
plt.rcParams['font.sans-serif'] = ['Microsoft YaHei', 'SimHei', 'SimSun', 'KaiTi', 'FangSong', 'Arial Unicode MS']
|
|
plt.rcParams['axes.unicode_minus'] = False # 正常显示负号
|
|
print("✓ 中文字体配置成功")
|
|
except Exception as e:
|
|
print(f"⚠ 字体配置警告: {e}")
|
|
print(" 如果图表中文显示异常,请运行 check_fonts.py 查看可用字体")
|
|
|
|
# 设置随机种子
|
|
numpyro.set_platform("cpu")
|
|
numpyro.set_host_device_count(4)
|
|
|
|
# =========================================================================
|
|
# 1. 参数化的 MCMC 模型
|
|
# =========================================================================
|
|
|
|
def simple_model_factor(y_obs, mu_loc, mu_scale, sigma_a, sigma_b):
|
|
"""
|
|
参数化的正态分布模型
|
|
先验参数作为函数参数传入
|
|
"""
|
|
# 先验分布:mu ~ Normal(mu_loc, mu_scale)
|
|
mu = numpyro.sample("mu", dist.Normal(mu_loc, mu_scale))
|
|
# 先验分布:sigma ~ Beta(sigma_a, sigma_b)
|
|
sigma = numpyro.sample("sigma", dist.Beta(sigma_a, sigma_b))
|
|
|
|
# 使用 numpyro.factor 添加对数似然
|
|
n = len(y_obs)
|
|
log_lik = -0.5 * n * jnp.log(2 * jnp.pi) - n * jnp.log(sigma) - jnp.sum((y_obs - mu) ** 2) / (2 * sigma ** 2)
|
|
numpyro.factor("log_likelihood", log_lik)
|
|
|
|
# =========================================================================
|
|
# 2. 生成模拟数据
|
|
# =========================================================================
|
|
np.random.seed(42)
|
|
true_mu = 2.5 # 真实均值
|
|
true_sigma = 0.2 # 真实标准差
|
|
n_data = 1000
|
|
y_data = np.random.normal(true_mu, true_sigma, n_data)
|
|
y_data_jax = jnp.array(y_data)
|
|
|
|
sample_mean = np.mean(y_data)
|
|
sample_std = np.std(y_data)
|
|
|
|
print(f"真实均值: {true_mu}, 真实标准差: {true_sigma}")
|
|
print(f"样本均值: {sample_mean:.3f}, 样本标准差: {sample_std:.3f}")
|
|
|
|
# =========================================================================
|
|
# 3. 运行两次 MCMC
|
|
# =========================================================================
|
|
|
|
# --- 运行 1: "合理"先验 (Reasonable Prior) ---
|
|
# mu ~ Normal(0, 10), sigma ~ Beta(2, 5) [均值 approx 0.28]
|
|
print("\n--- 正在运行 MCMC (合理先验) ---")
|
|
nuts_kernel_1 = NUTS(
|
|
partial(simple_model_factor, mu_loc=0., mu_scale=10., sigma_a=2., sigma_b=5.)
|
|
)
|
|
mcmc_1 = MCMC(nuts_kernel_1, num_warmup=100, num_samples=300, num_chains=4)
|
|
mcmc_1.run(jax.random.PRNGKey(0), y_obs=y_data_jax)
|
|
mcmc_samples_1 = mcmc_1.get_samples()
|
|
print("✓ MCMC (合理先验) 运行完毕")
|
|
# mcmc_1.print_summary()
|
|
|
|
# --- 运行 2: "错误/远处"先验 (Far Prior) ---
|
|
# mu ~ Normal(10, 1), sigma ~ Beta(5, 2) [均值 approx 0.71]
|
|
print("\n--- 正在运行 MCMC (错误先验) ---")
|
|
nuts_kernel_2 = NUTS(
|
|
partial(simple_model_factor, mu_loc=10., mu_scale=1., sigma_a=5., sigma_b=2.)
|
|
)
|
|
mcmc_2 = MCMC(nuts_kernel_2, num_warmup=1000, num_samples=3000, num_chains=4)
|
|
mcmc_2.run(jax.random.PRNGKey(1), y_obs=y_data_jax)
|
|
mcmc_samples_2 = mcmc_2.get_samples()
|
|
print("✓ MCMC (错误先验) 运行完毕")
|
|
# mcmc_2.print_summary()
|
|
|
|
# =========================================================================
|
|
# 4. 对比绘图 (修改版:使用 KDE 曲线避免遮挡)
|
|
# =========================================================================
|
|
|
|
# --- a. 创建用于绘图的网格 ---
|
|
mu_plot_grid = np.linspace(-5, 15, 400)
|
|
sigma_plot_grid = np.linspace(0.01, 1.0, 400)
|
|
|
|
# --- b. 计算先验的 PDF (不变) ---
|
|
# 合理先验
|
|
prior_mu_1_pdf = stats.norm(0, 10).pdf(mu_plot_grid) # type: ignore
|
|
prior_sigma_1_pdf = stats.beta(2, 5).pdf(sigma_plot_grid) # type: ignore
|
|
# 错误先验
|
|
prior_mu_2_pdf = stats.norm(10, 1).pdf(mu_plot_grid) # type: ignore
|
|
prior_sigma_2_pdf = stats.beta(5, 2).pdf(sigma_plot_grid) # type: ignore
|
|
|
|
# --- c. 【新】计算后验的 KDE (核密度估计) ---
|
|
# 这会根据 MCMC 样本生成平滑的概率密度函数
|
|
print("\n--- 正在计算 KDE (平滑曲线) ---")
|
|
kde_mu_1 = gaussian_kde(mcmc_samples_1['mu'])
|
|
kde_mu_2 = gaussian_kde(mcmc_samples_2['mu'])
|
|
kde_sigma_1 = gaussian_kde(mcmc_samples_1['sigma'])
|
|
kde_sigma_2 = gaussian_kde(mcmc_samples_2['sigma'])
|
|
|
|
# 在网格上计算 KDE 的 PDF 值
|
|
post_mu_1_pdf = kde_mu_1(mu_plot_grid)
|
|
post_mu_2_pdf = kde_mu_2(mu_plot_grid)
|
|
post_sigma_1_pdf = kde_sigma_1(sigma_plot_grid)
|
|
post_sigma_2_pdf = kde_sigma_2(sigma_plot_grid)
|
|
print("✓ KDE 计算完毕")
|
|
|
|
# --- d. 开始绘图 ---
|
|
plt.figure(figsize=(16, 8))
|
|
|
|
# --- 图 1: 对比 mu 的分布 ---
|
|
ax1 = plt.subplot(1, 2, 1)
|
|
# 绘制先验 (虚线, 稍透明)
|
|
ax1.plot(mu_plot_grid, prior_mu_1_pdf, 'b--', label='先验 1 (合理): N(0, 10)', linewidth=2, alpha=0.7)
|
|
ax1.plot(mu_plot_grid, prior_mu_2_pdf, 'r--', label='先验 2 (错误): N(10, 1)', linewidth=2, alpha=0.7)
|
|
|
|
# 绘制后验 (实线/点划线,不透明)
|
|
# 【修改点】用 plot 代替 hist
|
|
ax1.plot(mu_plot_grid, post_mu_1_pdf, 'b-', label='后验 1 (来自合理先验)', linewidth=3)
|
|
ax1.plot(mu_plot_grid, post_mu_2_pdf, 'r-.', label='后验 2 (来自错误先验)', linewidth=3) # 使用不同线型
|
|
|
|
# 绘制真实值和样本均值
|
|
ax1.axvline(true_mu, color='k', linestyle=':', linewidth=2.5, label=f'真实均值 = {true_mu}')
|
|
ax1.axvline(sample_mean, color='gray', linestyle='-', linewidth=2, label=f'样本均值 = {sample_mean:.3f}') # type: ignore
|
|
|
|
ax1.set_title(r"参数 $\mu$ 的先验与后验对比", fontsize=16)
|
|
ax1.set_xlabel(r"$\mu$ 的值")
|
|
ax1.set_ylabel("概率密度")
|
|
ax1.legend(fontsize=10)
|
|
ax1.grid(True, linestyle='--', alpha=0.6)
|
|
ax1.set_xlim(-5, 15)
|
|
ax1.set_ylim(bottom=0) # 确保 y 轴从 0 开始
|
|
|
|
# --- 图 2: 对比 sigma 的分布 ---
|
|
ax2 = plt.subplot(1, 2, 2)
|
|
# 绘制先验 (虚线, 稍透明)
|
|
ax2.plot(sigma_plot_grid, prior_sigma_1_pdf, 'b--', label='先验 1 (合理): Beta(2, 5)', linewidth=2, alpha=0.7)
|
|
ax2.plot(sigma_plot_grid, prior_sigma_2_pdf, 'r--', label='先验 2 (错误): Beta(5, 2)', linewidth=2, alpha=0.7)
|
|
|
|
# 绘制后验 (实线/点划线,不透明)
|
|
# 【修改点】用 plot 代替 hist
|
|
ax2.plot(sigma_plot_grid, post_sigma_1_pdf, 'b-', label='后验 1 (来自合理先验)', linewidth=3)
|
|
ax2.plot(sigma_plot_grid, post_sigma_2_pdf, 'r-.', label='后验 2 (来自错误先验)', linewidth=3) # 使用不同线型
|
|
|
|
# 绘制真实值和样本标准差
|
|
ax2.axvline(true_sigma, color='k', linestyle=':', linewidth=2.5, label=f'真实 $\sigma$ = {true_sigma}')
|
|
ax2.axvline(sample_std, color='gray', linestyle='-', linewidth=2, label=f'样本 $\sigma$ = {sample_std:.3f}') # type: ignore
|
|
|
|
ax2.set_title(r"参数 $\sigma$ 的先验与后验对比", fontsize=16)
|
|
ax2.set_xlabel(r"$\sigma$ 的值")
|
|
ax2.set_ylabel("概率密度")
|
|
ax2.legend(fontsize=10)
|
|
ax2.grid(True, linestyle='--', alpha=0.6)
|
|
ax2.set_xlim(0, 1.0)
|
|
ax2.set_ylim(bottom=0) # 确保 y 轴从 0 开始
|
|
|
|
plt.suptitle("先验信念 vs. 强大数据 (N=1000) [KDE平滑曲线]", fontsize=20, y=1.02)
|
|
plt.tight_layout()
|
|
plt.show()
|