Code cleanup

This commit is contained in:
Shangtong Zhang
2018-01-31 21:00:19 -07:00
parent 3dcda7b5d5
commit ad9189e25a
4 changed files with 32 additions and 112 deletions
+8 -15
View File
@@ -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
+21 -15
View File
@@ -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()
+1 -29
View File
@@ -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)
+2 -53
View File
@@ -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