initial commit
This commit is contained in:
@@ -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()
|
||||||
|
|
||||||
|
|
||||||
@@ -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}")
|
||||||
@@ -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()
|
||||||
|
|
||||||
|
|
||||||
Reference in New Issue
Block a user