Minor refactor

This commit is contained in:
Shangtong Zhang
2017-07-20 20:47:51 -06:00
parent 7b7d37694e
commit a145a01cb6
5 changed files with 66 additions and 51 deletions
+6 -4
View File
@@ -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:
+4 -1
View File
@@ -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
View File
@@ -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()
+38 -30
View File
@@ -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()
+16 -15
View File
@@ -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