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

def run(off_policy, Q=np.zeros((5, 2))):
    pi = lambda s: np.random.randint(2) if np.random.rand() < 0.2 else np.argmax(Q[s])
    for _ in range(300):
        s, ep_r = 0, 0
        while s < 4:
            a = pi(s)
            sn = min(4, max(0, s + (1 if a == 1 else -1)))
            r = 1.0 if s == 3 and a == 1 else -0.1
            target = Q[sn].max() if off_policy else Q[sn, pi(sn)]
            Q[s, a] += 0.1 * (r + 0.9 * target - Q[s, a])
            s, ep_r = sn, ep_r + r
        yield ep_r

pd.DataFrame({
    "Q-Learning": run(True),
    "SARSA": run(False)
}).rolling(20).mean().plot(title="SARSA vs Q-Learning")
plt.show()