修改代码及README

This commit is contained in:
2025-11-02 20:47:28 +08:00
parent 64ca5e257a
commit b6a9feddc7
14 changed files with 1785 additions and 236 deletions
+145 -99
View File
@@ -1,127 +1,173 @@
# JAX 入门示例代码
# 演示 JAX 的基本用法,包括 jnp 数组、jit 加速
import numpy as np
import numpyro
import jax
import jax.numpy as jnp
from jax import jit, grad, vmap
import time
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 来固定函数参数
# --- 1. 像 NumPy 一样简单 ---
print("--- 1. NumPy-like API ---")
key = jax.random.PRNGKey(0) # JAX 处理随机数需要一个密钥
x_jnp = jnp.arange(10)
y_jnp = jax.random.normal(key, (10,))
# =========================================================================
# 配置中文字体
# =========================================================================
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 查看可用字体")
print(f"JAX 数组 x: {x_jnp}")
print(f"JAX 数组 y (随机): {y_jnp}")
print(f"x 和 y 的点积: {jnp.dot(x_jnp, y_jnp)}")
# 设置随机种子
numpyro.set_platform("cpu")
numpyro.set_host_device_count(4)
# =========================================================================
# 1. 参数化的 MCMC 模型
# =========================================================================
# --- 2. jax.jit: 极致加速 ---
print("\n--- 2. @jit 加速 ---")
# 定义一个包含大量运算的函数
def slow_function(x):
# 这是一个比较耗时的操作 (矩阵乘法循环)
for _ in range(100):
x = jnp.dot(x, x.T)
return x
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)
# 创建 JIT 编译版本的函数
fast_function = jit(slow_function)
# =========================================================================
# 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)
# 创建数据
big_matrix = jax.random.normal(key, (200, 200))
sample_mean = np.mean(y_data)
sample_std = np.std(y_data)
# --- 计时对比 ---
# a) 运行普通 NumPy/Python 版本
start_time = time.time()
slow_function(big_matrix).block_until_ready() # .block_until_ready() 确保 JAX 计算完成
numpy_time = time.time() - start_time
print(f"普通 NumPy/Python 版本耗时: {numpy_time:.6f}")
print(f"真实均值: {true_mu}, 真实标准差: {true_sigma}")
print(f"样本均值: {sample_mean:.3f}, 样本标准差: {sample_std:.3f}")
# b) 运行 JIT 编译版本
# 第一次运行会包含编译时间
start_time = time.time()
fast_function(big_matrix).block_until_ready()
compile_time = time.time() - start_time
print(f"JIT 版本 (首次,含编译) 耗时: {compile_time:.6f}")
# =========================================================================
# 3. 运行两次 MCMC
# =========================================================================
# 第二次运行,将只体现执行速度
start_time = time.time()
fast_function(big_matrix).block_until_ready()
jit_time = time.time() - start_time
print(f"JIT 版本 (第二次) 耗时: {jit_time:.6f}")
# --- 运行 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()
print(f"加速比 (第二次运行 vs 普通版本): {numpy_time / jit_time:.2f}")
# --- 运行 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 曲线避免遮挡)
# =========================================================================
# --- 3. jax.grad: 自动求导 ---
print("\n--- 3. grad 自动求导 ---")
# 定义一个简单的函数 f(x) = x^3 + 2x^2 + 5
def my_func(x):
return x**3 + 2*x**2 + 5
# --- a. 创建用于绘图的网格 ---
mu_plot_grid = np.linspace(-5, 15, 400)
sigma_plot_grid = np.linspace(0.01, 1.0, 400)
# 使用 grad 创建一个计算 my_func 导数的函数
# f'(x) = 3x^2 + 4x
grad_my_func = grad(my_func)
# --- 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
x_val = 2.0
derivative = grad_my_func(x_val)
expected_derivative = 3 * x_val**2 + 4 * x_val
# --- 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'])
print(f"函数 f(x) = x^3 + 2x^2 + 5")
print(f"在 x = {x_val} 处的导数是: {derivative}")
print(f"理论上的导数值是: {expected_derivative}")
# 在网格上计算 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))
# --- 4. jax.vmap: 自动向量化 ---
print("\n--- 4. vmap 自动向量化 ---")
# 定义一个只能处理单个向量的函数 (向量点积)
def single_dot_product(a, b):
return jnp.dot(a, b)
# --- 图 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)
# 创建一批 (batch) 数据
# 假设我们有 5 对向量,每对向量长度为 3
batch_a = jnp.arange(15).reshape(5, 3)
batch_b = jnp.arange(15, 30).reshape(5, 3)
# 绘制后验 (实线/点划线,不透明)
# 【修改点】用 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) # 使用不同线型
# 使用 vmap 将函数向量化
# in_axes=(0, 0) 表示对 a 和 b 的第 0 轴 (批次轴) 进行映射
# out_axes=0 表示输出结果也沿着第 0 轴堆叠
batch_dot_product = vmap(single_dot_product, in_axes=(0, 0), out_axes=0)
# 绘制真实值和样本均值
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
results = batch_dot_product(batch_a, batch_b)
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 开始
print("批次数据 a:\n", batch_a)
print("批次数据 b:\n", batch_b)
print("vmap 向量化计算的点积结果:\n", results)
# --- 图 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)
# 对比手动 for 循环
manual_results = jnp.array([single_dot_product(a, b) for a, b in zip(batch_a, batch_b)])
print("手动 for 循环的结果:\n", manual_results)
print(f"vmap 结果与手动循环结果是否一致: {jnp.allclose(results, manual_results)}")
# 绘制后验 (实线/点划线,不透明)
# 【修改点】用 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
# --- 5. JAX 的随机数处理 ---
print("\n--- 5. JAX 的随机数 ---")
# 1. 创建一个初始密钥
key = jax.random.PRNGKey(42)
print(f"初始密钥: {key}")
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 开始
# 2. 使用密钥生成随机数
random_data_1 = jax.random.normal(key, (3,))
print(f"第一次生成的随机数据: {random_data_1}")
# 3. 再次使用同一个密钥,会得到完全相同的结果!
random_data_2 = jax.random.normal(key, (3,))
print(f"第二次使用相同密钥生成的数据: {random_data_2}")
# 4. 正确的做法:分割密钥
key, subkey = jax.random.split(key) # key 更新为新的主密钥, subkey 用于本次操作
random_data_3 = jax.random.normal(subkey, (3,))
print(f"\n分割密钥后,第一次生成的随机数据: {random_data_3}")
key, subkey = jax.random.split(key) # 再次分割
random_data_4 = jax.random.normal(subkey, (3,))
print(f"分割密钥后,第二次生成的随机数据: {random_data_4}")
plt.suptitle("先验信念 vs. 强大数据 (N=1000) [KDE平滑曲线]", fontsize=20, y=1.02)
plt.tight_layout()
plt.show()