From ad9189e25a76a8f23c2b77fa1e077f6fbc5860ae Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Wed, 31 Jan 2018 21:00:19 -0700 Subject: [PATCH] Code cleanup --- component/task.py | 23 ++++++---------- main.py | 36 ++++++++++++++----------- network/conv_network.py | 30 +-------------------- network/shallow_network.py | 55 ++------------------------------------ 4 files changed, 32 insertions(+), 112 deletions(-) diff --git a/component/task.py b/component/task.py index 6554d74..55f97b1 100644 --- a/component/task.py +++ b/component/task.py @@ -37,23 +37,14 @@ class BasicTask: def random_action(self): return self.env.action_space.sample() -class MountainCar(BasicTask): - name = 'MountainCar-v0' - success_threshold = -110 - - def __init__(self, max_steps=200): - BasicTask.__init__(self, max_steps) - self.env = gym.make(self.name) - self.env._max_episode_steps = sys.maxsize - -class CartPole(BasicTask): - name = 'CartPole-v0' - success_threshold = 195 - - def __init__(self, max_steps=200): +class ClassicalControl(BasicTask): + def __init__(self, name='CartPole-v0', max_steps=200): BasicTask.__init__(self, max_steps) + self.name = name self.env = gym.make(self.name) self.env._max_episode_steps = sys.maxsize + self.action_dim = self.env.action_space.n + self.state_dim = self.env.observation_space.shape[0] class LunarLander(BasicTask): name = 'LunarLander-v2' @@ -62,10 +53,12 @@ class LunarLander(BasicTask): def __init__(self, max_steps=sys.maxsize): BasicTask.__init__(self, max_steps) self.env = gym.make(self.name) + self.action_dim = self.env.action_space.n + self.state_dim = self.env.observation_space.shape[0] class PixelAtari(BasicTask): def __init__(self, name, no_op, frame_skip, normalized_state=True, - frame_size=84, max_steps=sys.maxsize): + frame_size=84, max_steps=10000): BasicTask.__init__(self, max_steps) self.normalized_state = normalized_state self.name = name diff --git a/main.py b/main.py index 8aa9b89..fc2a65e 100644 --- a/main.py +++ b/main.py @@ -12,7 +12,7 @@ import model.action_conditional_video_prediction as acvp def dqn_cart_pole(): config = Config() - config.task_fn = lambda: CartPole() + config.task_fn = lambda: ClassicalControl('CartPole-v0', max_steps=200) config.optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001) config.network_fn = lambda: FCNet([8, 50, 200, 2]) # config.network_fn = lambda: DuelingFCNet([8, 50, 200, 2]) @@ -31,7 +31,7 @@ def dqn_cart_pole(): def async_cart_pole(): config = Config() - config.task_fn= lambda: CartPole() + config.task_fn = lambda: ClassicalControl('CartPole-v0', max_steps=200) config.optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) config.network_fn = lambda: FCNet([4, 50, 200, 2]) config.policy_fn = lambda: GreedyPolicy(epsilon=0.5, final_step=5000, min_epsilon=0.1) @@ -50,14 +50,18 @@ def async_cart_pole(): def a3c_cart_pole(): config = Config() - config.task_fn = lambda: CartPole() + name = 'CartPole-v0' + # name = 'MountainCar-v0' + config.task_fn = lambda: ClassicalControl(name, max_steps=200) + # config.task_fn = lambda: LunarLander() + task = config.task_fn() config.optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) - config.network_fn = lambda: ActorCriticFCNet(4, 2) + config.network_fn = lambda: ActorCriticFCNet(task.state_dim, task.action_dim) config.policy_fn = SamplePolicy config.worker = AdvantageActorCritic config.discount = 0.99 config.max_episode_length = 200 - config.num_workers = 16 + config.num_workers = 7 config.update_interval = 6 config.test_interval = 1 config.test_repetitions = 30 @@ -69,20 +73,23 @@ def a3c_cart_pole(): def a2c_cart_pole(): config = Config() - task_fn = lambda: CartPole(max_steps=200) - config.num_workers = 3 + name = 'CartPole-v0' + # name = 'MountainCar-v0' + task_fn = lambda: ClassicalControl(name, max_steps=200) + # task_fn = lambda: LunarLander() + task = task_fn() + config.num_workers = 5 config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers) config.optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) - config.network_fn = lambda: ActorCriticFCNet(4, 2) + config.network_fn = lambda: ActorCriticFCNet(task.state_dim, task.action_dim) config.policy_fn = SamplePolicy config.discount = 0.99 - config.test_interval = 20 + config.test_interval = 200 config.test_repetitions = 10 config.logger = Logger('./log', logger) config.gae_tau = 1.0 config.entropy_weight = 0.01 - config.rollout_length = 50 - config.success_threshold = 195 + config.rollout_length = 20 run_episodes(A2CAgent(config)) def dqn_pixel_atari(name): @@ -140,12 +147,11 @@ def a3c_pixel_atari(name): task = config.task_fn() config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.0001) config.network_fn = lambda: OpenAIActorCriticConvNet( - config.history_length, task.env.action_space.n, LSTM=True) + config.history_length, task.env.action_space.n, LSTM=False) config.reward_shift_fn = lambda r: np.sign(r) config.policy_fn = SamplePolicy config.worker = AdvantageActorCritic config.discount = 0.99 - config.max_episode_length = 10000 config.num_workers = 6 config.update_interval = 20 config.test_interval = 50000 @@ -309,8 +315,8 @@ if __name__ == '__main__': mkdir('log') os.system('export OMP_NUM_THREADS=1') os.system('export CUDA_VISIBLE_DEVICES=0') - # logger.setLevel(logging.DEBUG) - logger.setLevel(logging.INFO) + logger.setLevel(logging.DEBUG) + # logger.setLevel(logging.INFO) # dqn_cart_pole() # async_cart_pole() diff --git a/network/conv_network.py b/network/conv_network.py index b43da71..aa2adbe 100644 --- a/network/conv_network.py +++ b/network/conv_network.py @@ -47,40 +47,12 @@ class DuelingNatureConvNet(nn.Module, DuelingNet): phi = F.relu(self.fc4(y)) return phi - -# Network for pixel Atari game with actor critic -class ActorCriticNatureConvNet(nn.Module, ActorCriticNet): - def __init__(self, - in_channels, - n_actions, - xentropy_weight=0.01, - grad_threshold=40, - gpu=True): - super(ActorCriticNatureConvNet, self).__init__() - self.conv1 = nn.Conv2d(in_channels, 32, kernel_size=8, stride=4) - self.conv2 = nn.Conv2d(32, 64, kernel_size=4, stride=2) - self.conv3 = nn.Conv2d(64, 64, kernel_size=3, stride=1) - self.fc4 = nn.Linear(7 * 7 * 64, 512) - self.fc_actor = nn.Linear(512, n_actions) - self.fc_critic = nn.Linear(512, 1) - self.xentropy_weight = xentropy_weight - self.grad_threshold = grad_threshold - BasicNet.__init__(self, gpu) - - def forward(self, x): - x = self.to_torch_variable(x) - y = F.elu(self.conv1(x)) - y = F.elu(self.conv2(y)) - y = F.elu(self.conv3(y)) - y = y.view(y.size(0), -1) - return F.elu(self.fc4(y)) - class OpenAIActorCriticConvNet(nn.Module, ActorCriticNet): def __init__(self, in_channels, n_actions, LSTM=False, - gpu=True): + gpu=False): super(OpenAIActorCriticConvNet, self).__init__() self.conv1 = nn.Conv2d(in_channels, 32, 3, stride=2, padding=1) self.conv2 = nn.Conv2d(32, 32, 3, stride=2, padding=1) diff --git a/network/shallow_network.py b/network/shallow_network.py index 15a599b..b2e77bf 100644 --- a/network/shallow_network.py +++ b/network/shallow_network.py @@ -44,8 +44,8 @@ class DuelingFCNet(nn.Module, DuelingNet): class ActorCriticFCNet(nn.Module, ActorCriticNet): def __init__(self, state_dim, action_dim): super(ActorCriticFCNet, self).__init__() - hidden_size1 = 50 - hidden_size2 = 200 + hidden_size1 = 64 + hidden_size2 = 64 self.fc1 = nn.Linear(state_dim, hidden_size1) self.fc2 = nn.Linear(hidden_size1, hidden_size2) self.fc_actor = nn.Linear(hidden_size2, action_dim) @@ -58,54 +58,3 @@ class ActorCriticFCNet(nn.Module, ActorCriticNet): x = F.relu(self.fc1(x)) phi = self.fc2(x) return phi - -class FruitHRFCNet(nn.Module, VanillaNet): - def __init__(self, state_dim, action_dim, head_weights, gpu=True): - super(FruitHRFCNet, self).__init__() - 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.head_weights = head_weights - BasicNet.__init__(self, gpu) - - def forward(self, x, heads_only): - x = self.to_torch_variable(x) - x = x.view(x.size(0), -1) - x = F.relu(self.fc1(x)) - head_q = [fc(x) for fc in self.fc2] - if not heads_only: - q = [h * w for h, w in zip(head_q, self.head_weights)] - q = torch.stack(q, dim=0) - q = q.sum(0).squeeze(0) - return q - else: - return head_q - - def predict(self, x, heads_only): - return self.forward(x, heads_only) - -class FruitMultiStatesFCNet(nn.Module, BasicNet): - def __init__(self, state_dim, action_dim, head_weights, gpu=True): - super(FruitMultiStatesFCNet, self).__init__() - 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.head_weights = head_weights - self.state_dim = state_dim - self.n_heads = head_weights.shape[0] - BasicNet.__init__(self, gpu) - - def predict(self, x, merge): - head_q = [] - for i in range(self.n_heads): - q = self.to_torch_variable(x[:, i, :]) - q = self.fc1[i](q) - q = F.relu(q) - q = self.fc2[i](q) - head_q.append(q) - if merge: - q = [q * w for q, w in zip(head_q, self.head_weights)] - q = torch.stack(q, dim=0) - q = q.sum(0).squeeze(0) - return q - return head_q