class MCControl(MCPrediction):
def __init__(self, env, gamma=0.9, epsilon=0.1):
super().__init__(env, gamma)
self.epsilon = epsilon
self.Q = np.zeros((env.observation_space.n, env.action_space.n))
self.returns = defaultdict(list)
def select_action(self, policy, state):
if np.random.rand() < self.epsilon:
return self.env.action_space.sample()
else:
return np.argmax(self.Q[state])
def update_value(self, episode):
G = 0
for t in reversed(range(len(episode))):
state, action, reward, next_state, done = episode[t]
G = self.gamma * G + reward
if (state, action) not in [(x[0], x[1]) for x in episode[:t]]:
self.returns[(state, action)].append(G)
self.Q[state][action] = np.mean(self.returns[(state, action)])
def improve_policy(self):
policy = {}
for state in range(self.env.observation_space.n):
policy[state] = np.argmax(self.Q[state])
return policy
def evaluate_policy(self, policy, num_episodes=1000):
nS = self.env.observation_space.n
nA = self.env.action_space.n
self.Q_track = np.zeros((nS, nA, num_episodes))
for e in range(num_episodes):
episode = self.generate_episode(policy)
self.update_value(episode)
self.Q_track[:, :, e] = self.Q.copy()
return self.Q, self.Q_track
def control(self, policy, num_episodes=1000):
for _ in range(num_episodes):
episode = self.generate_episode()
self.update_value(episode)
return self.improve_policy()
class SARSA(MCControl):
def __init__(self, env, gamma=0.9, alpha=0.1, epsilon=0.1):
super().__init__(env, gamma, epsilon)
self.alpha = alpha
def generate_episode(self, policy):
episode = []
state, info = self.env.reset()
action = self.select_action(policy, state)
done = False
while not done:
next_state, reward, done, _, _ = self.env.step(action)
next_action = self.select_action(policy, next_state)
episode.append((state, action, reward, next_state, next_action, done))
state, action = next_state, next_action
return episode
def update_value(self, episode):
for t in range(len(episode)):
state, action, reward, next_state, next_action, done = episode[t]
td_target = reward + (self.gamma * self.Q[next_state][next_action] * (not done))
td_error = td_target - self.Q[state][action]
self.Q[state][action] += self.alpha * td_error
def control(self, policy, num_episodes=1000):
for _ in range(num_episodes):
episode = self.generate_episode(policy)
self.update_value(episode)
return self.improve_policy()