MCPcopy Create free account
hub / github.com/rlcode/reinforcement-learning / QLearningAgent

Class QLearningAgent

1-grid-world/4-q_learning.py:7–46  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

5from env import Env
6
7class QLearningAgent:
8 def __init__(self, actions):
9 # actions = [0, 1, 2, 3]
10 self.actions = actions
11 self.learning_rate = 0.01
12 self.discount_factor = 0.9
13 self.epsilon = 0.1
14 self.q_table = defaultdict(lambda: [0.0, 0.0, 0.0, 0.0])
15
16 # update q function with sample <s, a, r, s'>
17 def learn(self, state, action, reward, next_state):
18 current_q = self.q_table[state][action]
19 # using Bellman Optimality Equation to update q function
20 new_q = reward + self.discount_factor * max(self.q_table[next_state])
21 self.q_table[state][action] += self.learning_rate * (new_q - current_q)
22
23 # get action for the state according to the q function table
24 # agent pick action of epsilon-greedy policy
25 def get_action(self, state):
26 if np.random.rand() < self.epsilon:
27 # take random action
28 action = np.random.choice(self.actions)
29 else:
30 # take action according to the q function table
31 state_action = self.q_table[state]
32 action = self.arg_max(state_action)
33 return action
34
35 @staticmethod
36 def arg_max(state_action):
37 max_index_list = []
38 max_value = state_action[0]
39 for index, value in enumerate(state_action):
40 if value > max_value:
41 max_index_list.clear()
42 max_value = value
43 max_index_list.append(index)
44 elif value == max_value:
45 max_index_list.append(index)
46 return random.choice(max_index_list)
47
48if __name__ == "__main__":
49 env = Env()

Callers 1

4-q_learning.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected