225 lines
7.6 KiB
Python
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()
|
|
|
|
|