###################################################################################
# PROGRAM: ex_GM.py
# DATE: 2026-8-31
# NOTICE: This program accompanies the book "マルコフ連鎖モンテカルロ法入門"
#                     by Koji Hukushima and Yoshihiko Nishikawa
###################################################################################
"""
========================================================================
【遊び方＆実験ガイド】レプリカ交換法 (Parallel Tempering) 
========================================================================
この方法は、複数の温度（β）のMCMCチェーンを並行して走らせ、チェーン同士
の状態を「交換」することで、局所解（1つの山）にトラップされるのを防ぐ手法である。

■ いろいろ書き換えて遊んでみよう。

1. USE_EXCHANGE (交換フラッグ)
   - True : 交換法あり。高温のチェーンが探索した広い範囲が、
            低温のチェーン（β=1.0、本来の目的の分布）に伝播し、
            両方の山からサンプリングできることがわかる。
   - False: 交換法なし（独立なMCMC）。β=1.0のチェーンが片方の山から
            抜け出せず、サンプリングが偏る様子が確認できる。

2. num_steps と burn_in
   - ステップ数を変えて、収束の様子を観察できる。
========================================================================
"""

import numpy as np
import matplotlib.pyplot as plt
from scipy.stats import multivariate_normal
from mpl_toolkits.axes_grid1 import make_axes_locatable

# ==========================================
# 読者が遊べるパラメータ設定
# ==========================================
USE_EXCHANGE = False         # ★ ここで True / False を切り替え
num_steps = 20000            # 総サンプリングステップ数
burn_in = 5000               # 捨てる初期サンプルの数
exchange_interval = 1        # 交換を試みる間隔（ステップ）

betas = np.array([0.0, 0.3, 0.7, 1.0])
num_chains = len(betas)

# ==========================================
# 1. 分布の定義 (ターゲットと初期分布)
# ==========================================
# ターゲット：2つの山を持つ混合ガウス分布 (β=1.0 で一致)
mu1, mu2 = np.array([0, 0]), np.array([4, 4])
cov1 = np.array([[1.0, 0.9], [0.9, 1.5]])
cov2 = np.array([[1.0, -0.8], [-0.8, 1.0]])
pi1, pi2 = 0.5, 0.5

# ベース：広く平坦なガウス分布 (β=0.0 で一致)
mu0 = np.array([2, 2])
cov0 = np.array([[2.0, 0.0], [0.0, 2.0]])

def log_p(x):
    """ターゲット分布の対数密度"""
    p1 = multivariate_normal.logpdf(x, mean=mu1, cov=cov1)
    p2 = multivariate_normal.logpdf(x, mean=mu2, cov=cov2)
    return np.logaddexp(np.log(pi1) + p1, np.log(pi2) + p2)

def log_p0(x):
    """ベース分布の対数密度"""
    return multivariate_normal.logpdf(x, mean=mu0, cov=cov0)

def log_interpolated_density(x, beta):
    """βによる補間分布"""
    return beta * log_p(x) + (1 - beta) * log_p0(x)

# ==========================================
# 2. Metropolis-Hastings 局所更新
# ==========================================
def metropolis_step(x, logpdf, delta=0.8):
    """単純なランダムウォークMHステップ"""
    x_new = x + np.random.uniform(-delta/2, delta/2, size=2)
    log_accept_ratio = logpdf(x_new) - logpdf(x)
    if np.log(np.random.rand()) < log_accept_ratio:
        return x_new
    return x

# ==========================================
# 3. レプリカ交換法（Parallel Tempering）本体
# ==========================================
# 初期状態（全系列ともベース分布から独立に生成）
chains = [np.random.multivariate_normal(mu0, cov0) for _ in betas]
samples = [[] for _ in betas]

# 交換成功回数のカウント用配列 (隣り合うペアごと)
ex_accpt = [0] * (num_chains - 1)
ex_attempts = [0] * (num_chains - 1)

np.random.seed(42) # 比較用にシードを固定

print(f"--- Sampling Started (Exchange: {USE_EXCHANGE}) ---")
for step in range(num_steps):
    # 1. 各チェーン（温度）ごとの局所更新
    for i, beta in enumerate(betas):
        chains[i] = metropolis_step(chains[i], lambda x, b=beta: log_interpolated_density(x, b))
        samples[i].append(chains[i])

    # 2. 交換ステップ（USE_EXCHANGE が True のときのみ実行）
    if USE_EXCHANGE and step % exchange_interval == 0:
        # 偶数番目・奇数番目のペアを交互に交換試行する (詳細釣り合いを満たすため)
        start_idx = step % 2
        for i in range(start_idx, num_chains - 1, 2):
            x1, x2 = chains[i], chains[i+1]
            beta1, beta2 = betas[i], betas[i+1]

            logpdf1 = lambda x: log_interpolated_density(x, beta1)
            logpdf2 = lambda x: log_interpolated_density(x, beta2)

            # 交換の受容確率の計算
            log_accept_ratio = (logpdf1(x2) + logpdf2(x1)) - (logpdf1(x1) + logpdf2(x2))
            
            ex_attempts[i] += 1
            if np.log(np.random.rand()) < log_accept_ratio:
                chains[i], chains[i+1] = x2, x1  # 交換成功
                ex_accpt[i] += 1 

if USE_EXCHANGE:
    print("ペアごとの交換成功率:")
    for i in range(num_chains - 1):
        rate = ex_accpt[i] / max(1, ex_attempts[i]) * 100
        print(f"  β={betas[i]:.1f} ⇔ β={betas[i+1]:.1f} : {rate:.1f}%")

# ==========================================
# 4. 可視化用データ整形と準備
# ==========================================
sample_data = {
    rf"$\beta={beta:.1f}$": np.vstack(samples[i][burn_in:])
    for i, beta in enumerate(betas)
}

x_grid = np.linspace(-4, 10, 200)
y_grid = np.linspace(-4, 10, 200)
X, Y = np.meshgrid(x_grid, y_grid)
grid_points = np.stack([X.ravel(), Y.ravel()], axis=-1)

# ==========================================
# 5. 描画
# ==========================================
plt.rcParams['mathtext.fontset'] = 'cm'
plt.rcParams["font.size"] = 18

fig, axes = plt.subplots(1, 4, figsize=(20, 5))

for i, (label, data) in enumerate(sample_data.items()):
    beta = betas[i]
    ax = axes[i]
    ax.set_title(label, fontsize=20, loc='right', pad=10)
    ax.set_xlim(-3, 8)
    ax.set_ylim(-3, 8)
    ax.set_xticks(np.arange(-2, 10, 2))
    ax.set_yticks(np.arange(-2, 10, 2))
    ax.grid(True, alpha=0.3)

    # 散布図
    ax.scatter(data[:, 0], data[:, 1], s=2, alpha=0.1, color='magenta')

    # 補間分布の等高線
    Z_log = np.array([log_interpolated_density(pt, beta) for pt in grid_points])
    Z = np.exp(Z_log - np.max(Z_log))
    Z = Z.reshape(X.shape)
    ax.contour(X, Y, Z, levels=10, cmap="gray", alpha=0.7)

    # ヒストグラム（周辺分布）
    divider = make_axes_locatable(ax)
    ax_histx = divider.append_axes("top", size=1.0, pad=0.1, sharex=ax)
    ax_histy = divider.append_axes("right", size=1.0, pad=0.1, sharey=ax)

    ax_histx.hist(data[:, 0], bins=40, range=(-3, 8), color="gray", density=True, alpha=0.7)
    ax_histy.hist(data[:, 1], bins=40, range=(-3, 8), orientation='horizontal', color="gray", density=True, alpha=0.7)

    # 理論的な周辺分布の計算 (数値積分)
    px_beta = np.trapz(Z, y_grid, axis=0)
    py_beta = np.trapz(Z, x_grid, axis=1)
    px_beta /= np.trapz(px_beta, x_grid)
    py_beta /= np.trapz(py_beta, y_grid)

    ax_histx.plot(x_grid, px_beta, color='red', lw=2)
    ax_histy.plot(py_beta, y_grid, color='red', lw=2)

    ax_histx.axis('off')
    ax_histy.axis('off')

# タイトル全体の追加
main_title = "Gaussian Mixture: " + ("With Replica Exchange" if USE_EXCHANGE else "Without Replica Exchange (Independent Chains)")
fig.suptitle(main_title, fontsize=24, y=1.05)
plt.tight_layout()

