Files
Baysian_Canonical_Identific…/demo2.py
T
2025-10-20 09:53:02 +08:00

225 lines
7.6 KiB
Python

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