From 6694450c70ad061086f6f3c5754d624f693e439f Mon Sep 17 00:00:00 2001 From: Hongru Date: Mon, 20 Oct 2025 09:53:02 +0800 Subject: [PATCH] initial commit --- demo.py | 105 ++++++++++++++++++++++++++ demo1.py | 124 ++++++++++++++++++++++++++++++ demo2.py | 224 +++++++++++++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 453 insertions(+) create mode 100644 demo.py create mode 100644 demo1.py create mode 100644 demo2.py diff --git a/demo.py b/demo.py new file mode 100644 index 0000000..b74d4f1 --- /dev/null +++ b/demo.py @@ -0,0 +1,105 @@ +## 实现一下重要抽样 +import numpy as np +import matplotlib.pyplot as plt +import scipy as stats +import seaborn as sns + +np.random.seed(42) + +def baysian_linear_regression_inportance_sampling(): + """ + 使用贝叶斯线性回归方法进行重要抽样 + + 本例需要估计一个简单的线性回归模型的后验分布 + + y = β0 + β1*x + ε, ε ~ N(0, σ²) + = [1 x] * [β0 β1]' + ε + """ + + # 1. 定义关键数据 + x = np.array([1, 2, 3, 4, 5, 6, 7, 8, 9, 10]) + y = np.array([2.1, 3.9, 6.2, 8.1, 9.8, 12.3, 13.9, 16.2, 17.8, 20.1]) + n = len(x) + + print(f"数据点数量:{n}") + print(f"x: {x}") + print(f"y: {y}") + + # 2. 定义先验分布 + # β0 ~ N(0, 10) + beta0_prior_mean = 0 + beta0_prior_var = 10 + + # β1 ~ N(1, 5) + beta1_prior_mean = 1 + beta1_prior_var = 5 + + # σ² ~ Inverse-Gamma(2, 1) + sigma2_prior_a = 2 + sigma2_prior_b = 1 + + # 3. 定义重要抽样分布 + # 使用最小二乘进行估计 + X = np.column_stack([np.ones(n), x]) + beta_estimate = np.linalg.inv(X.T @ X) @ X.T @ y + y_pred = X @ beta_estimate + residuals = y - y_pred + RSS = residuals.T @ residuals + sigma2_estimate = RSS / (n - 2) + + print(f"最小二乘结果:β0 = {beta_estimate[0]}, β1 = {beta_estimate[1]}, σ² = {sigma2_estimate}") + print(f"残差平方和 RSS = {RSS}") + + + + # 4. 定义联合后验分布(未归一化) + def unnormalized_log_posterior(param): + """ + 计算后验概率密度的未归一化值 + + param: [β0, β1, log_sigma2)] + """ + + beta0, beta1, log_sigma2 = param + + sigma2 = np.exp(log_sigma2) + + # 计算似然函数值 + y_pred = beta0 + beta1 * x + residuals = y - y_pred + log_likelihood = -0.5 * n * np.log(2 * np.pi * sigma2) * -0.5 * np.sum(residuals**2) / sigma2 + + # 计算先验概率 + log_prior_beta0 = stats.norm.logpdf(beta0, loc = beta0_prior_mean, scale = np.sqrt(beta0_prior_var)) + log_prior_beta1 = stats.norm.logpdf(beta1, loc = beta1_prior_mean, scale = np.sqrt(beta1_prior_var)) + log_prior_sigma2 = stats.invgamma.logpdf(sigma2, a = sigma2_prior_a, scale = sigma2_prior_b) + log_sigma2 + + return log_likelihood + log_prior_beta0 + log_prior_beta1 + log_prior_sigma2 + + # 5. 定义重要抽样分布 + proposal_mean = np.array([beta_estimate[0], beta_estimate[1], np.log(sigma2_estimate)]) + proposal_cov = np.diag([1.0, 1.0, 1.0]) + + def log_proposal_density(param): + return stats.multivariable_normal.logpdf(param, mean = proposal_mean, cov = proposal_cov) + + # 6. 进行重要抽样 + num_samples = 10000 + samples = np.zeros((num_samples, 3)) + weights = np.zeros(num_samples) + log_weights = np.zeros(num_samples) + + for i in range(num_samples): + samples[i] = np.random.multivariate_normal(proposal_mean, proposal_cov) + unnormalized_log_posterior = unnormalized_log_posterior(samples[i]) + log_proposal = log_proposal_density(samples[i]) + log_weights[i] = unnormalized_log_posterior - log_proposal + + weights = np.exp(log_weights) + + + +if __name__ == "__main__": + baysian_linear_regression_inportance_sampling() + + \ No newline at end of file diff --git a/demo1.py b/demo1.py new file mode 100644 index 0000000..d708046 --- /dev/null +++ b/demo1.py @@ -0,0 +1,124 @@ +import numpy as np +import jax +import jax.numpy as jnp +from jax import jit, grad, vmap +import time + +# --- 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,)) + +print(f"JAX 数组 x: {x_jnp}") +print(f"JAX 数组 y (随机): {y_jnp}") +print(f"x 和 y 的点积: {jnp.dot(x_jnp, y_jnp)}") + + +# --- 2. jax.jit: 极致加速 --- +print("\n--- 2. @jit 加速 ---") +# 定义一个包含大量运算的函数 +def slow_function(x): + # 这是一个比较耗时的操作 (矩阵乘法循环) + for _ in range(100): + x = jnp.dot(x, x.T) + return x + +# 创建 JIT 编译版本的函数 +fast_function = jit(slow_function) + +# 创建数据 +big_matrix = jax.random.normal(key, (200, 200)) + +# --- 计时对比 --- +# 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} 秒") + +# 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} 秒") + +# 第二次运行,将只体现执行速度 +start_time = time.time() +fast_function(big_matrix).block_until_ready() +jit_time = time.time() - start_time +print(f"JIT 版本 (第二次) 耗时: {jit_time:.6f} 秒") + +print(f"加速比 (第二次运行 vs 普通版本): {numpy_time / jit_time:.2f} 倍") + + +# --- 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 + +# 使用 grad 创建一个计算 my_func 导数的函数 +# f'(x) = 3x^2 + 4x +grad_my_func = grad(my_func) + +x_val = 2.0 +derivative = grad_my_func(x_val) +expected_derivative = 3 * x_val**2 + 4 * x_val + +print(f"函数 f(x) = x^3 + 2x^2 + 5") +print(f"在 x = {x_val} 处的导数是: {derivative}") +print(f"理论上的导数值是: {expected_derivative}") + + +# --- 4. jax.vmap: 自动向量化 --- +print("\n--- 4. vmap 自动向量化 ---") +# 定义一个只能处理单个向量的函数 (向量点积) +def single_dot_product(a, b): + return jnp.dot(a, b) + +# 创建一批 (batch) 数据 +# 假设我们有 5 对向量,每对向量长度为 3 +batch_a = jnp.arange(15).reshape(5, 3) +batch_b = jnp.arange(15, 30).reshape(5, 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) + +results = batch_dot_product(batch_a, batch_b) + +print("批次数据 a:\n", batch_a) +print("批次数据 b:\n", batch_b) +print("vmap 向量化计算的点积结果:\n", results) + +# 对比手动 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)}") + + +# --- 5. JAX 的随机数处理 --- +print("\n--- 5. JAX 的随机数 ---") +# 1. 创建一个初始密钥 +key = jax.random.PRNGKey(42) +print(f"初始密钥: {key}") + +# 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}") diff --git a/demo2.py b/demo2.py new file mode 100644 index 0000000..2f4ccb7 --- /dev/null +++ b/demo2.py @@ -0,0 +1,224 @@ +import numpy as np +import scipy.linalg +import matplotlib.pyplot as plt +import control +import jax +import jax.numpy as jnp +from jax.scipy.linalg import solve_triangular +import numpyro +import numpyro.distributions as dist +from numpyro.infer import MCMC, NUTS +import arviz as az + +plt.rcParams['font.sans-serif'] = ['SimHei'] # or 'Microsoft YaHei' etc. +plt.rcParams['axes.unicode_minus'] = False + +DX, DU, DY = 2, 1, 1 +T_sim = 400 +SIGMA_W_STD = 0.3 +SIGMA_Z_STD = 1e-12 # Nugget value instead of 0.0 + +# 定义 LTI 系统 +a0_true,a1_true = 0.6, -1.5 +b0_true,b1_true = 1.0, 0.5 +d0_true = 0.0 +Ac_true = jnp.array([[0.,1.], [-a0_true, -a1_true]]) +Bc_true = jnp.array([[0.],[1.]]) +Cc_true = jnp.array([[b0_true, b1_true]]) +Dc_true = jnp.array([[d0_true]]) + +A_true, B_true, C_true, D_true = Ac_true, Bc_true, Cc_true, Dc_true + +Sigma_true = (SIGMA_W_STD**2)*jnp.eye(DX) +Gamma_true = (SIGMA_Z_STD**2)*jnp.eye(DY) + +# 生成仿真数据 +key = jax.random.PRNGKey(0) +key_u, key_w, key_x0 = jax.random.split(key, 3) + +u_data = jax.random.normal(key_u, shape=(T_sim, DU)) +x0 = jnp.zeros(DX) + +x_hist = np.zeros((T_sim + 1, DX)) +y_hist = np.zeros((T_sim, DY)) +x_hist[0] = x0 + +for t in range(T_sim): + w_t = jax.random.multivariate_normal(key_w, jnp.zeros(DX), Sigma_true) + key_w, _ = jax.random.split(key_w) # 更新随机数种子 + x_next = A_true @ x_hist[t] + B_true @ u_data[t] + w_t + y_t = C_true @ x_hist[t] + D_true @ u_data[t] # 理想观测,无观测噪声,不考虑直通项影响 + y_t_observed = y_t + D_true @ u_data[t] # 理想观测,无观测噪声 + + x_hist[t+1] = x_next + y_hist[t] = y_t_observed + +y_data = y_hist + +def kalman_log_likelihood(params_A, params_B, params_C, params_D, params_Sigma, params_Gamma, u, y): + """ + 使用 JAX 实现的卡尔曼滤波器计算状态空间模型的对数似然。 + + Args: + params_A, params_B, params_C, params_D: 系统矩阵 (JAX arrays)。 + params_Sigma: 过程噪声协方差矩阵 (JAX array)。 + params_Gamma: 测量噪声协方差矩阵 (JAX array)。 + u: 输入数据 (JAX array, shape (T_sim, DU))。 + y: 输出数据 (JAX array, shape (T_sim, DY))。 + + Returns: + 总对数似然 (JAX scalar)。 + """ + # 确保输入是 JAX 数组 (防御性编程) + u = jnp.asarray(u) + y = jnp.asarray(y) + params_A = jnp.asarray(params_A) + params_B = jnp.asarray(params_B) + params_C = jnp.asarray(params_C) + params_D = jnp.asarray(params_D) + params_Sigma = jnp.asarray(params_Sigma) + params_Gamma = jnp.asarray(params_Gamma) + + # 初始状态估计和协方差 + # 假设初始状态均值为零,具有一定的先验不确定性 + x_hat_0_0 = jnp.zeros(DX) + P_0_0 = jnp.eye(DX) * 0.01 # 可以根据需要调整初始不确定性 + + def kf_step(carry, t): + """ + 卡尔曼滤波器的单步迭代函数,用于 jax.lax.scan。 + + Args: + carry: 上一步的状态 (x_hat_{t-1|t-1}, P_{t-1|t-1})。 + t: 当前时间步索引 (从 0 到 T_sim-1)。 + + Returns: + 新的状态 (x_hat_{t|t}, P_{t|t}) 和当前步的对数似然 log_lik_t。 + """ + x_hat_t_t_prev, P_t_t_prev = carry + + # --- 获取当前和过去的输入 --- + u_t_prev = jnp.where(t > 0, u[t - 1], jnp.zeros(DU)) + u_t_obs = u[t] + + # --- 获取当前观测 --- + y_t_obs = y[t] + + # --- 预测步骤 --- + x_hat_t_tm1 = params_A @ x_hat_t_t_prev + params_B @ u_t_prev + P_t_tm1 = params_A @ P_t_t_prev @ params_A.T + params_Sigma + + # --- 更新步骤 --- + nu_t = y_t_obs - (params_C @ x_hat_t_tm1 + params_D @ u_t_obs) + S_t = params_C @ P_t_tm1 @ params_C.T + params_Gamma + + try: + # 尝试直接求逆 + S_t_inv = jnp.linalg.inv(S_t) + K_t = P_t_tm1 @ params_C.T @ S_t_inv + except jnp.linalg.LinAlgError: + S_t_inv = jnp.full_like(S_t, jnp.nan) + K_t = jnp.full((DX, DY), jnp.nan) + + + x_hat_t_t = x_hat_t_tm1 + K_t @ nu_t + P_t_t = (jnp.eye(DX) - K_t @ params_C) @ P_t_tm1 + + sign, log_det_S = jnp.linalg.slogdet(S_t) + + # 确保 nu_t 是合适的形状进行二次型计算 + nu_t_col = nu_t.reshape(-1, 1) if DY > 0 else nu_t + quad_form = (nu_t_col.T @ S_t_inv @ nu_t_col).squeeze() + + log_lik_t = -0.5 * DY * jnp.log(2 * jnp.pi) - 0.5 * log_det_S - 0.5 * quad_form + + # 对于求逆失败的情况,返回一个非常差的似然 + log_lik_t = jnp.where(jnp.isnan(K_t).any(), -jnp.inf, log_lik_t) + + # 返回新的状态和当前步的对数似然 (确保是标量) + return (x_hat_t_t, P_t_t), jnp.asarray(log_lik_t).squeeze() + + initial_carry = (x_hat_0_0, P_0_0) + time_indices = jnp.arange(T_sim) + final_carry, log_lik_steps = jax.lax.scan(kf_step, initial_carry, time_indices) + + return jnp.sum(log_lik_steps) + + + +# 定义贝叶斯模型 +def model_standard(u, y = None): + # A 的先验 (DX*DX elements) + A_s_flat = numpyro.sample("A_s_flat", dist.Normal(0., 1.).expand([DX*DX])) + A_s = A_s_flat.reshape((DX, DX)) + # B 的先验 (DX*DU elements) + B_s_flag = numpyro.sample("B_s_flat", dist.Normal(0., 1.).expand([DX*DU])) + B_s = B_s_flag.reshape((DX, DU)) + # C 的先验 (DY*DX elements) + C_s_flat = numpyro.sample("C_s_flat", dist.Normal(0., 1.).expand([DY*DX])) + C_s = C_s_flat.reshape((DY, DX)) + # D 的先验 (仍然采样一个标量) + D_s_scalar = numpyro.sample("D_s", dist.Normal(0., 1.)) + # 将标量 reshape 为 (DY, DU) = (1, 1) 矩阵 + D_s = D_s_scalar.reshape((DY, DU)) + # Sigma 的先验 (DX*DX elements) + sigma_w = numpyro.sample("sigma_w", dist.HalfCauchy(1.)) + Sigma_w = (sigma_w**2) * jnp.eye(DX) + # Gamma 的先验 (DY*DY elements) + Gamma_z = (SIGMA_Z_STD**2) * jnp.eye(DY) + + # 计算似然 + log_lik = kalman_log_likelihood(A_s, B_s, C_s, D_s, Sigma_w, Gamma_z, u, y) + + # 观测数据的似然 + numpyro.factor("obs", log_lik) + + +# 运行 MCMC 进行推断 +kernel_s = NUTS(model_standard) +mcmc_s = MCMC(kernel_s, num_warmup=1000, num_samples=2000, num_chains=1) +mcmc_s.run(jax.random.PRNGKey(2), u=u_data, y=y_data) +samples_s = mcmc_s.get_samples() +idata_s = az.from_numpyro(mcmc_s) + +def model_canonical(u, y=None): + a0 = numpyro.sample("a0", dist.Normal(0., 1.)) + a1 = numpyro.sample("a1", dist.Normal(0., 1.)) + b0 = numpyro.sample("b0", dist.Normal(0., 1.)) + b1 = numpyro.sample("b1", dist.Normal(0., 1.)) + d0 = numpyro.sample("d0", dist.Normal(0., 1.)) + + A_c = jnp.array([[0., 1.], [-a0, -a1]]) + B_c = jnp.array([[0.],[1.]]) + C_c = jnp.array([[b0, b1]]) + D_c = jnp.array([[d0]]) + + sigma_w = numpyro.sample("sigma_w", dist.HalfCauchy(1.)) + Sigma_w = (sigma_w**2) * jnp.eye(DX) + Gamma_z = (SIGMA_Z_STD**2) * jnp.eye(DY) + + log_lik = kalman_log_likelihood(A_c, B_c, C_c, D_c, Sigma_w, Gamma_z, u, y) + + numpyro.factor("obs", log_lik) + + +# 运行 MCMC 进行推断 +kernel_c = NUTS(model_canonical) +mcmc_c = MCMC(kernel_c, num_warmup=1000, num_samples=2000, num_chains=1) +mcmc_c.run(jax.random.PRNGKey(3), u=u_data, y=y_data) +samples_c = mcmc_c.get_samples() +idata_c = az.from_numpyro(mcmc_c) + +# # Plot pair plot for standard parameters (select a few, e.g., A_s elements) +az.plot_pair(idata_s, var_names=['A_s_flat'], coords={'A_s_flat_dim_0': [0, 1, 2, 3]}) # Plot A11, A12, A21, A22 +plt.suptitle("后验分布 (标准参数化 $\Theta_s$)") +plt.tight_layout(rect=[0, 0.03, 1, 0.95]) # Adjust layout to prevent title overlap +plt.show() + +# Plot pair plot for canonical parameters +az.plot_pair(idata_c, var_names=['a0', 'a1', 'b0', 'b1']) +plt.suptitle("后验分布 (规范参数化 $\Theta_c$)") +plt.tight_layout(rect=[0, 0.03, 1, 0.95]) +plt.show() + +