###################################################################################
# PROGRAM: importance_sampling_dice.py
# DATE: 2026-8-30
# NOTICE: This program accompanies the book "マルコフ連鎖モンテカルロ法入門"
#                     by Koji Hukushima and Yoshihiko Nishikawa
###################################################################################

import numpy as np
import matplotlib.pyplot as plt

# ==========================================
# Step 1: 確率分布の設定
# ==========================================
np.random.seed(42)  # 再現性のためシードを固定
states = np.array([1, 2, 3, 4, 5, 6])

# ターゲット分布（本来知りたい分布：一様なサイコロ）
p_target = np.ones(6) / 6.0  

# 提案分布（実際にサンプリングできる偏ったサイコロ）
p_proposal = np.array([0.01, 0.02, 0.07, 0.10, 0.35, 0.45])

# シミュレーションの設定
sample_sizes = [10, 100, 1000, 10000, 100000]
n_runs_list = [100, 100, 100, 100, 100]  # 各サンプルサイズでの試行回数

# ==========================================
# Step 2: 重点サンプリングの実行
# ==========================================
means = []
std_errors = []
mean_ess_list = []  # 各サンプルサイズにおける平均ESSを記録

for n, n_runs in zip(sample_sizes, n_runs_list):
    estimates = []
    ess_values = []
    
    for _ in range(n_runs):
        # 1. 提案分布からサンプリング
        samples = np.random.choice(states, size=n, p=p_proposal)
        
        # 2. 重要度重み (Importance Weight) の計算
        # w = p_target(x) / p_proposal(x)
        weights = p_target[samples - 1] / p_proposal[samples - 1]
        
        # 3. 自己正規化重点サンプリングによる期待値の推定
        estimate = np.sum(samples * weights) / np.sum(weights)
        estimates.append(estimate)
        
        # 4. ESS (Effective Sample Size) の計算
        # ESS = (Σw_i)^2 / Σ(w_i^2)
        ess = (np.sum(weights)**2) / np.sum(weights**2)
        ess_values.append(ess)
        
    estimates = np.array(estimates)
    means.append(np.mean(estimates))
    std_errors.append(np.std(estimates, ddof=1) / np.sqrt(n_runs))  # 標準誤差
    mean_ess_list.append(np.mean(ess_values))

# ==========================================
# Step 3: ESSの結果出力
# ==========================================
print("=== 重点サンプリングの有効サンプルサイズ (ESS) ===")
for n, ess in zip(sample_sizes, mean_ess_list):
    efficiency = (ess / n) * 100
    print(f"名目サンプルサイズ: {n:7d} -> 実効サンプルサイズ (ESS): {ess:8.1f} (効率: {efficiency:5.1f}%)")

# ==========================================
# Step 4: グラフによる可視化
# ==========================================
# 描画用の全体設定
# plt.rcParams['font.family'] = 'IPAMincho' # 必要に応じて有効化
plt.rcParams['mathtext.fontset'] = 'cm'
plt.rcParams["font.size"] = 20          # 汎用性を考慮し少しだけサイズダウン
plt.rcParams['ytick.labelsize'] = 18
plt.rcParams['xtick.labelsize'] = 18

fig, ax = plt.subplots(figsize=(8, 6))

# 期待値の推定値と標準誤差のエラーバーをプロット
ax.errorbar(sample_sizes, means, yerr=std_errors, 
            ms=10, fmt='o', capsize=5, color='magenta', ecolor='magenta', alpha=0.8, 
            label='Estimated Value (Importance Sampling)')

# 真の期待値（一様分布のサイコロの期待値 = 3.5）
ax.axhline(3.5, color='black', linestyle='--', linewidth=2, label='True Expected Value (3.5)')

ax.set_xscale('log')
ax.set_xlabel("Sample Size $M$")
ax.set_ylabel("Estimated Expected Value")
ax.set_yticks([3.5, 3.7, 3.9, 4.1])

legend = ax.legend(edgecolor='black', fontsize=14)
legend.get_frame().set_linewidth(1.5)

plt.tight_layout()
plt.show()
