diff --git a/async_agent.py b/async_agent.py index bb6ac5e..7233051 100644 --- a/async_agent.py +++ b/async_agent.py @@ -12,28 +12,32 @@ import time from task import * from network import * from torch.autograd import Variable -import thread +import threading class AsyncAgent: def __init__(self, task_fn, network_fn, optimizer_fn, policy_fn, discount, step_limit, target_network_update_freq, n_workers): self.network_fn = network_fn self.learning_network = network_fn() - self.learning_network.share_memory() + # self.learning_network.share_memory() self.target_network = network_fn() - self.target_network.share_memory() + # self.target_network.share_memory() self.target_network.load_state_dict(self.learning_network.state_dict()) - self.optimizer_fn = optimizer_fn + # self.optimizer_fn = optimizer_fn # self.optimizer = optimizer_fn(self.learning_network.parameters()) self.task_fn = task_fn self.step_limit = step_limit self.discount = discount self.target_network_update_freq = target_network_update_freq - self.policy = policy_fn() - self.total_steps = mp.Value('i', 0) - self.lock = mp.Lock() + # self.policy = policy_fn() + self.policy_fn = policy_fn + # self.total_steps = mp.Value('i', 0) + self.total_steps = 0 + self.lock = threading.Lock() + # self.lock = mp.Lock() self.n_workers = n_workers + self.batch_size = 5 def async_update(self, worker_network, optimizer): with self.lock: @@ -46,56 +50,52 @@ class AsyncAgent: # print list(self.learning_network.parameters())[1].data def worker(self, id): - worker_network = self.network_fn() + # worker_network = self.network_fn() task = self.task_fn() episode = 0 - optimizer = self.optimizer_fn(self.learning_network.parameters()) + # optimizer = self.optimizer_fn(self.learning_network.parameters()) + policy = self.policy_fn() + terminal = True + episode_steps = 0 while True: - worker_network.load_state_dict(self.learning_network.state_dict()) - worker_network.zero_grad() - state = np.reshape(task.reset(), (1, -1)) - steps = 0 - total_reward = 0 - while not self.step_limit or steps < self.step_limit: - value = worker_network.predict(state) - action = self.policy.sample(value.flatten()) - next_state, reward, done, info = task.step(action) - next_state = np.reshape(next_state, (1, -1)) - q_next = self.learning_network.predict(next_state) - # q_next = self.target_network.predict(next_state) - q_next = np.max(q_next, axis=1) - if done: - q_next = 0 - q_next = reward + self.discount * q_next - value[0, action] = q_next - # if (steps > 0): - # print list(worker_network.parameters())[1].grad.data - # worker_network.zero_grad() - worker_network.gradient(state, value) - # print list(worker_network.parameters())[1].grad - # worker_network.zero_grad() - # worker_network.gradient(state, value) - # print list(worker_network.parameters())[1].grad - # worker_network.gradient(state, value) - # print list(worker_network.parameters())[1].grad - - steps += 1 - total_reward += reward + batch_states, batch_actions, batch_rewards = [], [], [] + if terminal: + if id == 0: + print 'episode %d, epsilon %f, steps: %d' % \ + (episode, policy.epsilon, episode_steps) + episode_steps = 0 + episode += 1 + policy.update_epsilon() + terminal = False + state = task.reset() + state = np.reshape(state, (1, -1)) + while not terminal and len(batch_states) < self.batch_size: + episode_steps += 1 with self.lock: - self.total_steps.value += 1 - if self.total_steps.value % self.target_network_update_freq == 0: - self.target_network.load_state_dict(self.learning_network.state_dict()) - if done: - break - state = next_state - print 'worker %d, episode %d, rewards %d' % (id, episode, total_reward) - episode += 1 - self.async_update(worker_network, optimizer) - self.policy.update_epsilon() + self.total_steps += 1 + batch_states.append(state) + value = self.learning_network.predict(state) + action = policy.sample(value.flatten()) + batch_actions.append(action) + state, reward, terminal, _ = task.step(action) + state = np.reshape(state, (1, -1)) + if not terminal: + q_next = np.max(self.target_network.predict(state)) + reward += self.discount * q_next + batch_rewards.append(reward) + + self.learning_network.learn_from_raw(np.vstack(batch_states), batch_actions, batch_rewards) + + if self.total_steps % self.target_network_update_freq == 0: + with self.lock: + self.target_network.load_state_dict(self.learning_network.state_dict()) def run(self): - procs = [mp.Process(target=self.worker, args=(i, )) for i in range(self.n_workers)] + procs = [threading.Thread(target=self.worker, args=(i, )) for i in range(self.n_workers)] + # procs = [mp.Process(target=self.worker, args=(i, )) for i in range(self.n_workers)] for p in procs: p.start() + # while True: + # time.sleep(0.01) for p in procs: p.join() # print list(self.learning_network.parameters())[1].data @@ -103,22 +103,25 @@ class Test: def __init__(self): # self.data = np.zeros((2, 3)) self.data = torch.zeros((2, 3)) - self.data.share_memory_() + # self.data.share_memory_() print self.data - self.val = mp.Value('i', 0) - self.lock = mp.Lock() + # self.val = mp.Value('i', 0) + # self.lock = mp.Lock() + self.lock = threading.Lock() def fun(self): # self.data += rank for i in range(50): time.sleep(0.01) - with self.lock: - self.val.value += 1 + # with self.lock: + # self.val.value += 1 + self.data += 1 def run(self): processes = [] for i in range(3): - processes.append(mp.Process(target=self.fun)) + # processes.append(mp.Process(target=self.fun)) + processes.append(threading.Thread(target=self.fun)) for p in processes: p.start() for p in processes: @@ -143,13 +146,17 @@ class TestNet(nn.Module): if __name__ == '__main__': task_fn = lambda: CartPole() - network_fn = lambda: FullyConnectedNet([4, 50, 200, 2]) - optimizer_fn = lambda params: torch.optim.SGD(params, 0.01) - policy_fn = lambda: GreedyPolicy(epsilon=0.5, decay_factor=0.99, min_epsilon=0.01) + optimizer_fn = lambda params: torch.optim.SGD(params, 0.001) + network_fn = lambda: FullyConnectedNet([4, 50, 200, 2], optimizer_fn = optimizer_fn) + policy_fn = lambda: GreedyPolicy(epsilon=1.0, end_episode=500, min_epsilon=0.1) # config = {'discount': 0.99, 'step_limit': 5000, 'target_network_update_freq': 200} - agent = AsyncAgent(task_fn, network_fn, optimizer_fn, policy_fn, 0.99, 0, 200, 6) + agent = AsyncAgent(task_fn, network_fn, optimizer_fn, policy_fn, 0.99, 0, 200, 8) agent.run() + # t = Test() + # t.run() + # print t.data + # t = TestNet() # x = Variable(torch.from_numpy(np.array([[0.1, 0.2]], dtype='float32'))) # target = Variable(torch.from_numpy(np.array([1], dtype='float32'))) diff --git a/network.py b/network.py index ee656b7..ea8a8bd 100644 --- a/network.py +++ b/network.py @@ -43,12 +43,21 @@ class FullyConnectedNet(nn.Module): def predict(self, x): return self.forward(x).cpu().data.numpy() + def learn_from_raw(self, x, actions, rewards): + y = self.forward(x) + target = np.copy(y.data.numpy()) + target[np.arange(target.shape[0]), actions] = np.asarray(rewards) + target = Variable(torch.from_numpy(target)) + loss = self.criterion(y, target) + self.zero_grad() + loss.backward() + self.optimizer.step() + def learn(self, x, target): target = torch.from_numpy(target) if self.gpu: target = target.cuda() target = Variable(target) - y = self.forward(x) loss = self.criterion(y, target) self.zero_grad() diff --git a/policy.py b/policy.py index f3da4fe..4c8025a 100644 --- a/policy.py +++ b/policy.py @@ -7,10 +7,11 @@ import numpy as np class GreedyPolicy: - def __init__(self, epsilon, decay_factor, min_epsilon): - self.epsilon = epsilon - self.decay_factor = decay_factor + def __init__(self, epsilon, end_episode, min_epsilon): + self.init_epsilon = self.epsilon = epsilon + self.current_episode = 0 self.min_epsilon = min_epsilon + self.end_episode = end_episode def sample(self, state): if np.random.rand() < self.epsilon: @@ -18,5 +19,6 @@ class GreedyPolicy: return np.argmax(state) def update_epsilon(self): - if self.epsilon > self.min_epsilon: - self.epsilon *= self.decay_factor + self.epsilon = self.init_epsilon - float(self.current_episode) / self.end_episode * (self.init_epsilon - self.min_epsilon) + self.epsilon = max(self.epsilon, self.min_epsilon) + self.current_episode += 1