###################################################################################
# PROGRAM: GibbsSampler_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

# ==========================================
# 遊べるパラメータ設定
# ==========================================
# ★ここの「距離」を変更して、グラフ右側のZの推移がどう変わるか観察しよう
# 例1: distance = 4.0 （山が近いので、頻繁にクラスタ間を行き来する）
# 例2: distance = 8.0 （山が離れすぎると、一度片方に落ちたら抜け出せなくなる）
distance = 4.0  
n_samples = 5000  # トレースを見やすくするため、少し少なめに設定

# ==========================================
# Step 1: ターゲット分布（混合ガウス分布）のパラメータ定義
# ==========================================
mu1 = np.array([0.0, 0.0])
mu2 = np.array([distance, distance]) # 距離に応じて右上の山を移動
cov1 = np.array([[1.0, 0.9], [0.9, 1.5]])  # 左下の山
cov2 = np.array([[1.0, -0.8], [-0.8, 1.0]]) # 右上の山

# ==========================================
# Step 2: 隠れ変数 Z を用いた Gibbsサンプラー
# ==========================================
def gibbs_sampler(n_steps):
    X_samples = []
    Z_samples = []
    
    # 初期化
    X = np.array([2.0, 2.0])
    Z = 0  
    
    for _ in range(n_steps):
        # ------------------------------------------------
        # 1. 現在の X が与えられた下での Z の更新 (Z | X)
        # ------------------------------------------------
        # 現在のXが、それぞれの山から発生した確率密度を計算
        p0 = multivariate_normal.pdf(X, mean=mu1, cov=cov1) * 0.5
        p1 = multivariate_normal.pdf(X, mean=mu2, cov=cov2) * 0.5
        
        # Z = 1 (右上の山) になる確率
        prob_z1 = p1 / (p0 + p1)
        
        # コイン投げ（ベルヌーイ試行）で新しい Z を決定
        Z = 1 if np.random.rand() < prob_z1 else 0

        # ------------------------------------------------
        # 2. 現在の Z が与えられた下での X の更新 (X | Z)
        # ------------------------------------------------
        # 選択された山の分布から、直接新しい X を生成
        if Z == 0:
            X = np.random.multivariate_normal(mu1, cov1)
        else:
            X = np.random.multivariate_normal(mu2, cov2)

        X_samples.append(X.copy())
        Z_samples.append(Z)

    return np.array(X_samples), np.array(Z_samples)

# ==========================================
# 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 Gibbs Sampling!")
X_samples, Z_samples = gibbs_sampler(n_samples)

# 描画設定
plt.rcParams['mathtext.fontset'] = 'cm'
plt.rcParams["font.size"] = 16

fig, axes = plt.subplots(1, 2, figsize=(14, 5))

# --- 左図: X の散布図 ---
x1 = X_samples[:, 0]
x2 = X_samples[:, 1]
axes[0].scatter(x1[Z_samples == 0], x2[Z_samples == 0], color='magenta', alpha=0.5, s=15, label='Z = 0')
axes[0].scatter(x1[Z_samples == 1], x2[Z_samples == 1], color='gray', alpha=0.5, s=15, label='Z = 1')
axes[0].set_xlabel("$X_1$")
axes[0].set_ylabel("$X_2$")

# 距離に応じて描画範囲を動的に調整
lim_max = max(8, distance + 4)
axes[0].set_xlim(-3, lim_max)
axes[0].set_ylim(-3, lim_max) 
axes[0].set_aspect('equal')
axes[0].legend(loc='upper left')
axes[0].set_title("Gibbs samples of X")

# --- 右図: Z の推移 (Trace) ---
# Z=0,1 を見やすくするため 1,2 にシフトして描画
axes[1].plot(np.arange(1, n_samples + 1), Z_samples + 1, color='black', lw=1.5, alpha=0.8) 
axes[1].set_xlabel("MC Steps")
axes[1].set_ylabel("Latent variable $Z$")
axes[1].set_yticks([1, 2])
axes[1].set_yticklabels(['Z = 0', 'Z = 1'])
axes[1].set_ylim(0.5, 2.5)
axes[1].set_title("Trace of Z")

plt.tight_layout()
plt.show()
