修改代码及README
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user