From 69b7ea53ccec9f250f4291a7dffb4ac2b34e5322 Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Thu, 21 Dec 2017 22:22:36 -0700 Subject: [PATCH] Minor update for DQN --- agent/DQN_agent.py | 22 +++++++++------------- main.py | 14 ++++++-------- network/base_network.py | 2 -- network/conv_network.py | 6 ++---- network/shallow_network.py | 4 ---- 5 files changed, 17 insertions(+), 31 deletions(-) diff --git a/agent/DQN_agent.py b/agent/DQN_agent.py index c7e18d3..199604b 100644 --- a/agent/DQN_agent.py +++ b/agent/DQN_agent.py @@ -16,29 +16,25 @@ import torch class DQNAgent: def __init__(self, config): self.config = config - self.learning_network = config.network_fn(config.optimizer_fn) - self.target_network = config.network_fn(config.optimizer_fn) + self.learning_network = config.network_fn() + self.target_network = config.network_fn() + self.optimizer = config.optimizer_fn(self.learning_network.parameters()) + self.criterion = nn.MSELoss() self.target_network.load_state_dict(self.learning_network.state_dict()) self.task = config.task_fn() self.replay = config.replay_fn() self.policy = config.policy_fn() self.total_steps = 0 - self.history_buffer = None def episode(self, deterministic=False): episode_start_time = time.time() state = self.task.reset() - if self.history_buffer is None: - self.history_buffer = [np.zeros_like(state)] * self.config.history_length - else: - self.history_buffer.pop(0) - self.history_buffer.append(state) + self.history_buffer = [state] * self.config.history_length state = np.vstack(self.history_buffer) total_reward = 0.0 steps = 0 while True: - value = self.learning_network.predict(np.stack([self.task.normalize_state(state)]), False) - value = value.cpu().data.numpy().flatten() + value = self.learning_network.predict(np.stack([self.task.normalize_state(state)]), True).flatten() if deterministic: action = np.argmax(value) elif self.total_steps < self.config.exploration_steps: @@ -97,10 +93,10 @@ class DQNAgent: actions = self.learning_network.to_torch_variable(actions, 'int64').unsqueeze(1) q = self.learning_network.predict(states, False) q = q.gather(1, actions).squeeze(1) - loss = self.learning_network.criterion(q, q_next) - self.learning_network.zero_grad() + loss = self.criterion(q, q_next) + self.optimizer.zero_grad() loss.backward() - self.learning_network.optimizer.step() + self.optimizer.step() if not deterministic and self.total_steps % self.config.target_network_update_freq == 0: self.target_network.load_state_dict(self.learning_network.state_dict()) if not deterministic and self.total_steps > self.config.exploration_steps: diff --git a/main.py b/main.py index c5f36b1..9b74e82 100644 --- a/main.py +++ b/main.py @@ -14,7 +14,7 @@ def dqn_cart_pole(): config = Config() config.task_fn = lambda: CartPole() config.optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001) - config.network_fn = lambda optimizer_fn: FCNet([8, 50, 200, 2], optimizer_fn) + config.network_fn = lambda: FCNet([8, 50, 200, 2]) # config.network_fn = lambda optimizer_fn: DuelingFCNet([8, 50, 200, 2], optimizer_fn) config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1) config.replay_fn = lambda: Replay(memory_size=10000, batch_size=10) @@ -75,7 +75,7 @@ def dqn_pixel_atari(name): config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False) action_dim = config.task_fn().action_dim config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01) - config.network_fn = lambda optimizer_fn: NatureConvNet(config.history_length, action_dim, optimizer_fn) + config.network_fn = lambda: NatureConvNet(config.history_length, action_dim) # config.network_fn = lambda optimizer_fn: DuelingNatureConvNet(config.history_length, n_actions, optimizer_fn) config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1) config.replay_fn = lambda: Replay(memory_size=1000000, batch_size=32, dtype=np.uint8) @@ -141,8 +141,7 @@ def dqn_fruit(): config.optimizer_fn = lambda params: torch.optim.SGD(params, 0.01, momentum=0.9) config.reward_weight = np.ones(10) / 10 config.hybrid_reward = False - config.network_fn = lambda optimizer_fn: FruitHRFCNet( - 98, 4, config.reward_weight, optimizer_fn) + config.network_fn = lambda: FruitHRFCNet(98, 4, config.reward_weight) config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1) config.replay_fn = lambda: Replay(memory_size=10000, batch_size=15) config.discount = 0.95 @@ -163,8 +162,7 @@ def hrdqn_fruit(): config.hybrid_reward = True config.reward_weight = np.ones(10) / 10 config.optimizer_fn = lambda params: torch.optim.SGD(params, 0.01, momentum=0.9) - config.network_fn = lambda optimizer_fn: FruitHRFCNet( - 98, 4, config.reward_weight, optimizer_fn) + config.network_fn = lambda optimizer_fn: FruitHRFCNet(98, 4, config.reward_weight) config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1) config.replay_fn = lambda: HybridRewardReplay(memory_size=10000, batch_size=15) config.discount = 0.95 @@ -280,7 +278,7 @@ if __name__ == '__main__': # logger.setLevel(logging.DEBUG) logger.setLevel(logging.INFO) - dqn_cart_pole() + # dqn_cart_pole() # async_cart_pole() # a3c_cart_pole() # a3c_continuous() @@ -290,7 +288,7 @@ if __name__ == '__main__': # dqn_fruit() # hrdqn_fruit() - # dqn_pixel_atari('PongNoFrameskip-v4') + dqn_pixel_atari('PongNoFrameskip-v4') # async_pixel_atari('PongNoFrameskip-v4') # a3c_pixel_atari('PongNoFrameskip-v4') diff --git a/network/base_network.py b/network/base_network.py index 748228f..21086a7 100644 --- a/network/base_network.py +++ b/network/base_network.py @@ -13,8 +13,6 @@ import numpy as np # Base class for all kinds of network class BasicNet: def __init__(self, optimizer_fn, gpu, LSTM=False): - if optimizer_fn is not None: - self.optimizer = optimizer_fn(self.parameters()) self.gpu = gpu and torch.cuda.is_available() self.LSTM = LSTM if self.gpu: diff --git a/network/conv_network.py b/network/conv_network.py index 4951b4f..331b734 100644 --- a/network/conv_network.py +++ b/network/conv_network.py @@ -15,8 +15,7 @@ class NatureConvNet(nn.Module, VanillaNet): self.conv3 = nn.Conv2d(64, 64, kernel_size=3, stride=1) self.fc4 = nn.Linear(7 * 7 * 64, 512) self.fc5 = nn.Linear(512, n_actions) - self.criterion = nn.MSELoss() - BasicNet.__init__(self, optimizer_fn, gpu) + BasicNet.__init__(self, None, gpu) def forward(self, x): x = self.to_torch_variable(x) @@ -37,8 +36,7 @@ class DuelingNatureConvNet(nn.Module, DuelingNet): self.fc4 = nn.Linear(7 * 7 * 64, 512) self.fc_advantage = nn.Linear(512, n_actions) self.fc_value = nn.Linear(512, 1) - self.criterion = nn.MSELoss() - BasicNet.__init__(self, optimizer_fn, gpu) + BasicNet.__init__(self, None, gpu) def forward(self, x): x = self.to_torch_variable(x) diff --git a/network/shallow_network.py b/network/shallow_network.py index 148a2bb..265a9cb 100644 --- a/network/shallow_network.py +++ b/network/shallow_network.py @@ -13,7 +13,6 @@ class FCNet(nn.Module, VanillaNet): self.fc1 = nn.Linear(dims[0], dims[1]) self.fc2 = nn.Linear(dims[1], dims[2]) self.fc3 = nn.Linear(dims[2], dims[3]) - self.criterion = nn.MSELoss() BasicNet.__init__(self, optimizer_fn, gpu) def forward(self, x): @@ -32,7 +31,6 @@ class DuelingFCNet(nn.Module, DuelingNet): self.fc2 = nn.Linear(dims[1], dims[2]) self.fc_value = nn.Linear(dims[2], 1) self.fc_advantage = nn.Linear(dims[2], dims[3]) - self.criterion = nn.MSELoss() BasicNet.__init__(self, optimizer_fn, gpu) def forward(self, x): @@ -67,7 +65,6 @@ class FruitHRFCNet(nn.Module, VanillaNet): hidden_size = 250 self.fc1 = nn.Linear(state_dim, hidden_size) self.fc2 = nn.ModuleList([nn.Linear(hidden_size, action_dim) for _ in head_weights]) - self.criterion = nn.MSELoss() self.head_weights = head_weights BasicNet.__init__(self, optimizer_fn, gpu) @@ -93,7 +90,6 @@ class FruitMultiStatesFCNet(nn.Module, BasicNet): hidden_size = 250 self.fc1 = nn.ModuleList([nn.Linear(state_dim, hidden_size) for _ in head_weights]) self.fc2 = nn.ModuleList([nn.Linear(hidden_size, action_dim) for _ in head_weights]) - self.criterion = nn.MSELoss() self.head_weights = head_weights self.state_dim = state_dim self.n_heads = head_weights.shape[0]