Expose original Atari game reward

This commit is contained in:
Shangtong Zhang
2017-12-23 20:52:17 -07:00
parent 69b7ea53cc
commit 964734755f
9 changed files with 11 additions and 8 deletions
+2 -1
View File
@@ -46,10 +46,11 @@ class DQNAgent:
self.history_buffer.pop(0)
self.history_buffer.append(next_state)
next_state = np.vstack(self.history_buffer)
total_reward += np.sum(reward * self.config.reward_weight)
reward = self.config.reward_shift_fn(reward)
if not deterministic:
self.replay.feed([state, action, reward, next_state, int(done)])
self.total_steps += 1
total_reward += np.sum(reward * self.config.reward_weight)
steps += 1
state = next_state
if done:
+1
View File
@@ -33,6 +33,7 @@ class AdvantageActorCritic:
steps += 1
total_reward += reward
reward = config.reward_shift_fn(reward)
if deterministic:
if terminal:
+1
View File
@@ -34,6 +34,7 @@ class NStepQLearning:
steps += 1
total_reward += reward
reward = config.reward_shift_fn(reward)
if deterministic:
if terminal:
+1
View File
@@ -34,6 +34,7 @@ class OneStepQLearning:
steps += 1
total_reward += reward
reward = config.reward_shift_fn(reward)
if deterministic:
if terminal:
+1
View File
@@ -37,6 +37,7 @@ class OneStepSarsa:
steps += 1
total_reward += reward
reward = config.reward_shift_fn(reward)
if deterministic:
if terminal:
-5
View File
@@ -176,11 +176,6 @@ class ProcessFrame(gym.Wrapper):
def _reset(self):
return self.process_fn(self.env.reset())
class ClippedRewardsWrapper(gym.Wrapper):
def _step(self, action):
obs, reward, done, info = self.env.step(action)
return obs, np.sign(reward), done, info
class NormalizeFrame(gym.Wrapper):
def __init__(self, env=None):
super(NormalizeFrame, self).__init__(env)
+1 -2
View File
@@ -70,8 +70,7 @@ class PixelAtari(BasicTask):
env = MaxAndSkipEnv(env, skip=frame_skip)
if 'FIRE' in env.unwrapped.get_action_meanings():
env = FireResetEnv(env)
env = ProcessFrame(env, frame_size)
self.env = ClippedRewardsWrapper(env)
self.env = ProcessFrame(env, frame_size)
self.action_dim = self.env.action_space.n
def normalize_state(self, state):
+3
View File
@@ -79,6 +79,7 @@ def dqn_pixel_atari(name):
# config.network_fn = lambda optimizer_fn: DuelingNatureConvNet(config.history_length, n_actions, optimizer_fn)
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1)
config.replay_fn = lambda: Replay(memory_size=1000000, batch_size=32, dtype=np.uint8)
config.reward_shift_fn = lambda r: np.sign(r)
config.discount = 0.99
config.target_network_update_freq = 10000
config.max_episode_length = 0
@@ -104,6 +105,7 @@ def async_pixel_atari(name):
# config.worker = OneStepSarsa
# config.worker = NStepQLearning
config.worker = OneStepQLearning
config.reward_shift_fn = lambda r: np.sign(r)
config.discount = 0.99
config.target_network_update_freq = 10000
config.max_episode_length = 10000
@@ -123,6 +125,7 @@ def a3c_pixel_atari(name):
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.reward_shift_fn = lambda r: np.sign(r)
config.policy_fn = SamplePolicy
config.worker = AdvantageActorCritic
config.discount = 0.99
+1
View File
@@ -37,6 +37,7 @@ class Config:
self.noise_decay_interval = 0
self.target_network_mix = 0.001
self.action_shift_fn = lambda a: a
self.reward_shift_fn = lambda r: r
self.reward_weight = 1
self.hybrid_reward = False
self.target_type = self.q_target