mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Code cleanup
This commit is contained in:
+8
-15
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user