###################################################################################
# PROGRAM: HMC_GM.py
# DATE: 2026-8-30
# NOTICE: This program accompanies the book "マルコフ連鎖モンテカルロ法入門"
#                     by Koji Hukushima and Yoshihiko Nishikawa
###################################################################################
"""
========================================================================
【遊び方＆実験ガイド】HMC（ハミルトニアン・モンテカルロ）を混合正規分布で体感する
========================================================================
■ いろいろ書き換えて遊んでみよう

1. リープフロッグ回数 (L) を変えてみる
   - Lが小さすぎる（例: 2）と、MALAやランダムウォークと同じく移動距離が短くなる。
   - Lを大きくする（例: 20）と、等高線に沿ってように移動し、
     遠くの山へ一気にワープできるようになる。

2. 運動量のリフレッシュ率 (rho) の効果
   - rho = 1.0 : 毎回、ランダムな方向に動き直します（標準のHMC）。
   - rho < 1.0 (例: 0.6) : 前回の運動量を一部引き継ぐ。
     これにより、ジグザグに進むのを防ぎ、同じ方向へ進みやすくなる
     （Generalized HMC / 部分リフレッシュ法）

3. ステップサイズ (eps) の限界
   - epsを大きくしすぎると、シミュレーション誤差が大きくなりすぎて
     エネルギー保存則が崩れ、棄却（Reject）されやすくなる。
========================================================================
"""

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

# ==========================================
# 読者が遊べるパラメータ設定
# ==========================================
# (eps, L, rho) の組み合わせを3パターン設定して比較します
params = [
    {"eps": 0.05, "L": 10, "rho": 1.0},  # 標準HMC
    {"eps": 0.05, "L": 10, "rho": 0.8},  # 部分リフレッシュ (軽め)
    {"eps": 0.05, "L": 10, "rho": 0.6},  # 部分リフレッシュ (強め)
]

n_samples = 1500
initial_state = np.array([6.0, -2.0])    # スタート地点

# ==========================================
# 1. 目標分布 (2次元・混合ガウス) と勾配の定義
# ==========================================
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

inv_cov1, inv_cov2 = np.linalg.inv(cov1), np.linalg.inv(cov2)

def logp(x):
    """対数確率密度関数（HMCでは -logp がポテンシャルエネルギー U(x) となる）"""
    p1 = multivariate_normal.pdf(x, mu1, cov1)
    p2 = multivariate_normal.pdf(x, mu2, cov2)
    return np.log(pi1*p1 + pi2*p2 + 1e-12)

def grad_logp(x):
    """対数確率密度の勾配（HMCでは、これがボールを加速させる力 F となる）"""
    p1 = multivariate_normal.pdf(x, mu1, cov1)
    p2 = multivariate_normal.pdf(x, mu2, cov2)
    tot = pi1*p1 + pi2*p2 + 1e-12
    w1, w2 = pi1*p1/tot, pi2*p2/tot
    # ガウス分布の微分
    grad1 = -inv_cov1 @ (x - mu1)
    grad2 = -inv_cov2 @ (x - mu2)
    return w1 * grad1 + w2 * grad2

# ==========================================
# 2. リープフロッグ積分（物理シミュレーション）
# ==========================================
def leapfrog(x, p, eps, L):
    """力学の法則に従って、L歩だけボールを転がす"""
    # 運動量 p は確率が高い方向（grad_logp）へ加速される
    p = p + 0.5 * eps * grad_logp(x)
    for i in range(L):
        x = x + eps * p
        if i != L - 1:
            p = p + eps * grad_logp(x)
    p = p + 0.5 * eps * grad_logp(x)
    
    return x, -p  # MCMCの可逆性の条件を満たすために運動量を反転させる

# ==========================================
# 3. HMC 本体（部分リフレッシュ付き）
# ==========================================
def hmc_sampler(eps, L, rho, n_samples, x0):
    """rho=1: 完全リフレッシュ (HMC),  rho<1: 部分リフレッシュ (GHMC)"""
    x = x0.copy()
    p_prev = np.random.randn(2)  # 最初の運動量
    samples, acc = [], 0

    for _ in range(n_samples):
        # --- 1. 運動量のリフレッシュ (新しい勢いをつける) ---
        xi = np.random.randn(2)
        # rhoが1未満なら、前回の勢い(p_prev)をブレンドして引き継ぐ
        p = np.sqrt(1 - rho**2) * p_prev + rho * xi

        # --- 2. ハミルトン軌道 (ボールを転がす) ---
        x_new, p_new = leapfrog(x.copy(), p.copy(), eps, L)

        # --- 3. Metropolis 判定 (エネルギーが保存されていれば受容) ---
        # 全エネルギー H = ポテンシャルエネルギー(-logp) + 運動エネルギー(0.5*p^2)
        H_curr = -logp(x)     + 0.5 * np.sum(p**2)
        H_prop = -logp(x_new) + 0.5 * np.sum(p_new**2)
        
        if np.log(np.random.rand()) < H_curr - H_prop:
            x, p_prev = x_new, p_new      # 採択：次の周期は p_new を引き継ぐ
            acc += 1
        else:
            p_prev = -p                   # 棄却：運動量を反転させて持ち越す

        samples.append(x.copy())

    return np.asarray(samples), acc / n_samples

# ==========================================
# 4. シミュレーションとグラフ描画
# ==========================================
np.random.seed(42)

# 等高線描画用のグリッド
xv, yv = np.mgrid[-3:8:.05, -3:8:.05]
pos = np.dstack((xv, yv))
z = 0.5 * multivariate_normal.pdf(pos, mean=mu1, cov=cov1) + \
    0.5 * multivariate_normal.pdf(pos, mean=mu2, cov=cov2)

plt.rcParams['mathtext.fontset'] = 'cm'
plt.rcParams["font.size"] = 16
fig, axes = plt.subplots(1, 3, figsize=(16, 6))

for ax, p in zip(axes, params):
    eps, L, rho = p["eps"], p["L"], p["rho"]
    samples, acc_rate = hmc_sampler(eps, L, rho, n_samples, initial_state)
    
    # 等高線の描画
    ax.contour(xv, yv, z, levels=15, colors='black', alpha=0.3)
    
    # 軌跡とサンプルの描画
    ax.plot(samples[:, 0], samples[:, 1], '-', lw=0.5, alpha=0.3, color='gray')
    ax.plot(samples[:, 0], samples[:, 1], 'o', markersize=3, alpha=0.4, color='magenta')
    
    # スタート地点を赤い星マークで強調
    ax.plot(initial_state[0], initial_state[1], 
            marker='*', color='red', markersize=15, markeredgecolor='black', label='Start')
    
    # グラフの装飾
    ax.set_title(f"$\\epsilon$={eps}, $L$={L}, $\\rho$={rho}\nAcceptance: {acc_rate:.2f}", loc='left', fontsize=18)
    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.set_aspect('equal')
    
    if p == params[0]:
        ax.legend(loc='upper left')

plt.tight_layout()
plt.show()

