Skip to content

Instantly share code, notes, and snippets.

@vaibkumr
Created March 20, 2019 10:03
Show Gist options
  • Select an option

  • Save vaibkumr/5bd93710e118633e0793dc5d0b92b19a to your computer and use it in GitHub Desktop.

Select an option

Save vaibkumr/5bd93710e118633e0793dc5d0b92b19a to your computer and use it in GitHub Desktop.
import gym
import numpy as np
import time
"""
SARSA on policy learning python implementation.
This is a python implementation of the SARSA algorithm in the Sutton and Barto's book on
RL. It's called SARSA because - (state, action, reward, state, action). The only difference
between SARSA and Qlearning is that SARSA takes the next action based on the current policy
while qlearning takes the action with maximum utility of next state.
Using the simplest gym environment for brevity: https://gym.openai.com/envs/FrozenLake-v0/
"""
def init_q(s, a, type="ones"):
"""
@param s the number of states
@param a the number of actions
@param type random, ones or zeros for the initialization
"""
if type == "ones":
return np.ones((s, a))
elif type == "random":
return np.random.random((s, a))
elif type == "zeros":
return np.zeros((s, a))
def epsilon_greedy(Q, epsilon, n_actions, s, train=False):
"""
@param Q Q values state x action -> value
@param epsilon for exploration
@param s number of states
@param train if true then no random actions selected
"""
if train or np.random.rand() < epsilon:
action = np.argmax(Q[s, :])
else:
action = np.random.randint(0, n_actions)
return action
def expected_sarsa(alpha, gamma, epsilon, episodes, max_steps, n_tests, render = False, test=False):
"""
@param alpha learning rate
@param gamma decay factor
@param epsilon for exploration
@param max_steps for max step in each episode
@param n_tests number of test episodes
"""
env = gym.make('Taxi-v2')
n_states, n_actions = env.observation_space.n, env.action_space.n
Q = init_q(n_states, n_actions, type="ones")
timestep_reward = []
for episode in range(episodes):
print(f"Episode: {episode}")
total_reward = 0
s = env.reset()
t = 0
done = False
while t < max_steps:
if render:
env.render()
t += 1
a = epsilon_greedy(Q, epsilon, n_actions, s)
s_, reward, done, info = env.step(a)
total_reward += reward
if done:
Q[s, a] += alpha * ( reward - Q[s, a] )
else:
expected_value = np.mean(Q[s_,:])
# print(Q[s,:], sum(Q[s,:]), len(Q[s,:]), expected_value)
Q[s, a] += alpha * (reward + (gamma * expected_value) - Q[s, a])
s = s_
if done:
if True:
print(f"This episode took {t} timesteps and reward {total_reward}")
timestep_reward.append(total_reward)
break
if render:
print(f"Here are the Q values:\n{Q}\nTesting now:")
if test:
test_agent(Q, env, n_tests, n_actions)
return timestep_reward
def test_agent(Q, env, n_tests, n_actions, delay=0.1):
for test in range(n_tests):
print(f"Test #{test}")
s = env.reset()
done = False
epsilon = 0
total_reward = 0
while True:
time.sleep(delay)
env.render()
a = epsilon_greedy(Q, epsilon, n_actions, s, train=True)
print(f"Chose action {a} for state {s}")
s, reward, done, info = env.step(a)
total_reward += reward
if done:
print(f"Episode reward: {total_reward}")
time.sleep(1)
break
if __name__ =="__main__":
alpha = 0.1
gamma = 0.9
epsilon = 0.9
episodes = 1000
max_steps = 2500
n_tests = 20
timestep_reward = expected_sarsa(alpha, gamma, epsilon,
episodes, max_steps, n_tests,
render=False, test=True
)
print(timestep_reward)
@konichuvak

konichuvak commented Jun 12, 2020

Copy link
Copy Markdown

Hi @TimeTraveller-San

I think you got a bit confused about expected part of the algorithm.
expected_value = np.mean(Q[s_,:]) as you have it computes a simple average, which is different from a weighted average of Q values with the weights corresponding to the probabilities of actions in state s.

What you want is something like

a_greedy = np.argmax(Q[s_]) # note that argmax in numpy does not break ties randomly
probs = np.ones(n_actions) * epsilon / (n_actions - 1)
probs[a_greedy] = 1 - epsilon
expected_value = Q[s_].dot(probs)

This might be the case why Expected Sarsa performed worse than the other algorithms in your study :)

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment