import numpy as np
import pandas as pd
import matplotlib.pyplot as plt

S, A, g = 5, 2, 0.9
P, R = np.random.dirichlet(np.ones(S), size=(S, A)), np.random.randn(S, A)
pi, history = np.zeros(S, dtype=int), []

for _ in range(10):
    V = np.linalg.solve(np.eye(S) - g * P[np.arange(S), pi], R[np.arange(S), pi])
    Q = np.column_stack([R[:, a] + g * P[:, a] @ V for a in range(A)])
    pi = np.argmax(Q, axis=1)
    history.append(V.mean())

pd.Series(history).plot(title="Policy Iteration Value Convergence", xlabel="Iteration", ylabel="Mean V")
plt.show()