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()