-
Notifications
You must be signed in to change notification settings - Fork 8
Expand file tree
/
Copy pathdqn_keras.py
More file actions
118 lines (95 loc) · 4.83 KB
/
Copy pathdqn_keras.py
File metadata and controls
118 lines (95 loc) · 4.83 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
from keras.layers import Dense, Activation, Conv2D, Flatten
from keras.models import Sequential, load_model
from keras.optimizers import Adam
import numpy as np
class Buffer(object):
def __init__(self, max_size, input_shape):
self.mem_size = max_size
self.mem_cntr = 0
self.state_memory = np.zeros((self.mem_size, *input_shape), dtype=np.float32)
self.new_state_memory = np.zeros((self.mem_size, *input_shape), dtype=np.float32)
self.action_memory = np.zeros(self.mem_size, dtype=np.int32)
self.reward_memory = np.zeros(self.mem_size, dtype=np.float32)
self.terminal_memory = np.zeros(self.mem_size, dtype=np.uint8)
def store_transition(self, state, action, reward, state_, done):
index = self.mem_cntr % self.mem_size
self.state_memory[index] = state
self.new_state_memory[index] = state_
self.action_memory[index] = action
self.reward_memory[index] = reward
self.terminal_memory[index] = done
self.mem_cntr += 1
def sample_buffer(self, batch_size):
max_mem = min(self.mem_cntr, self.mem_size)
batch = np.random.choice(max_mem, batch_size, replace=False)
states = self.state_memory[batch]
actions = self.action_memory[batch]
rewards = self.reward_memory[batch]
states_ = self.new_state_memory[batch]
terminal = self.terminal_memory[batch]
return states, actions, rewards, states_, terminal
def build_dqn(lr, n_actions, input_dims, fc1_dims):
model = Sequential()
'''
model.add(Conv2D(filters=32, kernel_size=8, strides=4, activation='relu', input_shape=(*input_dims,), data_format='channels_first'))
model.add(Conv2D(filters=64, kernel_size=4, strides=2, activation='relu', data_format='channels_first'))
model.add(Conv2D(filters=64, kernel_size=3, strides=1, activation='relu', data_format='channels_first'))
'''
model.add(Conv2D(filters=32, kernel_size=8, strides=4, activation='relu', input_shape=input_dims))
model.add(Conv2D(filters=64, kernel_size=4, strides=2, activation='relu'))
model.add(Conv2D(filters=64, kernel_size=3, strides=1, activation='relu'))
model.add(Flatten())
model.add(Dense(fc1_dims, activation='relu'))
model.add(Dense(n_actions))
model.compile(optimizer=Adam(lr=lr), loss='mean_squared_error')
return model
class Agent(object):
def __init__(self, alpha, gamma, n_actions, epsilon, batch_size, input_dims, replace_target=10000, eps_dec=0.996, eps_min=0.01,
mem_size=1000000, q_eval_fname='q_eval.h5', q_target_fname='q_next.h5'):
self.action_space = [i for i in range(n_actions)]
self.gamma = gamma
self.epsilon = epsilon
self.eps_dec = eps_dec
self.eps_min = eps_min
self.batch_size = batch_size
self.replace_target = replace_target
self.q_target_model_file = q_target_fname
self.q_eval_model_file = q_eval_fname
self.learn_step = 0
self.memory = Buffer(mem_size, input_dims)
self.q_eval = build_dqn(alpha, n_actions, input_dims, 512)
self.q_next = build_dqn(alpha, n_actions, input_dims, 512)
def replace_target_network(self):
if self.replace_target is not None and self.learn_step % self.replace_target == 0:
self.q_next.set_weights(self.q_eval.get_weights())
def store_transition(self, state, action, reward, new_state, done):
self.memory.store_transition(state, action, reward, new_state, done)
def choose_action(self, observation):
if np.random.random() < self.epsilon:
action = np.random.choice(self.action_space)
else:
state = np.array(observation, copy=False, dtype=np.float32)
actions = self.q_eval.predict(state)
action = np.argmax(actions)
return action
def learn(self):
if self.memory.mem_cntr > self.batch_size:
state, action, reward, new_state, done = self.memory.sample_buffer(self.batch_size)
self.replace_target_network()
q_eval = self.q_eval.predict(state)
q_next = self.q_next.predict(new_state)
q_next[done] = 0.0
q_target = q_eval[:]
indices = np.arange(self.batch_size)
q_target[indices, action] = reward + self.gamma*np.max(q_next,axis=1)
self.q_eval.train_on_batch(state, q_target)
self.epsilon = self.epsilon - self.eps_dec if self.epsilon > self.eps_min else self.eps_min
self.learn_step += 1
def save_models(self):
self.q_eval.save(self.q_eval_model_file)
self.q_next.save(self.q_target_model_file)
print('... saving models ...')
def load_models(self):
self.q_eval = load_model(self.q_eval_model_file)
self.q_next = load_model(self.q_target_model_file)
print('... loading models ...')