Minor update for DQN

This commit is contained in:
Shangtong Zhang
2017-12-21 22:22:36 -07:00
parent 8839614c44
commit 69b7ea53cc
5 changed files with 17 additions and 31 deletions
+9 -13
View File
@@ -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:
+6 -8
View File
@@ -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')
-2
View File
@@ -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:
+2 -4
View File
@@ -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)
-4
View File
@@ -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]