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

V, g, a, V_true = np.zeros(5), 0.9, 0.1, 0.9 ** np.arange(3, -1, -1)
errs = []

for _ in range(300):
    s = 0
    while s < 4:
        sn, r = s + 1, float(s == 3)
        V[s] += a * (r + g * V[sn] - V[s])
        s = sn
    errs.append(np.sqrt(np.mean((V[:4] - V_true) ** 2)))

pd.Series(errs).plot(title="TD(0) RMSE Convergence", xlabel="Episode", ylabel="RMSE")
plt.show()