Files
Baysian_Canonical_Identific…/main.py
T

314 lines
12 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) 转换为控制器规范型参数。
返回:
dict: 包含 'a0', 'a1', 'b0', 'b1' 的字典
"""
dx = A_s.shape[0]
if dx != 2:
raise ValueError("此转换函数仅为 dx=2 的情况实现。")
# 1. 计算特征多项式系数: p(λ) = λ^2 + a1*λ + a0
a1_true = -np.trace(A_s)
a0_true = np.linalg.det(A_s)
# 2. 构造转换矩阵 T_c
I = np.eye(dx)
f1 = (A_s + a1_true * I) @ B_s
f2 = B_s
Tc_inv = np.hstack([f1, f2])
if np.linalg.matrix_rank(Tc_inv) < dx:
raise np.linalg.LinAlgError("系统 (A_s, B_s) 不是可控的,无法转换为控制器规范型。")
Tc = np.linalg.inv(Tc_inv)
# 3. 转换 C 矩阵: C_c = C_s * T_c
C_c = C_s @ Tc
b0_true = C_c[0, 0]
b1_true = C_c[0, 1]
true_params = {'a0': a0_true, 'a1': a1_true, 'b0': b0_true, 'b1': b1_true}
print("\n--- Ground Truth 转换结果 ---")
print(f"真实规范型参数: {true_params}")
return true_params
# =========================================================================
# 可视化函数
# =========================================================================
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 = 400, 0.3, 0.5
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)
# --- (修改) 在真实值附近添加小的随机扰动 (±10%) ---
init_key = init_noise_key_c # 使用独立的 key
noise_scale = 0.1 # 10% 的扰动
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,
init_params=init_params_c_noisy # <-- 使用带扰动的初始值
)
if choice in [2, 3]: # 运行标准型
print("\n--- 运行标准型模型 MCMC (在真实值附近初始化) ---")
mcmc_key_s, mcmc_key = jax.random.split(mcmc_key)
# --- (修改) 在真实值附近添加小的随机扰动 (±10%) ---
init_key = init_noise_key_s # 使用独立的 key
noise_scale = 0.1 # 10% 的扰动
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,
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程序执行完成!")