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

# ==========================================
# 共通のグラフ設定
# ==========================================
# チュートリアル環境での文字化けを防ぐため、標準的なフォント設定にしています
plt.rcParams['mathtext.fontset'] = 'cm'
plt.rcParams["font.size"] = 18

# ==========================================
# MH法によるサンプラー
# ==========================================
def mh_dice_sampler(x0, M):
    """一様なサイコロの目を生成するメトロポリス法"""
    x = x0
    samples = []
    
    # ターゲット分布（サイコロなので、どの目も等しく 1/6）
    def target_prob(state):
        return 1.0 / 6.0
        
    for _ in range(M):
        # 現在の状態を記録
        samples.append(x)
        
        # 1. 提案分布 q(x' | x) に従って次の候補 x' を選ぶ
        # (-1, 0, 1) を等確率で選び、1〜6の範囲をループさせる (mod 6)
        epsilon = np.random.choice([-1, 0, 1])
        x_prime = ((x + epsilon - 1) % 6) + 1  
        
        # 2. 採択確率 (Acceptance probability) alpha を計算
        # 今回は提案分布が対称なのでメトロポリスの判定基準を使用
        alpha = min(1.0, target_prob(x_prime) / target_prob(x))
        
        # 3. 確率 alpha で提案を採択 (Accept)、 1 - alpha で棄却 (Reject)
        u = np.random.rand()
        if u <= alpha:
            x = x_prime  # 採択: 状態を更新（今回は alpha=1 なので必ずここを通る）
        else:
            x = x        # 棄却: 現在の状態に留まる
            
    return samples

# ==========================================
# ヒストグラムの描画
# ==========================================
def plot_histograms(sample_sizes, x0=3):
    fig, axes = plt.subplots(1, len(sample_sizes), figsize=(15, 4), sharey=True)
    
    for ax, M in zip(axes, sample_sizes):
        samples = mh_dice_sampler(x0, M)
        
        # ヒストグラムの計算と描画
        counts, _ = np.histogram(samples, bins=np.arange(1, 8), density=True)
        ax.bar(np.arange(1, 7), counts, width=0.8, align='center', 
               edgecolor='black', color='magenta', alpha=0.3)
        
        # 理論値（1/6）のライン
        ax.axhline(1/6, color='magenta', linestyle='--')
        
        ax.set_xticks(np.arange(1, 7))
        ax.set_title(f'$M = {M}$')
        ax.set_ylim(0, 0.41)
        
    # 最初のグラフのみY軸ラベルを設定
    axes[0].set_ylabel('Probability')
    
    plt.tight_layout()
    plt.show()

# ==========================================
# トレースプロット（対数スケール）の描画
# ==========================================
def plot_trace_logx(samples):
    steps = np.arange(1, len(samples) + 1)
    
    fig, ax = plt.subplots(figsize=(12, 4))
    
    ax.plot(steps, samples, marker='o', markersize=3, 
            linestyle='-', color='magenta', alpha=0.3)
    
    ax.set_xscale('log')
    ax.set_xlabel('MC Steps')
    ax.set_ylabel('State (Dice value)')
    ax.set_yticks(range(1, 7))
    ax.grid(True, which="both", ls="--", alpha=0.5)
    
    plt.tight_layout()
    plt.show()

# ==========================================
# 実行部分
# ==========================================
# 1. サンプルサイズごとのヒストグラムの変化
print("=== ヒストグラムの描画 ===")
plot_histograms(sample_sizes=[1, 10, 100, 1000, 10000])

# 2. M=10000 のトレースプロット
print("=== トレースプロットの描画 ===")
samples = mh_dice_sampler(x0=3, M=10000)
plot_trace_logx(samples)

