###################################################################################
# PROGRAM: MALA_GM.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
from scipy.stats import multivariate_normal

# ==========================================
# 読者が遊べるパラメータ設定
# ==========================================
# ★ステップ幅（epsilon）を変更して、MALAの挙動を観察しましょう！
# epsilons = [0.1, 0.5, 1.0] のように設定します。
epsilons = [0.1, 0.5, 1.0]

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

# ==========================================
# Step 1: ターゲット分布とその「勾配」の定義
# ==========================================
mu1 = np.array([0, 0])
mu2 = 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]])

def target_density(x):
    p1 = multivariate_normal.pdf(x, mean=mu1, cov=cov1)
    p2 = multivariate_normal.pdf(x, mean=mu2, cov=cov2)
    return 0.5 * p1 + 0.5 * p2

def log_target_density(x):
    # log計算時のゼロ割りを防ぐため微小値(1e-12)を足す
    return np.log(target_density(x) + 1e-12)

def grad_log_target_density(x):
    """対数確率密度関数の勾配（どちらへ行けば確率が高くなるか）を計算"""
    p1 = multivariate_normal.pdf(x, mean=mu1, cov=cov1)
    p2 = multivariate_normal.pdf(x, mean=mu2, cov=cov2)
    
    # ガウス分布の対数微分
    grad1 = -np.linalg.solve(cov1, x - mu1)
    grad2 = -np.linalg.solve(cov2, x - mu2)
    
    denom = 0.5 * p1 + 0.5 * p2 + 1e-12
    return (0.5 * p1 * grad1 + 0.5 * p2 * grad2) / denom

# ==========================================
# Step 2: MALA サンプリング
# ==========================================
def mala_sampling(epsilon, n_steps, x0):
    samples = []
    x = x0.copy()
    accepted = 0
    
    for _ in range(n_steps):
        # 1. 勾配方向へシフトした提案分布（ランジュバン方程式に基づく）
        grad_current = grad_log_target_density(x)
        mu_forward = x + 0.5 * epsilon * grad_current
        
        # 提案をサンプリング
        x_proposal = np.random.multivariate_normal(mu_forward, epsilon * np.eye(2))

        # 2. 受容確率（Acceptance Ratio）の計算
        log_p_current = log_target_density(x)
        log_p_proposal = log_target_density(x_proposal)

        # 逆方向の遷移確率 (x_proposal -> x)
        grad_proposal = grad_log_target_density(x_proposal)
        mu_reverse = x_proposal + 0.5 * epsilon * grad_proposal
        
        log_q_forward = multivariate_normal.logpdf(x_proposal, mean=mu_forward, cov=epsilon * np.eye(2))
        log_q_reverse = multivariate_normal.logpdf(x, mean=mu_reverse, cov=epsilon * np.eye(2))

        # MH比（対数スケールで計算）
        acceptance_log_ratio = (log_p_proposal + log_q_reverse) - (log_p_current + log_q_forward)
        
        # 3. 受容判定
        if np.log(np.random.rand()) < acceptance_log_ratio:
            x = x_proposal
            accepted += 1

        samples.append(x.copy())
        
    return np.array(samples), accepted / n_steps

# ==========================================
# Step 3: シミュレーションとグラフ描画
# ==========================================
import hashlib
def set_string_seed(seed_str):
    seed_int = int(hashlib.sha256(seed_str.encode('utf-8')).hexdigest(), 16) % (2**32)
    np.random.seed(seed_int)

# シードを固定して、パラメータを変えたときの違いを比較しやすくする
set_string_seed("I love MALA!")

plt.rcParams['mathtext.fontset'] = 'cm'
plt.rcParams["font.size"] = 16

fig, axes = plt.subplots(1, 3, figsize=(15, 5))

# 等高線描画用のグリッド作成
xv, yv = np.mgrid[-3:8:.05, -3:8:.05]
pos = np.dstack((xv, yv))
z = target_density(pos)

print(f"=== MALAサンプリング ===")
print(f"ステップ数: {n_samples}, 初期値: {initial_state}")

for ax, eps in zip(axes, epsilons):
    samples, acc_rate = mala_sampling(eps, n_samples, initial_state)
    print(f"ステップ幅 ε={eps:3.1f} -> 受容率: {acc_rate:.3f}")
    
    # ターゲット分布の等高線
    ax.contour(xv, yv, z, levels=15, cmap='gray', alpha=0.5)
    
    # 軌跡（線）とサンプル（点）の描画
    ax.plot(samples[:, 0], samples[:, 1], '-', lw=0.5, alpha=0.6, color='gray')
    ax.plot(samples[:, 0], samples[:, 1], 'o', markersize=4, alpha=0.3, 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}\nAcceptance Rate: {acc_rate:.2f}", loc='left')
    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 eps == epsilons[0]:
        ax.legend(loc='upper left')

plt.tight_layout()
plt.show()
