345 lines
13 KiB
Python
345 lines
13 KiB
Python
import numpy as np
|
|
import jax
|
|
import jax.numpy as jnp
|
|
import matplotlib.pyplot as plt
|
|
from matplotlib.patches import Circle
|
|
import seaborn as sns
|
|
import arviz as az
|
|
import numpyro
|
|
|
|
from generateGroudTruth import generate_ground_truth_system
|
|
from generateSimData import simulate_lti_data
|
|
# 确保 models_and_mcmc.py 中的 run_mcmc 接受 init_params 参数
|
|
from models_and_mcmc import run_mcmc, model_canonical, model_standard
|
|
|
|
# =========================================================================
|
|
# 配置中文字体
|
|
# =========================================================================
|
|
try:
|
|
# Windows 系统优先尝试这些字体
|
|
plt.rcParams['font.sans-serif'] = ['Microsoft YaHei', 'SimHei', 'SimSun', 'KaiTi', 'FangSong', 'Arial Unicode MS']
|
|
plt.rcParams['axes.unicode_minus'] = False # 正常显示负号
|
|
print("✓ 中文字体配置成功")
|
|
except Exception as e:
|
|
print(f"⚠ 字体配置警告: {e}")
|
|
print(" 如果图表中文显示异常,请运行 check_fonts.py 查看可用字体")
|
|
|
|
# =========================================================================
|
|
# 实现标准型到规范型的转换
|
|
# =========================================================================
|
|
def standard_to_canonical(A_s, B_s, C_s):
|
|
"""
|
|
将一个 2x2 的标准状态空间系统 (A_s, B_s, C_s) 转换为控制器规范型参数。
|
|
|
|
此实现基于论文附录 D.2 的推导,
|
|
通过计算特征多项式和可控性矩阵来找到变换矩阵 T_c。
|
|
|
|
参数:
|
|
A_s (np.ndarray): 2x2 状态矩阵
|
|
B_s (np.ndarray): 2x1 输入矩阵
|
|
C_s (np.ndarray): 1x2 观测矩阵
|
|
|
|
返回:
|
|
dict: 包含 'a0', 'a1', 'b0', 'b1' 的字典
|
|
"""
|
|
|
|
# 确保输入是 numpy 数组
|
|
A_s = np.asarray(A_s)
|
|
B_s = np.asarray(B_s)
|
|
C_s = np.asarray(C_s)
|
|
|
|
if A_s.shape != (2, 2) or B_s.shape != (2, 1) or C_s.shape != (1, 2):
|
|
raise ValueError(f"输入维度不正确: A_s {A_s.shape}, B_s {B_s.shape}, C_s {C_s.shape}")
|
|
|
|
# 1. 计算特征多项式系数 (来自 Ac)
|
|
# p(λ) = λ^2 - tr(A_s)λ + det(A_s)
|
|
# 规范型 p(λ) = λ^2 + a1*λ + a0
|
|
# 比较系数: a1 = -tr(A_s), a0 = det(A_s)
|
|
a1_true = -np.trace(A_s)
|
|
a0_true = np.linalg.det(A_s)
|
|
|
|
# 2. 构造逆转换矩阵 T_c^{-1} = [f1, f2]
|
|
I = np.eye(2)
|
|
|
|
# 根据附录 D.2 (1157), f_k = (A_s + a_{d-1}I)f_{k+1} + ...
|
|
f2 = B_s
|
|
f1 = (A_s + a1_true * I) @ B_s
|
|
|
|
Tc_inv = np.hstack([f1, f2]) # 这是 T_c
|
|
print(f"[standard_to_canonical] 恢复的 T_c (即 Tc_inv):\n{Tc_inv}")
|
|
|
|
# 3. 检查可控性 (Controllability)
|
|
if np.linalg.matrix_rank(Tc_inv) < 2:
|
|
raise ValueError("系统不可控 (Uncontrollable), 无法转换为控制器规范型。")
|
|
|
|
# 4. 计算转换矩阵 T_c (这是 T_c^{-1})
|
|
Tc = np.linalg.inv(Tc_inv) # 这是 T_c^{-1}
|
|
print(f"[standard_to_canonical] 恢复的 T_c^{{-1}} (即 Tc):\n{Tc}\n")
|
|
|
|
# 5. 应用变换找到 C_c = [b0, b1] (来自 Cc)
|
|
# 正确的公式是 C_c = C_s * T_c
|
|
# 在我们的变量名中, T_c 是 Tc_inv
|
|
C_c = C_s @ Tc_inv # 这是正确行 (C_s * T_c)
|
|
|
|
b0_true = C_c[0, 0]
|
|
b1_true = C_c[0, 1]
|
|
|
|
# 打印矩阵
|
|
print(f"[standard_to_canonical] 计算得到的规范型参数:")
|
|
print(f" a0: {a0_true}, a1: {a1_true}, b0: {b0_true}, b1: {b1_true}\n")
|
|
|
|
return {
|
|
'a0': a0_true,
|
|
'a1': a1_true,
|
|
'b0': b0_true,
|
|
'b1': b1_true
|
|
}
|
|
|
|
# =========================================================================
|
|
# 可视化函数
|
|
# =========================================================================
|
|
|
|
def plot_canonical_results(mcmc_canonical, true_params_c):
|
|
"""
|
|
可视化规范型模型的 MCMC 结果,重现图 2。
|
|
"""
|
|
print("\n--- 正在生成规范型模型的后验分布图 (图 2)... ---")
|
|
|
|
idata_c = az.from_numpyro(mcmc_canonical)
|
|
samples_c = mcmc_canonical.get_samples()
|
|
|
|
# 图 2(a): 参数的配对图
|
|
az.plot_pair(
|
|
idata_c,
|
|
var_names=['a0', 'a1', 'b0', 'b1'],
|
|
kind='kde',
|
|
marginals=True,
|
|
point_estimate='mean',
|
|
reference_values=true_params_c,
|
|
figsize=(10, 10)
|
|
)
|
|
plt.suptitle("图 2(a) 复现: 规范型参数的后验分布", y=1.02, fontsize=16)
|
|
|
|
# 图 2(b): 特征值在复平面上的分布
|
|
true_eigenvalues = np.roots([1, true_params_c['a1'], true_params_c['a0']])
|
|
|
|
posterior_eigenvalues = []
|
|
# 避免使用过多样本导致计算缓慢,可以对样本进行降采样
|
|
num_plot_samples = min(5000, len(samples_c['a0']))
|
|
plot_indices = np.random.choice(len(samples_c['a0']), num_plot_samples, replace=False)
|
|
|
|
for i in plot_indices:
|
|
a0 = samples_c['a0'][i]
|
|
a1 = samples_c['a1'][i]
|
|
posterior_eigenvalues.extend(np.roots([1, a1, a0]))
|
|
posterior_eigenvalues = np.array(posterior_eigenvalues, dtype=np.complex128)
|
|
|
|
mean_params = {k: np.mean(v) for k, v in samples_c.items()}
|
|
map_eigenvalues = np.roots([1, mean_params['a1'], mean_params['a0']])
|
|
|
|
plt.figure(figsize=(8, 8))
|
|
sns.kdeplot(x=posterior_eigenvalues.real, y=posterior_eigenvalues.imag,
|
|
fill=True, cmap="Blues", levels=10, alpha=0.7) # 添加透明度
|
|
plt.plot(true_eigenvalues.real, true_eigenvalues.imag, 'ro', markersize=10,
|
|
label=f'真实特征值: {true_eigenvalues[0]:.3f}')
|
|
plt.plot(map_eigenvalues.real, map_eigenvalues.imag, 'gs', markersize=10,
|
|
label=f'后验均值估计: {map_eigenvalues[0]:.3f}')
|
|
circle = Circle((0, 0), 1, color='gray', fill=False, linestyle='--')
|
|
plt.gca().add_artist(circle)
|
|
|
|
plt.title('图 2(b) 复现: 主特征值的后验分布', fontsize=16)
|
|
plt.xlabel('Re(λ)')
|
|
plt.ylabel('Im(λ)')
|
|
plt.legend()
|
|
plt.axis('equal')
|
|
plt.grid(True)
|
|
|
|
|
|
def plot_standard_results(mcmc_standard, true_params_s):
|
|
"""
|
|
可视化标准型模型的 MCMC 结果,重现图 3。
|
|
"""
|
|
print("\n--- 正在生成标准型模型的后验分布图 (图 3)... ---")
|
|
|
|
idata_s = az.from_numpyro(mcmc_standard)
|
|
|
|
# 为了清晰起见,只选择论文图 3a 中显示的几个参数子集
|
|
plot_vars = ['A11', 'A12', 'A21', 'A22', 'B1', 'B2', 'C1', 'C2']
|
|
# 限制绘制的参数数量,避免图像过于拥挤
|
|
az.plot_pair(
|
|
idata_s,
|
|
var_names=['A11', 'A12', 'A21','A22', 'B1', 'B2', 'C1','C2'], # 选择部分参数展示
|
|
kind='kde',
|
|
marginals=True,
|
|
# point_estimate='mean', # 对于多峰分布,均值可能误导,不显示
|
|
reference_values={k: v for k, v in true_params_s.items() if k in plot_vars},
|
|
figsize=(12, 12) # 稍微增大图像尺寸
|
|
)
|
|
plt.suptitle("图 3(a) 复现: 标准型参数的后验分布 (部分)", y=1.02, fontsize=16)
|
|
|
|
def get_user_choice():
|
|
"""
|
|
获取用户选择运行哪个模型
|
|
"""
|
|
print("\n" + "="*60)
|
|
print("请选择要运行的模型类型:")
|
|
print("1. 只运行规范型模型 (Canonical)")
|
|
print("2. 只运行标准型模型 (Standard ABCD)")
|
|
print("3. 两个模型都运行")
|
|
print("="*60)
|
|
|
|
while True:
|
|
try:
|
|
choice = input("请输入选择 (1/2/3): ").strip()
|
|
if choice in ['1', '2', '3']:
|
|
return int(choice)
|
|
else:
|
|
print("请输入有效的选择 (1, 2 或 3)")
|
|
except KeyboardInterrupt:
|
|
print("\n程序被用户中断")
|
|
exit()
|
|
except:
|
|
print("请输入有效的选择 (1, 2 或 3)")
|
|
|
|
# =========================================================================
|
|
# 主程序
|
|
# =========================================================================
|
|
if __name__ == '__main__':
|
|
# 获取用户选择
|
|
choice = get_user_choice()
|
|
|
|
# --- 0. 设置 ---
|
|
numpyro.set_host_device_count(4)
|
|
main_rng_key = jax.random.PRNGKey(42)
|
|
|
|
# --- 1. 生成 Ground Truth 系统 ---
|
|
print("\n--- (步骤 1) 生成真实系统 ---")
|
|
gt_key, sim_key, mcmc_key, init_noise_key = jax.random.split(main_rng_key, 4) # 多分配一个 key
|
|
A_true, B_true, C_true, D_true = generate_ground_truth_system(
|
|
rng_seed=int(gt_key[0])
|
|
)
|
|
|
|
# --- 2. 仿真数据 ---
|
|
print("\n--- (步骤 2) 生成仿真数据 ---")
|
|
T_steps, sigma_proc, sigma_meas = 800, 0.05, 0.05
|
|
u_data, y_data = simulate_lti_data(
|
|
A_true, B_true, C_true, D_true, T_steps, sigma_proc, sigma_meas,
|
|
rng_seed=int(sim_key[0])
|
|
)
|
|
u_data_jax, y_data_jax = jnp.array(u_data), jnp.array(y_data)
|
|
|
|
# --- 步骤 2.5: 计算 MCMC 初始值 ---
|
|
true_params_c = standard_to_canonical(A_true, B_true, C_true)
|
|
true_params_s = {
|
|
'A11': A_true[0, 0], 'A12': A_true[0, 1],
|
|
'A21': A_true[1, 0], 'A22': A_true[1, 1],
|
|
'B1': B_true[0, 0], 'B2': B_true[1, 0],
|
|
'C1': C_true[0, 0], 'C2': C_true[0, 1]
|
|
}
|
|
|
|
# --- 3 & 4. 运行 MCMC ---
|
|
mcmc_canonical = None
|
|
mcmc_standard = None
|
|
|
|
# 分配用于初始值噪声的 key
|
|
init_noise_key_c, init_noise_key_s = jax.random.split(init_noise_key)
|
|
|
|
if choice in [1, 3]: # 运行规范型
|
|
print("\n--- 运行规范型模型 MCMC (在真实值附近初始化) ---")
|
|
mcmc_key_c, mcmc_key = jax.random.split(mcmc_key)
|
|
|
|
# --- (修改) 在真实值附近添加小的随机扰动 (±30%) ---
|
|
init_key = init_noise_key_c # 使用独立的 key
|
|
noise_scale = 0.3 # 30% 的扰动
|
|
init_params_c_noisy = {}
|
|
for param_name, true_value in true_params_c.items():
|
|
# 确保即使 true_value 为 0 也有扰动,添加一个小的基准值
|
|
base_value = abs(true_value) if abs(true_value) > 1e-6 else 1.0
|
|
noise = jax.random.normal(init_key, shape=()) * base_value * noise_scale
|
|
# 确保 a0, a1 扰动后仍在有效范围内
|
|
if param_name == 'a0':
|
|
noisy_val = np.clip(true_value + float(noise), -0.99, 0.99) # 限制在 (-1, 1) 内
|
|
elif param_name == 'a1':
|
|
# 需要知道 a0 的扰动值来确定 a1 的范围
|
|
a0_noisy = init_params_c_noisy.get('a0', true_params_c['a0']) # 获取已扰动的 a0
|
|
low_bound = -1 - a0_noisy + 1e-6 # 加一点边界防止卡住
|
|
high_bound = 1 + a0_noisy - 1e-6
|
|
noisy_val = np.clip(true_value + float(noise), low_bound, high_bound)
|
|
else:
|
|
noisy_val = true_value + float(noise)
|
|
|
|
init_params_c_noisy[param_name] = noisy_val
|
|
init_key, _ = jax.random.split(init_key) # 更新 key
|
|
|
|
print(f"规范型初始参数 (带扰动): {init_params_c_noisy}")
|
|
|
|
mcmc_canonical = run_mcmc(
|
|
model_canonical,
|
|
mcmc_key_c,
|
|
u_data_jax,
|
|
y_data_jax,
|
|
sigma_proc,
|
|
sigma_meas,
|
|
init_params = None # <-- 使用带扰动的初始值
|
|
)
|
|
|
|
if choice in [2, 3]: # 运行标准型
|
|
print("\n--- 运行标准型模型 MCMC (在真实值附近初始化) ---")
|
|
mcmc_key_s, mcmc_key = jax.random.split(mcmc_key)
|
|
|
|
# --- (修改) 在真实值附近添加小的随机扰动 (±30%) ---
|
|
init_key = init_noise_key_s # 使用独立的 key
|
|
noise_scale = 0.3 # 30% 的扰动
|
|
init_params_s_noisy = {}
|
|
for param_name, true_value in true_params_s.items():
|
|
# 确保即使 true_value 为 0 也有扰动
|
|
base_value = abs(true_value) if abs(true_value) > 1e-6 else 1.0
|
|
noise = jax.random.normal(init_key, shape=()) * base_value * noise_scale
|
|
init_params_s_noisy[param_name] = true_value + float(noise)
|
|
init_key, _ = jax.random.split(init_key) # 更新 key
|
|
|
|
print(f"标准型初始参数 (带扰动): {init_params_s_noisy}")
|
|
|
|
mcmc_standard = run_mcmc(
|
|
model_standard,
|
|
mcmc_key_s,
|
|
u_data_jax,
|
|
y_data_jax,
|
|
sigma_proc,
|
|
sigma_meas,
|
|
init_params=init_params_s_noisy # <-- 使用带扰动的初始值
|
|
)
|
|
|
|
# --- 5. 分析与可视化 ---
|
|
|
|
# (计算真实参数的代码已移至 步骤 2.5)
|
|
|
|
# 根据选择显示结果
|
|
if choice in [1, 3] and mcmc_canonical is not None:
|
|
plot_canonical_results(mcmc_canonical, true_params_c)
|
|
|
|
print("\n" + "="*50)
|
|
print("规范型模型分析:")
|
|
print("1. 后验分布应为单峰 (unimodal) 且近似高斯。")
|
|
print("2. 参数解释简单,后验均值是好的点估计。")
|
|
print("3. 特征值分布应集中围绕真实值。")
|
|
print("="*50)
|
|
|
|
if choice in [2, 3] and mcmc_standard is not None:
|
|
plot_standard_results(mcmc_standard, true_params_s)
|
|
|
|
print("\n" + "="*50)
|
|
print("标准型模型分析:")
|
|
print("1. 后验分布可能呈现复杂的多峰 (multi-modal) 和强相关性。")
|
|
print("2. 这是由参数的非唯一性 (non-identifiability) 造成的。")
|
|
print("3. MCMC 采样效率可能较低,点估计(如均值)可能无意义。")
|
|
print("="*50)
|
|
|
|
# 显示所有图像
|
|
if choice != 0:
|
|
print("\n显示所有图像...")
|
|
plt.show()
|
|
|
|
print("\n程序执行完成!")
|
|
|