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

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

pd.Series(sarsa()).rolling(20).mean().plot(title="SARSA Control Learning Curve")
plt.show()