mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-11 11:53:01 +08:00
Minor refactor
This commit is contained in:
+6
-4
@@ -62,8 +62,10 @@ class DDPGAgent:
|
|||||||
while not self.step_limit or steps < self.step_limit:
|
while not self.step_limit or steps < self.step_limit:
|
||||||
action = self.actor.predict(np.stack([state])).flatten()
|
action = self.actor.predict(np.stack([state])).flatten()
|
||||||
if not deterministic:
|
if not deterministic:
|
||||||
action += self.random_process.sample()
|
if self.total_steps < self.exploration_steps:
|
||||||
action = np.clip(action, -1, 1)
|
action = np.random.uniform(-1, 1, action.shape)
|
||||||
|
else:
|
||||||
|
action += self.random_process.sample()
|
||||||
next_state, reward, done, info = self.task.step(action)
|
next_state, reward, done, info = self.task.step(action)
|
||||||
if not deterministic:
|
if not deterministic:
|
||||||
self.replay.feed([state, action, reward, next_state, int(done)])
|
self.replay.feed([state, action, reward, next_state, int(done)])
|
||||||
@@ -122,7 +124,7 @@ class DDPGAgent:
|
|||||||
|
|
||||||
if self.test_interval and ep % self.test_interval == 0:
|
if self.test_interval and ep % self.test_interval == 0:
|
||||||
self.logger.info('Testing...')
|
self.logger.info('Testing...')
|
||||||
self.save('data/%sdqn-model-%s.bin' % (self.tag, self.task.name))
|
self.save('data/%sddpg-model-%s.bin' % (self.tag, self.task.name))
|
||||||
test_rewards = []
|
test_rewards = []
|
||||||
for _ in range(self.test_repetitions):
|
for _ in range(self.test_repetitions):
|
||||||
test_rewards.append(self.episode(True))
|
test_rewards.append(self.episode(True))
|
||||||
@@ -130,7 +132,7 @@ class DDPGAgent:
|
|||||||
avg_test_rewards.append(avg_reward)
|
avg_test_rewards.append(avg_reward)
|
||||||
self.logger.info('Avg reward %f(%f)' % (
|
self.logger.info('Avg reward %f(%f)' % (
|
||||||
avg_reward, np.std(test_rewards) / np.sqrt(self.test_repetitions)))
|
avg_reward, np.std(test_rewards) / np.sqrt(self.test_repetitions)))
|
||||||
with open('data/%sdqn-statistics-%s.bin' % (self.tag, self.task.name), 'wb') as f:
|
with open('data/%sddpg-statistics-%s.bin' % (self.tag, self.task.name), 'wb') as f:
|
||||||
pickle.dump({'rewards': rewards,
|
pickle.dump({'rewards': rewards,
|
||||||
'test_rewards': avg_test_rewards}, f)
|
'test_rewards': avg_test_rewards}, f)
|
||||||
if avg_reward > self.task.success_threshold:
|
if avg_reward > self.task.success_threshold:
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ class DQNAgent:
|
|||||||
double_q,
|
double_q,
|
||||||
test_interval,
|
test_interval,
|
||||||
test_repetitions,
|
test_repetitions,
|
||||||
|
tag,
|
||||||
logger):
|
logger):
|
||||||
self.learning_network = network_fn(optimizer_fn)
|
self.learning_network = network_fn(optimizer_fn)
|
||||||
self.target_network = network_fn(optimizer_fn)
|
self.target_network = network_fn(optimizer_fn)
|
||||||
@@ -45,7 +46,7 @@ class DQNAgent:
|
|||||||
self.test_repetitions = test_repetitions
|
self.test_repetitions = test_repetitions
|
||||||
self.history_buffer = None
|
self.history_buffer = None
|
||||||
self.double_q = double_q
|
self.double_q = double_q
|
||||||
self.tag = ''
|
self.tag = tag
|
||||||
|
|
||||||
def episode(self, deterministic=False):
|
def episode(self, deterministic=False):
|
||||||
episode_start_time = time.time()
|
episode_start_time = time.time()
|
||||||
@@ -62,6 +63,8 @@ class DQNAgent:
|
|||||||
value = self.learning_network.predict(np.stack([self.task.normalize_state(state)]), True)
|
value = self.learning_network.predict(np.stack([self.task.normalize_state(state)]), True)
|
||||||
if deterministic:
|
if deterministic:
|
||||||
action = np.argmax(value.flatten())
|
action = np.argmax(value.flatten())
|
||||||
|
elif self.total_steps < self.explore_steps:
|
||||||
|
action = np.random.randint(0, len(value.flatten()))
|
||||||
else:
|
else:
|
||||||
action = self.policy.sample(value.flatten())
|
action = self.policy.sample(value.flatten())
|
||||||
next_state, reward, done, info = self.task.step(action)
|
next_state, reward, done, info = self.task.step(action)
|
||||||
+2
-1
@@ -31,6 +31,7 @@ class AsyncAgent:
|
|||||||
test_interval,
|
test_interval,
|
||||||
test_repetitions,
|
test_repetitions,
|
||||||
history_length,
|
history_length,
|
||||||
|
tag,
|
||||||
logger):
|
logger):
|
||||||
self.network_fn = network_fn
|
self.network_fn = network_fn
|
||||||
self.learning_network = network_fn()
|
self.learning_network = network_fn()
|
||||||
@@ -58,7 +59,7 @@ class AsyncAgent:
|
|||||||
self.test_repetitions = test_repetitions
|
self.test_repetitions = test_repetitions
|
||||||
self.logger = logger
|
self.logger = logger
|
||||||
self.history_length = history_length
|
self.history_length = history_length
|
||||||
self.tag = ''
|
self.tag = tag
|
||||||
|
|
||||||
def deterministic_episode(self, task, network):
|
def deterministic_episode(self, task, network):
|
||||||
state = task.reset()
|
state = task.reset()
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
from async_agent import *
|
from async_agent import *
|
||||||
from dqn_agent import *
|
from DQN_agent import *
|
||||||
from DDPG_agent import *
|
from DDPG_agent import *
|
||||||
import logging
|
import logging
|
||||||
import traceback
|
import traceback
|
||||||
@@ -15,7 +15,7 @@ def dqn_cart_pole():
|
|||||||
config['replay_fn'] = lambda: Replay(memory_size=10000, batch_size=10)
|
config['replay_fn'] = lambda: Replay(memory_size=10000, batch_size=10)
|
||||||
config['discount'] = 0.99
|
config['discount'] = 0.99
|
||||||
config['target_network_update_freq'] = 200
|
config['target_network_update_freq'] = 200
|
||||||
config['step_limit'] = 0
|
config['step_limit'] = 200
|
||||||
config['explore_steps'] = 1000
|
config['explore_steps'] = 1000
|
||||||
config['logger'] = gym.logger
|
config['logger'] = gym.logger
|
||||||
config['history_length'] = 2
|
config['history_length'] = 2
|
||||||
@@ -23,6 +23,7 @@ def dqn_cart_pole():
|
|||||||
config['test_repetitions'] = 50
|
config['test_repetitions'] = 50
|
||||||
# config['double_q'] = True
|
# config['double_q'] = True
|
||||||
config['double_q'] = False
|
config['double_q'] = False
|
||||||
|
config['tag'] = ''
|
||||||
agent = DQNAgent(**config)
|
agent = DQNAgent(**config)
|
||||||
agent.run()
|
agent.run()
|
||||||
|
|
||||||
@@ -37,13 +38,14 @@ def async_cart_pole():
|
|||||||
# config['worker_fn'] = OneStepSarsa
|
# config['worker_fn'] = OneStepSarsa
|
||||||
config['discount'] = 0.99
|
config['discount'] = 0.99
|
||||||
config['target_network_update_freq'] = 200
|
config['target_network_update_freq'] = 200
|
||||||
config['step_limit'] = 0
|
config['step_limit'] = 200
|
||||||
config['n_workers'] = 16
|
config['n_workers'] = 16
|
||||||
config['update_interval'] = 6
|
config['update_interval'] = 6
|
||||||
config['test_interval'] = 4000
|
config['test_interval'] = 4000
|
||||||
config['test_repetitions'] = 50
|
config['test_repetitions'] = 50
|
||||||
config['history_length'] = 1
|
config['history_length'] = 1
|
||||||
config['logger'] = gym.logger
|
config['logger'] = gym.logger
|
||||||
|
config['tag'] = ''
|
||||||
agent = AsyncAgent(**config)
|
agent = AsyncAgent(**config)
|
||||||
agent.run()
|
agent.run()
|
||||||
|
|
||||||
@@ -57,13 +59,14 @@ def a3c_cart_pole():
|
|||||||
config['worker_fn'] = AdvantageActorCritic
|
config['worker_fn'] = AdvantageActorCritic
|
||||||
config['discount'] = 0.99
|
config['discount'] = 0.99
|
||||||
config['target_network_update_freq'] = 200
|
config['target_network_update_freq'] = 200
|
||||||
config['step_limit'] = 0
|
config['step_limit'] = 200
|
||||||
config['n_workers'] = 16
|
config['n_workers'] = 16
|
||||||
config['update_interval'] = update_interval
|
config['update_interval'] = update_interval
|
||||||
config['history_length'] = 1
|
config['history_length'] = 1
|
||||||
config['test_interval'] = 4000
|
config['test_interval'] = 4000
|
||||||
config['test_repetitions'] = 50
|
config['test_repetitions'] = 50
|
||||||
config['logger'] = gym.logger
|
config['logger'] = gym.logger
|
||||||
|
config['tag'] = ''
|
||||||
agent = AsyncAgent(**config)
|
agent = AsyncAgent(**config)
|
||||||
agent.run()
|
agent.run()
|
||||||
|
|
||||||
@@ -87,8 +90,8 @@ def dqn_pixel_atari(name):
|
|||||||
config['test_repetitions'] = 1
|
config['test_repetitions'] = 1
|
||||||
# config['double_q'] = True
|
# config['double_q'] = True
|
||||||
config['double_q'] = False
|
config['double_q'] = False
|
||||||
|
config['tag'] = ''
|
||||||
agent = DQNAgent(**config)
|
agent = DQNAgent(**config)
|
||||||
agent.tag = 'dueling_'
|
|
||||||
agent.run()
|
agent.run()
|
||||||
|
|
||||||
def async_pixel_atari(name):
|
def async_pixel_atari(name):
|
||||||
@@ -115,8 +118,8 @@ def async_pixel_atari(name):
|
|||||||
config['test_repetitions'] = 1
|
config['test_repetitions'] = 1
|
||||||
config['history_length'] = history_length
|
config['history_length'] = history_length
|
||||||
config['logger'] = gym.logger
|
config['logger'] = gym.logger
|
||||||
|
config['tag'] = ''
|
||||||
agent = AsyncAgent(**config)
|
agent = AsyncAgent(**config)
|
||||||
agent.tag = 'Centered-target-network-'
|
|
||||||
agent.run()
|
agent.run()
|
||||||
|
|
||||||
def a3c_pixel_atari(name):
|
def a3c_pixel_atari(name):
|
||||||
@@ -139,29 +142,8 @@ def a3c_pixel_atari(name):
|
|||||||
config['test_repetitions'] = 1
|
config['test_repetitions'] = 1
|
||||||
config['history_length'] = history_length
|
config['history_length'] = history_length
|
||||||
config['logger'] = gym.logger
|
config['logger'] = gym.logger
|
||||||
agent = AsyncAgent(**config)
|
|
||||||
agent.tag = ''
|
|
||||||
agent.run()
|
|
||||||
|
|
||||||
def ddpg_montain_car():
|
|
||||||
config = dict()
|
|
||||||
config['task_fn'] = lambda: ContinuousMountainCar()
|
|
||||||
config['actor_network_fn'] = lambda: DDPGActorNet(2, 1)
|
|
||||||
config['critic_network_fn'] = lambda: DDPGCriticNet(2, 1)
|
|
||||||
config['actor_optimizer_fn'] = lambda params: torch.optim.Adam(params, lr=1e-4)
|
|
||||||
config['critic_optimizer_fn'] =\
|
|
||||||
lambda params: torch.optim.Adam(params, lr=1e-3, weight_decay=0.01)
|
|
||||||
config['replay_fn'] = lambda: HighDimActionReplay(memory_size=1000000, batch_size=64)
|
|
||||||
config['discount'] = 0.99
|
|
||||||
config['step_limit'] = 2500
|
|
||||||
config['tau'] = 0.001
|
|
||||||
config['exploration_steps'] = 100
|
|
||||||
config['random_process_fn'] = lambda: OrnsteinUhlenbeckProcess(theta=0.15, sigma=0.2)
|
|
||||||
config['test_interval'] = 50
|
|
||||||
config['test_repetitions'] = 10
|
|
||||||
config['tag'] = ''
|
config['tag'] = ''
|
||||||
config['logger'] = gym.logger
|
agent = AsyncAgent(**config)
|
||||||
agent = DDPGAgent(**config)
|
|
||||||
agent.run()
|
agent.run()
|
||||||
|
|
||||||
def ddpg_pendulum():
|
def ddpg_pendulum():
|
||||||
@@ -188,6 +170,32 @@ def ddpg_pendulum():
|
|||||||
agent = DDPGAgent(**config)
|
agent = DDPGAgent(**config)
|
||||||
agent.run()
|
agent.run()
|
||||||
|
|
||||||
|
def ddpg_bipedal_walker():
|
||||||
|
task_fn = lambda: BipedalWalker()
|
||||||
|
task = task_fn()
|
||||||
|
action_dim = task.env.action_space.shape[0]
|
||||||
|
state_dim = task.env.observation_space.shape[0]
|
||||||
|
config = dict()
|
||||||
|
config['task_fn'] = task_fn
|
||||||
|
config['actor_network_fn'] = lambda: DDPGActorNet(state_dim, action_dim, gpu=True)
|
||||||
|
config['critic_network_fn'] = lambda: DDPGCriticNet(state_dim, action_dim, gpu=True)
|
||||||
|
config['actor_optimizer_fn'] = lambda params: torch.optim.Adam(params, lr=1e-4)
|
||||||
|
config['critic_optimizer_fn'] =\
|
||||||
|
lambda params: torch.optim.Adam(params, lr=1e-3, weight_decay=0.01)
|
||||||
|
config['replay_fn'] = lambda: HighDimActionReplay(memory_size=1000000, batch_size=64)
|
||||||
|
config['discount'] = 0.99
|
||||||
|
config['step_limit'] = 1000
|
||||||
|
config['tau'] = 0.001
|
||||||
|
config['exploration_steps'] = 100
|
||||||
|
config['random_process_fn'] = \
|
||||||
|
lambda: OrnsteinUhlenbeckProcess(size=action_dim, theta=0.15, sigma=0.2)
|
||||||
|
config['test_interval'] = 10
|
||||||
|
config['test_repetitions'] = 10
|
||||||
|
config['tag'] = ''
|
||||||
|
config['logger'] = gym.logger
|
||||||
|
agent = DDPGAgent(**config)
|
||||||
|
agent.run()
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
gym.logger.setLevel(logging.DEBUG)
|
gym.logger.setLevel(logging.DEBUG)
|
||||||
# gym.logger.setLevel(logging.INFO)
|
# gym.logger.setLevel(logging.INFO)
|
||||||
@@ -204,5 +212,5 @@ if __name__ == '__main__':
|
|||||||
# async_pixel_atari('BreakoutNoFrameskip-v3')
|
# async_pixel_atari('BreakoutNoFrameskip-v3')
|
||||||
# a3c_pixel_atari('BreakoutNoFrameskip-v3')
|
# a3c_pixel_atari('BreakoutNoFrameskip-v3')
|
||||||
|
|
||||||
# ddpg_montain_car()
|
# ddpg_pendulum()
|
||||||
ddpg_pendulum()
|
ddpg_bipedal_walker()
|
||||||
|
|||||||
@@ -27,8 +27,6 @@ class BasicTask:
|
|||||||
next_state = self.normalize_state(next_state)
|
next_state = self.normalize_state(next_state)
|
||||||
return next_state, np.sign(reward), done, info
|
return next_state, np.sign(reward), done, info
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class MountainCar(BasicTask):
|
class MountainCar(BasicTask):
|
||||||
name = 'MountainCar-v0'
|
name = 'MountainCar-v0'
|
||||||
success_threshold = -110
|
success_threshold = -110
|
||||||
@@ -45,6 +43,7 @@ class CartPole(BasicTask):
|
|||||||
def __init__(self):
|
def __init__(self):
|
||||||
BasicTask.__init__(self)
|
BasicTask.__init__(self)
|
||||||
self.env = gym.make(self.name)
|
self.env = gym.make(self.name)
|
||||||
|
self.env._max_episode_steps = sys.maxsize
|
||||||
|
|
||||||
class LunarLander(BasicTask):
|
class LunarLander(BasicTask):
|
||||||
name = 'LunarLander-v2'
|
name = 'LunarLander-v2'
|
||||||
@@ -74,17 +73,6 @@ class PixelAtari(BasicTask):
|
|||||||
def normalize_state(self, state):
|
def normalize_state(self, state):
|
||||||
return np.asarray(state, dtype=np.float32) / 255.0
|
return np.asarray(state, dtype=np.float32) / 255.0
|
||||||
|
|
||||||
|
|
||||||
class ContinuousMountainCar(BasicTask):
|
|
||||||
name = 'MountainCarContinuous-v0'
|
|
||||||
success_threshold = 1000
|
|
||||||
|
|
||||||
def __init__(self):
|
|
||||||
BasicTask.__init__(self)
|
|
||||||
self.env = gym.make(self.name)
|
|
||||||
self.env._max_episode_steps = sys.maxsize
|
|
||||||
|
|
||||||
|
|
||||||
class Pendulum(BasicTask):
|
class Pendulum(BasicTask):
|
||||||
name = 'Pendulum-v0'
|
name = 'Pendulum-v0'
|
||||||
success_threshold = 200
|
success_threshold = 200
|
||||||
@@ -95,7 +83,20 @@ class Pendulum(BasicTask):
|
|||||||
self.env._max_episode_steps = sys.maxsize
|
self.env._max_episode_steps = sys.maxsize
|
||||||
|
|
||||||
def step(self, action):
|
def step(self, action):
|
||||||
|
action = 2 * np.clip(action, -1, 1)
|
||||||
|
next_state, reward, done, info = self.env.step(action)
|
||||||
|
return next_state, reward, done, info
|
||||||
|
|
||||||
|
class BipedalWalker(BasicTask):
|
||||||
|
name = 'BipedalWalker-v2'
|
||||||
|
success_threshold = 2000
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
BasicTask.__init__(self)
|
||||||
|
self.env = gym.make(self.name)
|
||||||
|
self.env._max_episode_steps = sys.maxsize
|
||||||
|
|
||||||
|
def step(self, action):
|
||||||
|
action = np.clip(action, -1, 1)
|
||||||
next_state, reward, done, info = self.env.step(action)
|
next_state, reward, done, info = self.env.step(action)
|
||||||
if self.normalized_state:
|
|
||||||
next_state = self.normalize_state(next_state)
|
|
||||||
return next_state, reward, done, info
|
return next_state, reward, done, info
|
||||||
|
|||||||
Reference in New Issue
Block a user