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)

def solve(pi):
    V, pol = np.zeros(S), np.arange(S) % A
    for _ in range(12):
        Q = np.column_stack([R[:, a] + g * P[:, a] @ V for a in range(A)])
        pol = np.argmax(Q, axis=1)
        V = np.linalg.solve(
            np.eye(S) - g * P[np.arange(S), pol],
            R[np.arange(S), pol]
        ) if pi else Q.max(axis=1)
        yield V.mean()

pd.DataFrame({
    "Policy Iteration": solve(True),
    "Value Iteration": solve(False)
}).plot(title="PI vs VI Convergence")
plt.show()