mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-11 11:53:01 +08:00
Update DDPG params and evaluation protocol
This commit is contained in:
@@ -46,7 +46,7 @@ class BaseAgent:
|
||||
if done:
|
||||
break
|
||||
total_rewards += reward
|
||||
return
|
||||
self.config.logger.info('evaluation episode return: %f' % (total_rewards))
|
||||
|
||||
def evaluation_episodes(self):
|
||||
interval = self.config.evaluation_episodes_interval
|
||||
|
||||
@@ -51,7 +51,6 @@ class DDPGAgent(BaseAgent):
|
||||
if not deterministic:
|
||||
action += self.random_process.sample()
|
||||
next_state, reward, done, info = self.task.step(action)
|
||||
# torchvision.utils.save_image(torch.tensor(np.asarray(next_state)).unsqueeze(1), 'data/image/%s.png' % get_time_str())
|
||||
next_state = self.config.state_normalizer(next_state)
|
||||
total_reward += reward
|
||||
reward = self.config.reward_normalizer(reward)
|
||||
|
||||
@@ -5,9 +5,27 @@ class RandomProcess(object):
|
||||
pass
|
||||
|
||||
class GaussianProcess(RandomProcess):
|
||||
def __init__(self, size, std_schedule):
|
||||
def __init__(self, size, std):
|
||||
self.size = size
|
||||
self.std_schedule = std_schedule
|
||||
self.std = std
|
||||
|
||||
def sample(self):
|
||||
return np.random.randn(self.size) * self.std_schedule()
|
||||
return np.random.randn(*self.size) * self.std()
|
||||
|
||||
class OrnsteinUhlenbeckProcess(RandomProcess):
|
||||
def __init__(self, size, std, theta=.15, dt=1e-2, x0=None):
|
||||
self.theta = theta
|
||||
self.mu = 0
|
||||
self.std = std
|
||||
self.dt = dt
|
||||
self.x0 = x0
|
||||
self.size = size
|
||||
self.reset_states()
|
||||
|
||||
def sample(self):
|
||||
x = self.x_prev + self.theta * (self.mu - self.x_prev) * self.dt + self.std() * np.sqrt(self.dt) * np.random.randn(*self.size)
|
||||
self.x_prev = x
|
||||
return x
|
||||
|
||||
def reset_states(self):
|
||||
self.x_prev = self.x0 if self.x0 is not None else np.zeros(self.size)
|
||||
|
||||
+10
-7
@@ -335,9 +335,12 @@ def ddpg_low_dim_state():
|
||||
config = Config()
|
||||
log_dir = get_default_log_dir(ddpg_low_dim_state.__name__)
|
||||
# config.task_fn = lambda **kwargs: Pendulum(log_dir=log_dir)
|
||||
config.task_fn = lambda **kwargs: Bullet('AntBulletEnv-v0', **kwargs)
|
||||
# config.task_fn = lambda **kwargs: Roboschool('RoboschoolAnt-v1', **kwargs)
|
||||
# config.task_fn = lambda **kwargs: Bullet('AntBulletEnv-v0', **kwargs)
|
||||
config.task_fn = lambda **kwargs: Roboschool('RoboschoolHopper-v1', **kwargs)
|
||||
config.evaluation_env = config.task_fn(log_dir=log_dir)
|
||||
config.max_steps = int(1e6)
|
||||
config.evaluation_episodes_interval = int(1e4)
|
||||
config.evaluation_episodes = 20
|
||||
|
||||
config.network_fn = lambda state_dim, action_dim: DeterministicActorCriticNet(
|
||||
state_dim, action_dim,
|
||||
@@ -348,8 +351,8 @@ def ddpg_low_dim_state():
|
||||
|
||||
config.replay_fn = lambda: Replay(memory_size=1000000, batch_size=64)
|
||||
config.discount = 0.99
|
||||
config.state_normalizer = RunningStatsNormalizer()
|
||||
config.random_process_fn = lambda action_dim: GaussianProcess(action_dim, LinearSchedule(0.3, 0, 1e6))
|
||||
config.random_process_fn = lambda action_dim: OrnsteinUhlenbeckProcess(
|
||||
size=(action_dim, ), std=LinearSchedule(0.2))
|
||||
config.min_memory_size = 64
|
||||
config.target_network_mix = 1e-3
|
||||
config.logger = get_logger()
|
||||
@@ -374,8 +377,8 @@ def ddpg_pixel():
|
||||
config.discount = 0.99
|
||||
config.state_normalizer = ImageNormalizer()
|
||||
config.max_steps = 1e7
|
||||
config.random_process_fn = lambda action_dim: GaussianProcess(
|
||||
action_dim, LinearSchedule(0.3, 0, config.max_steps))
|
||||
config.random_process_fn = lambda action_dim: OrnsteinUhlenbeckProcess(
|
||||
size=(action_dim, ), std=LinearSchedule(0.2))
|
||||
config.min_memory_size = 64
|
||||
config.target_network_mix = 1e-3
|
||||
config.logger = get_logger(file_name=ddpg_pixel.__name__)
|
||||
@@ -435,7 +438,7 @@ if __name__ == '__main__':
|
||||
|
||||
# ddpg_low_dim_state()
|
||||
# ddpg_pixel()
|
||||
# ppo_continuous()
|
||||
ppo_continuous()
|
||||
|
||||
# action_conditional_video_prediction()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user