"""迷你 2 状态 MDP：Value Iteration 教学示例。

状态: S0, S1
动作: stay, go
目的: 展示 model-based DP 如何算出 V* 与 π*
"""
from __future__ import annotations

import numpy as np

# 状态与动作
STATES = ["S0", "S1"]
ACTIONS = ["stay", "go"]
S, A = range(len(STATES)), range(len(ACTIONS))

# 转移 P[s, a, s'] 与奖励 R[s, a, s']
# 约定:
#   stay: 大概率留在原地，小奖励
#   go:   尝试转移到另一状态，成功有较高奖励
P = np.zeros((2, 2, 2))
R = np.zeros((2, 2, 2))

# S0, stay -> 多半留在 S0
P[0, 0, 0], P[0, 0, 1] = 0.9, 0.1
R[0, 0, 0], R[0, 0, 1] = 1.0, 0.0
# S0, go -> 多半去 S1
P[0, 1, 0], P[0, 1, 1] = 0.2, 0.8
R[0, 1, 0], R[0, 1, 1] = 0.0, 5.0
# S1, stay
P[1, 0, 0], P[1, 0, 1] = 0.1, 0.9
R[1, 0, 0], R[1, 0, 1] = 0.0, 1.0
# S1, go -> 多半回 S0
P[1, 1, 0], P[1, 1, 1] = 0.8, 0.2
R[1, 1, 0], R[1, 1, 1] = 3.0, 0.0


def value_iteration(gamma: float = 0.9, theta: float = 1e-8, max_iter: int = 1000):
    V = np.zeros(2)
    history = []

    for i in range(max_iter):
        delta = 0.0
        V_new = np.zeros_like(V)
        for s in S:
            q_sa = []
            for a in A:
                q = np.sum(P[s, a] * (R[s, a] + gamma * V))
                q_sa.append(q)
            V_new[s] = max(q_sa)
            delta = max(delta, abs(V_new[s] - V[s]))
        V = V_new
        history.append(V.copy())
        if delta < theta:
            print(f"Value Iteration 在第 {i + 1} 轮收敛")
            break

    # 由 V* 导出贪心策略
    pi = []
    Q = np.zeros((2, 2))
    for s in S:
        for a in A:
            Q[s, a] = np.sum(P[s, a] * (R[s, a] + gamma * V))
        pi.append(ACTIONS[int(np.argmax(Q[s]))])

    return V, Q, pi, history


if __name__ == "__main__":
    V_star, Q_star, pi_star, hist = value_iteration()
    print("V* =", {STATES[i]: round(float(V_star[i]), 4) for i in S})
    print("Q* =")
    for s in S:
        for a in A:
            print(f"  Q({STATES[s]}, {ACTIONS[a]}) = {Q_star[s, a]:.4f}")
    print("π* =", {STATES[i]: pi_star[i] for i in S})
