diff --git a/agent/A2C_agent.py b/agent/A2C_agent.py index b87c9ed..33a719b 100644 --- a/agent/A2C_agent.py +++ b/agent/A2C_agent.py @@ -52,6 +52,7 @@ class A2CAgent: rollout = [] states = self.states for i in range(config.rollout_length): + states = self.task.normalize_state(states) prob, log_prob, value = self.network.predict(states) actions = [self.policy.sample(p) for p in prob.data.cpu().numpy()] actions = config.action_shift_fn(actions) @@ -68,7 +69,7 @@ class A2CAgent: states = next_states self.states = states - _, _, pending_value = self.network.predict(states) + _, _, pending_value = self.network.predict(self.task.normalize_state(states)) rollout.append([None, None, pending_value, None, None, None]) processed_rollout = [None] * (len(rollout) - 1) diff --git a/agent/NStepDQN_agent.py b/agent/NStepDQN_agent.py index f2bf05f..ae4588f 100644 --- a/agent/NStepDQN_agent.py +++ b/agent/NStepDQN_agent.py @@ -40,7 +40,7 @@ class NStepDQNAgent: rollout = [] states = self.states for i in range(config.rollout_length): - q = self.learning_network.predict(states) + q = self.learning_network.predict(self.task.normalize_state(states)) actions = [self.policy.sample(v) for v in q.data.cpu().numpy()] actions = config.action_shift_fn(actions) next_states, rewards, terminals, _ = self.task.step(actions) @@ -63,7 +63,7 @@ class NStepDQNAgent: self.states = states processed_rollout = [None] * (len(rollout)) - returns = self.target_network.predict(states).data + returns = self.target_network.predict(self.task.normalize_state(states)).data returns, _ = torch.max(returns, dim=1, keepdim=True) for i in reversed(range(len(rollout))): q, actions, rewards, terminals = rollout[i] diff --git a/agent/__init__.py b/agent/__init__.py index 7d48e63..41f28d5 100644 --- a/agent/__init__.py +++ b/agent/__init__.py @@ -4,4 +4,4 @@ from .DDPG_agent import * from .A2C_agent import * from .CategoricalDQN_agent import * from .NStepDQN_agent import * -from .QuantileRegressionDQN_agent import * \ No newline at end of file +from .QuantileRegressionDQN_agent import * diff --git a/component/__init__.py b/component/__init__.py index b51b19c..d9c2df2 100644 --- a/component/__init__.py +++ b/component/__init__.py @@ -1,4 +1,4 @@ -from .atari_wrapper import * +from .atari_wrapper import * from .policy import * from .replay import * from .task import * diff --git a/component/atari_wrapper.py b/component/atari_wrapper.py index 667964d..f626b88 100644 --- a/component/atari_wrapper.py +++ b/component/atari_wrapper.py @@ -1,54 +1,70 @@ +# based on https://github.com/openai/baselines/blob/master/baselines/common/atari_wrappers.py + import numpy as np from collections import deque import gym from gym import spaces -from skimage import color, transform +from gym.spaces import Box +import cv2 +cv2.ocl.setUseOpenCL(False) class NoopResetEnv(gym.Wrapper): - def __init__(self, env=None, noop_max=30): + def __init__(self, env, noop_max=30): """Sample initial states by taking random number of no-ops on reset. No-op is assumed to be action 0. """ - super(NoopResetEnv, self).__init__(env) + gym.Wrapper.__init__(self, env) self.noop_max = noop_max + self.override_num_noops = None + self.noop_action = 0 assert env.unwrapped.get_action_meanings()[0] == 'NOOP' - def reset(self): + def reset(self, **kwargs): """ Do no-op action for a number of steps in [1, noop_max].""" - self.env.reset() - noops = np.random.randint(1, self.noop_max + 1) + self.env.reset(**kwargs) + if self.override_num_noops is not None: + noops = self.override_num_noops + else: + noops = self.unwrapped.np_random.randint(1, self.noop_max + 1) #pylint: disable=E1101 + assert noops > 0 + obs = None for _ in range(noops): - obs, _, _, _ = self.env.step(0) + obs, _, done, _ = self.env.step(self.noop_action) + if done: + obs = self.env.reset(**kwargs) return obs - def step(self, action): - return self.env.step(action) + def step(self, ac): + return self.env.step(ac) class FireResetEnv(gym.Wrapper): - def __init__(self, env=None): + def __init__(self, env): """Take action on reset for environments that are fixed until firing.""" - super(FireResetEnv, self).__init__(env) + gym.Wrapper.__init__(self, env) assert env.unwrapped.get_action_meanings()[1] == 'FIRE' assert len(env.unwrapped.get_action_meanings()) >= 3 - def reset(self): - self.env.reset() - obs, _, _, _ = self.env.step(1) - obs, _, _, _ = self.env.step(2) + def reset(self, **kwargs): + self.env.reset(**kwargs) + obs, _, done, _ = self.env.step(1) + if done: + self.env.reset(**kwargs) + obs, _, done, _ = self.env.step(2) + if done: + self.env.reset(**kwargs) return obs - def step(self, action): - return self.env.step(action) + def step(self, ac): + return self.env.step(ac) class EpisodicLifeEnv(gym.Wrapper): - def __init__(self, env=None): + def __init__(self, env): """Make end-of-life == end-of-episode, but only reset on true game over. Done by DeepMind for the DQN and co. since it helps value estimation. """ - super(EpisodicLifeEnv, self).__init__(env) + gym.Wrapper.__init__(self, env) self.lives = 0 self.was_real_done = True - self.was_realreset = False def step(self, action): obs, reward, done, info = self.env.step(action) @@ -57,56 +73,53 @@ class EpisodicLifeEnv(gym.Wrapper): # then update lives to handle bonus lives lives = self.env.unwrapped.ale.lives() if lives < self.lives and lives > 0: - # for Qbert somtimes we stay in lives == 0 condtion for a few frames + # for Qbert sometimes we stay in lives == 0 condtion for a few frames # so its important to keep lives > 0, so that we only reset once # the environment advertises done. done = True self.lives = lives return obs, reward, done, info - def reset(self): + def reset(self, **kwargs): """Reset only when lives are exhausted. This way all states are still reachable even though lives are episodic, and the learner need not know about any of this behind-the-scenes. """ if self.was_real_done: - obs = self.env.reset() - self.was_realreset = True + obs = self.env.reset(**kwargs) else: # no-op step to advance from terminal/lost life state obs, _, _, _ = self.env.step(0) - self.was_realreset = False self.lives = self.env.unwrapped.ale.lives() return obs class MaxAndSkipEnv(gym.Wrapper): - def __init__(self, env=None, skip=4): + def __init__(self, env, skip=4): """Return only every `skip`-th frame""" - super(MaxAndSkipEnv, self).__init__(env) + gym.Wrapper.__init__(self, env) # most recent raw observations (for max pooling across time steps) - self._obs_buffer = deque(maxlen=2) + self._obs_buffer = np.zeros((2,)+env.observation_space.shape, dtype=np.uint8) self._skip = skip def step(self, action): + """Repeat action, sum reward, and max over last observations.""" total_reward = 0.0 done = None - for _ in range(self._skip): + for i in range(self._skip): obs, reward, done, info = self.env.step(action) - self._obs_buffer.append(obs) + if i == self._skip - 2: self._obs_buffer[0] = obs + if i == self._skip - 1: self._obs_buffer[1] = obs total_reward += reward if done: break - - max_frame = np.max(np.stack(self._obs_buffer), axis=0) + # Note that the observation on the done=True frame + # doesn't matter + max_frame = self._obs_buffer.max(axis=0) return max_frame, total_reward, done, info - def reset(self): - """Clear past frame buffer and init. to first obs. from inner env.""" - self._obs_buffer.clear() - obs = self.env.reset() - self._obs_buffer.append(obs) - return obs + def reset(self, **kwargs): + return self.env.reset(**kwargs) class SkipEnv(gym.Wrapper): def __init__(self, env=None, skip=4): @@ -129,6 +142,92 @@ class SkipEnv(gym.Wrapper): obs = self.env.reset() return obs +class ClipRewardEnv(gym.RewardWrapper): + def __init__(self, env): + gym.RewardWrapper.__init__(self, env) + + def reward(self, reward): + """Bin reward to {+1, 0, -1} by its sign.""" + return np.sign(reward) + +class WarpFrame(gym.ObservationWrapper): + def __init__(self, env): + """Warp frames to 84x84 as done in the Nature paper and later work.""" + gym.ObservationWrapper.__init__(self, env) + self.width = 84 + self.height = 84 + self.observation_space = spaces.Box(low=0, high=255, + shape=(self.height, self.width, 1), dtype=np.uint8) + + def observation(self, frame): + frame = cv2.cvtColor(frame, cv2.COLOR_RGB2GRAY) + frame = cv2.resize(frame, (self.width, self.height), interpolation=cv2.INTER_AREA) + return frame[:, :, None] + +class LazyFrames(object): + def __init__(self, frames): + """This object ensures that common frames between the observations are only stored once. + It exists purely to optimize memory usage which can be huge for DQN's 1M frames replay + buffers. + + This object should only be converted to numpy array before being passed to the model. + + You'd not believe how complex the previous solution was.""" + self._frames = frames + self._out = None + + def _force(self): + if self._out is None: + self._out = np.concatenate(self._frames, axis=2) + self._frames = None + return self._out + + def __array__(self, dtype=None): + out = self._force() + if dtype is not None: + out = out.astype(dtype) + return out + + def __len__(self): + return len(self._force()) + + def __getitem__(self, i): + return self._force()[i] + +class StackFrame(gym.Wrapper): + def __init__(self, env=None, history_length=1): + super(StackFrame, self).__init__(env) + self.history_length = history_length + self.buffer = None + + def reset(self): + state = self.env.reset() + self.buffer = [state] * self.history_length + # return LazyFrames(self.buffer) + return np.asarray(np.vstack(self.buffer)) + + def step(self, action): + state, reward, done, info = self.env.step(action) + self.buffer.pop(0) + self.buffer.append(state) + # return LazyFrames(self.buffer), reward, done, info + return np.asarray(np.vstack(self.buffer)), reward, done, info + +class WrapPyTorch(gym.ObservationWrapper): + # from https://github.com/ikostrikov/pytorch-a2c-ppo-acktr/blob/master/envs.py + def __init__(self, env=None): + super(WrapPyTorch, self).__init__(env) + obs_shape = self.observation_space.shape + self.observation_space = Box( + self.observation_space.low[0,0,0], + self.observation_space.high[0,0,0], + [obs_shape[2], obs_shape[1], obs_shape[0]], + dtype=np.uint8 + ) + + def observation(self, observation): + return observation.transpose(2, 0, 1) + class DatasetEnv(gym.Wrapper): def __init__(self, env=None): super(DatasetEnv, self).__init__(env) @@ -153,52 +252,24 @@ class DatasetEnv(gym.Wrapper): self.saved_obs.append(obs) return obs -class ProcessFrame(gym.Wrapper): - def __init__(self, env=None, frame_size=84): - super(ProcessFrame, self).__init__(env) - self.frame_size = frame_size - self.observation_space = spaces.Box(low=0, high=255, shape=(1, frame_size, frame_size), dtype=np.uint8) +def make_atari(env_id, frame_skip=4): + env = gym.make(env_id) + assert 'NoFrameskip' in env.spec.id + env = NoopResetEnv(env, noop_max=30) + env = MaxAndSkipEnv(env, skip=4) + return env - def process(self, obs): - obs = color.rgb2gray(obs) - obs = transform.resize(obs, (self.frame_size, self.frame_size), mode='constant') - obs = (255 * obs).astype(np.uint8).reshape((1, ) + obs.shape) - return obs - - def step(self, action): - obs, reward, done, info = self.env.step(action) - return self.process(obs), reward, done, info - - def reset(self): - return self.process(self.env.reset()) - -class NormalizeFrame(gym.Wrapper): - def __init__(self, env=None): - super(NormalizeFrame, self).__init__(env) - - def _normalize(self, obs): - return np.asarray(obs, dtype=np.float32) / 255.0 - - def step(self, action): - obs, reward, done, info = self.env.step(action) - return self._normalize(obs), reward, done, info - - def reset(self): - return self._normalize(self.env.reset()) - -class StackFrame(gym.Wrapper): - def __init__(self, env=None, history_length=1): - super(StackFrame, self).__init__(env) - self.history_length = history_length - self.buffer = None - - def reset(self): - state = self.env.reset() - self.buffer = [state] * self.history_length - return np.asarray(np.vstack(self.buffer)) - - def step(self, action): - state, reward, done, info = self.env.step(action) - self.buffer.pop(0) - self.buffer.append(state) - return np.asarray(np.vstack(self.buffer)), reward, done, info +def wrap_deepmind(env, episode_life=True, clip_rewards=True, history_length=0): + """Configure environment for DeepMind-style Atari. + """ + if episode_life: + env = EpisodicLifeEnv(env) + if 'FIRE' in env.unwrapped.get_action_meanings(): + env = FireResetEnv(env) + env = WarpFrame(env) + if clip_rewards: + env = ClipRewardEnv(env) + env = WrapPyTorch(env) + if history_length: + env = StackFrame(env, history_length) + return env diff --git a/component/replay.py b/component/replay.py index 6b787b5..18225b7 100644 --- a/component/replay.py +++ b/component/replay.py @@ -225,12 +225,17 @@ class GeneralReplay: def feed(self, experiences): for experience in zip(*experiences): - self.buffer.append(experience) - if len(self.buffer) > self.memory_size: - del self.buffer[0] + self.feed_single(experience) - def sample(self): - sampled = zip(*random.sample(self.buffer, self.batch_size)) + def feed_single(self, experience): + self.buffer.append(experience) + if len(self.buffer) > self.memory_size: + del self.buffer[0] + + def sample(self, batch_size=None): + if batch_size is None: + batch_size = self.batch_size + sampled = zip(*random.sample(self.buffer, batch_size)) return sampled def clear(self): @@ -238,3 +243,42 @@ class GeneralReplay: def full(self): return len(self.buffer) == self.memory_size + + def size(self): + return len(self.buffer) + + def empty(self): + return not len(self.buffer) + +class SkewedReplay: + def __init__(self, memory_size, batch_size): + memory_size = memory_size / 2 + self.non_zero_reward = GeneralReplay(memory_size, batch_size / 2) + self.zero_reward = GeneralReplay(memory_size, batch_size / 2) + self.batch_size = batch_size + + def feed(self, experiences): + experiences = zip(*experiences) + for exp in experiences: + if np.abs(exp[2]) < 1e-5: + self.zero_reward.feed_single(exp) + else: + self.non_zero_reward.feed_single(exp) + + def sample(self): + if self.zero_reward.empty(): + batch = self.non_zero_reward.sample(self.batch_size) + elif self.non_zero_reward.empty(): + batch = self.zero_reward.sample(self.batch_size) + else: + non_zero_batch_size = min(self.non_zero_reward.size(), self.batch_size / 2) + zero_batch_size = min(self.zero_reward.size(), self.batch_size / 2) + batch1 = self.zero_reward.sample(zero_batch_size) + batch2 = self.non_zero_reward.sample(non_zero_batch_size) + batch = list(map(lambda seq: np.concatenate([np.asarray(x) for x in seq], axis=0), zip(batch1, batch2))) + batch = list(map(lambda x: np.asarray(x), batch)) + return batch + + + + diff --git a/component/task.py b/component/task.py index 6e747a9..f9c586f 100644 --- a/component/task.py +++ b/component/task.py @@ -11,6 +11,8 @@ import multiprocessing as mp import sys from .bench import Monitor from utils import * +import datetime +import uuid class BasicTask: def __init__(self, max_steps=sys.maxsize): @@ -34,9 +36,6 @@ class BasicTask: def random_action(self): return self.env.action_space.sample() - def set_monitor(self, filename): - self.env = Monitor(self.env, filename) - class ClassicalControl(BasicTask): def __init__(self, name='CartPole-v0', max_steps=200): BasicTask.__init__(self, max_steps) @@ -57,23 +56,22 @@ class LunarLander(BasicTask): 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=10000, history_length=1): + def __init__(self, name, seed=0, log_file=None, max_steps=sys.maxsize, + frame_skip=4, history_length=4): BasicTask.__init__(self, max_steps) - self.normalized_state = normalized_state - self.name = name - env = gym.make(name) - assert 'NoFrameskip' in env.spec.id - env = EpisodicLifeEnv(env) - env = NoopResetEnv(env, noop_max=no_op) - env = MaxAndSkipEnv(env, skip=frame_skip) - if 'FIRE' in env.unwrapped.get_action_meanings(): - env = FireResetEnv(env) - env = ProcessFrame(env, frame_size) - if normalized_state: - env = NormalizeFrame(env) - self.env = StackFrame(env, history_length) + env = make_atari(name, frame_skip) + env.seed(seed) + if log_file is None: + log_dir = '%s-%s' % ( + name, + datetime.datetime.now().strftime("%y%m%d-%-H%M%S")) + mkdir('./log/%s' % log_dir) + log_file = './log/%s/%s' % (log_dir, uuid.uuid1()) + env = Monitor(env, log_file) + env = wrap_deepmind(env, history_length=history_length) + self.env = env self.action_dim = self.env.action_space.n + self.name = name def normalize_state(self, state): return np.asarray(state) / 255.0 @@ -143,12 +141,12 @@ class Roboschool(BasicTask): def step(self, action): return BasicTask.step(self, np.clip(action, -1, 1)) -def sub_task(parent_pipe, pipe, task_fn, filename=None): +def sub_task(parent_pipe, pipe, task_fn, rank, log_dir): + np.random.seed() + seed = np.random.randint(0, sys.maxsize) parent_pipe.close() - task = task_fn() - if filename is not None: - task.set_monitor(filename) - task.env.seed(np.random.randint(0, sys.maxsize)) + task = task_fn(log_file=os.path.join(log_dir, str(rank))) + task.env.seed(seed) while True: op, data = pipe.recv() if op == 'step': @@ -166,13 +164,11 @@ class ParallelizedTask: self.task_fn = task_fn self.task = task_fn() self.name = self.task.name - # date = datetime.datetime.now().strftime("%I:%M%p-on-%B-%d-%Y") - mkdir('./log/%s-%s' % (self.name, tag)) - filenames = ['./log/%s-%s/worker-%d' % (self.name, tag, i) - for i in range(num_workers)] + log_dir = './log/%s-%s' % (self.name, tag) + mkdir(log_dir) self.pipes, worker_pipes = zip(*[mp.Pipe() for _ in range(num_workers)]) - args = [(p, wp, task_fn, filename) - for p, wp, filename in zip(self.pipes, worker_pipes, filenames)] + args = [(p, wp, task_fn, rank, log_dir) + for rank, (p, wp) in enumerate(zip(self.pipes, worker_pipes))] self.workers = [mp.Process(target=sub_task, args=arg) for arg in args] for p in self.workers: p.start() for p in worker_pipes: p.close() @@ -200,3 +196,6 @@ class ParallelizedTask: for pipe in self.pipes: pipe.send(('exit', None)) for p in self.workers: p.join() + + def normalize_state(self, state): + return self.task.normalize_state(state) diff --git a/main.py b/main.py index 657de7d..71635c4 100644 --- a/main.py +++ b/main.py @@ -29,54 +29,11 @@ def dqn_cart_pole(): # config.double_q = False run_episodes(DQNAgent(config)) -def async_cart_pole(): - config = Config() - 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) - # config.worker = OneStepQLearning - config.worker = NStepQLearning - # config.worker = OneStepSarsa - config.discount = 0.99 - config.target_network_update_freq = 200 - config.num_workers = 16 - config.update_interval = 6 - config.test_interval = 1 - config.test_repetitions = 50 - config.logger = Logger('./log', logger) - agent = AsyncAgent(config) - agent.run() - -def a3c_cart_pole(): - config = Config() - 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(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 = 7 - config.update_interval = 6 - config.test_interval = 1 - config.test_repetitions = 30 - config.logger = Logger('./log', logger) - config.gae_tau = 1.0 - config.entropy_weight = 0.01 - agent = AsyncAgent(config) - agent.run() - def a2c_cart_pole(): config = Config() name = 'CartPole-v0' # name = 'MountainCar-v0' - task_fn = lambda: ClassicalControl(name, max_steps=200) - # task_fn = lambda: LunarLander() + task_fn = lambda **kwargs: ClassicalControl(name, max_steps=200) task = task_fn() config.num_workers = 5 config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers) @@ -84,8 +41,6 @@ def a2c_cart_pole(): config.network_fn = lambda: ActorCriticFCNet(task.state_dim, task.action_dim) config.policy_fn = SamplePolicy config.discount = 0.99 - config.test_interval = 200 - config.test_repetitions = 10 config.logger = Logger('./log', logger) config.gae_tau = 1.0 config.entropy_weight = 0.01 @@ -95,14 +50,13 @@ def a2c_cart_pole(): def dqn_pixel_atari(name): config = Config() config.history_length = 4 - config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False, - history_length=config.history_length) + config.task_fn = lambda: PixelAtari(name, frame_skip=4, history_length=config.history_length) action_dim = config.task_fn().action_dim config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01) config.network_fn = lambda: NatureConvNet(config.history_length, action_dim, gpu=0) # config.network_fn = lambda: DuelingNatureConvNet(config.history_length, action_dim) 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.replay_fn = lambda: Replay(memory_size=100000, 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 @@ -136,63 +90,14 @@ def dqn_ram_atari(name): # config.double_q = False run_episodes(DQNAgent(config)) -def async_pixel_atari(name): - config = Config() - config.history_length = 1 - config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42) - task = config.task_fn() - config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.0001) - config.network_fn = lambda: OpenAIConvNet( - config.history_length, task.env.action_space.n) - config.policy_fn = lambda: StochasticGreedyPolicy( - epsilons=[0.7, 0.7, 0.7], final_step=2000000, min_epsilons=[0.1, 0.01, 0.5], - probs=[0.4, 0.3, 0.3]) - # 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 - config.num_workers = 6 - config.update_interval = 20 - config.test_interval = 50000 - config.test_repetitions = 1 - config.logger = Logger('./log', logger) - agent = AsyncAgent(config) - agent.run() - -def a3c_pixel_atari(name): - config = Config() - config.history_length = 1 - config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42) - 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=False) - config.reward_shift_fn = lambda r: np.sign(r) - config.policy_fn = SamplePolicy - config.worker = AdvantageActorCritic - config.discount = 0.99 - config.num_workers = 6 - config.update_interval = 20 - config.test_interval = 50000 - config.test_repetitions = 1 - config.logger = Logger('./log', logger) - agent = AsyncAgent(config) - agent.run() - def a2c_pixel_atari(name): config = Config() config.history_length = 4 config.num_workers = 5 - task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=84, - history_length=config.history_length) - config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers) + task_fn = lambda **kwargs: PixelAtari(name, frame_skip=4, history_length=config.history_length) + config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers, tag=a2c_pixel_atari.__name__) task = config.task_fn() config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.0007) - # config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.0001) - # config.network_fn = lambda: OpenAIActorCriticConvNet( config.network_fn = lambda: NatureActorCriticConvNet( config.history_length, task.task.env.action_space.n, gpu=3) config.reward_shift_fn = lambda r: np.sign(r) @@ -208,98 +113,98 @@ def a2c_pixel_atari(name): config.logger = Logger('./log', logger, skip=True) run_iterations(A2CAgent(config)) -def a3c_continuous(): - config = Config() - config.task_fn = lambda: Pendulum() - # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2') - # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2') - # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2') - task = config.task_fn() - config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.0001) - config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) - config.network_fn = lambda: DisjointActorCriticNet( - # lambda: GaussianActorNet(task.state_dim, task.action_dim, unit_std=False, action_gate=F.tanh, action_scale=2.0), - lambda: GaussianActorNet(task.state_dim, task.action_dim, unit_std=True), - lambda: GaussianCriticNet(task.state_dim)) - config.policy_fn = lambda: GaussianPolicy() - config.worker = ContinuousAdvantageActorCritic - config.discount = 0.99 - config.num_workers = 8 - config.update_interval = 20 - config.test_interval = 1 - config.test_repetitions = 1 - config.entropy_weight = 0 - config.gradient_clip = 40 - config.logger = Logger('./log', logger) - agent = AsyncAgent(config) - agent.run() - -def p3o_continuous(): - config = Config() - config.task_fn = lambda: Pendulum() - # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2') - # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2') - # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2') - # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1') - # config.task_fn = lambda: Roboschool('RoboschoolAnt-v1') - task = config.task_fn() - config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim, - gpu=-1, unit_std=True) - config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim, gpu=-1) - config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn) - config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) - config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) - - config.policy_fn = lambda: GaussianPolicy() - config.replay_fn = lambda: GeneralReplay(memory_size=2048, batch_size=2048) - config.worker = ProximalPolicyOptimization - config.discount = 0.99 - config.gae_tau = 0.97 - config.num_workers = 6 - config.test_interval = 1 - config.test_repetitions = 1 - config.entropy_weight = 0 - config.gradient_clip = 20 - config.rollout_length = 10000 - config.optimize_epochs = 1 - config.ppo_ratio_clip = 0.2 - config.logger = Logger('./log', logger) - agent = AsyncAgent(config) - agent.run() - -def d3pg_continuous(): - config = Config() - config.task_fn = lambda: Pendulum() - # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2') - # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2') - # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2') - # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1') - # config.task_fn = lambda: Roboschool('RoboschoolReacher-v1') - task = config.task_fn() - config.actor_network_fn = lambda: DeterministicActorNet( - task.state_dim, task.action_dim, F.tanh, 2, non_linear=F.relu, batch_norm=False) - config.critic_network_fn = lambda: DeterministicCriticNet( - task.state_dim, task.action_dim, non_linear=F.relu, batch_norm=False) - config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn) - 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-4) - config.replay_fn = lambda: SharedReplay(memory_size=1000000, batch_size=64, - state_shape=(task.state_dim, ), action_shape=(task.action_dim, )) - config.discount = 0.99 - config.random_process_fn = \ - lambda: OrnsteinUhlenbeckProcess(size=task.action_dim, theta=0.15, sigma=0.2, - n_steps_annealing=100000) - config.worker = DeterministicPolicyGradient - config.num_workers = 6 - config.min_memory_size = 50 - config.target_network_mix = 0.001 - config.test_interval = 500 - config.test_repetitions = 1 - config.gradient_clip = 20 - config.logger = Logger('./log', logger) - agent = AsyncAgent(config) - agent.run() +# def a3c_continuous(): +# config = Config() +# config.task_fn = lambda: Pendulum() +# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2') +# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2') +# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2') +# task = config.task_fn() +# config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.0001) +# config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) +# config.network_fn = lambda: DisjointActorCriticNet( +# # lambda: GaussianActorNet(task.state_dim, task.action_dim, unit_std=False, action_gate=F.tanh, action_scale=2.0), +# lambda: GaussianActorNet(task.state_dim, task.action_dim, unit_std=True), +# lambda: GaussianCriticNet(task.state_dim)) +# config.policy_fn = lambda: GaussianPolicy() +# config.worker = ContinuousAdvantageActorCritic +# config.discount = 0.99 +# config.num_workers = 8 +# config.update_interval = 20 +# config.test_interval = 1 +# config.test_repetitions = 1 +# config.entropy_weight = 0 +# config.gradient_clip = 40 +# config.logger = Logger('./log', logger) +# agent = AsyncAgent(config) +# agent.run() +# +# def p3o_continuous(): +# config = Config() +# config.task_fn = lambda: Pendulum() +# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2') +# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2') +# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2') +# # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1') +# # config.task_fn = lambda: Roboschool('RoboschoolAnt-v1') +# task = config.task_fn() +# config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim, +# gpu=-1, unit_std=True) +# config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim, gpu=-1) +# config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn) +# config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) +# config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) +# +# config.policy_fn = lambda: GaussianPolicy() +# config.replay_fn = lambda: GeneralReplay(memory_size=2048, batch_size=2048) +# config.worker = ProximalPolicyOptimization +# config.discount = 0.99 +# config.gae_tau = 0.97 +# config.num_workers = 6 +# config.test_interval = 1 +# config.test_repetitions = 1 +# config.entropy_weight = 0 +# config.gradient_clip = 20 +# config.rollout_length = 10000 +# config.optimize_epochs = 1 +# config.ppo_ratio_clip = 0.2 +# config.logger = Logger('./log', logger) +# agent = AsyncAgent(config) +# agent.run() +# +# def d3pg_continuous(): +# config = Config() +# config.task_fn = lambda: Pendulum() +# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2') +# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2') +# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2') +# # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1') +# # config.task_fn = lambda: Roboschool('RoboschoolReacher-v1') +# task = config.task_fn() +# config.actor_network_fn = lambda: DeterministicActorNet( +# task.state_dim, task.action_dim, F.tanh, 2, non_linear=F.relu, batch_norm=False) +# config.critic_network_fn = lambda: DeterministicCriticNet( +# task.state_dim, task.action_dim, non_linear=F.relu, batch_norm=False) +# config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn) +# 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-4) +# config.replay_fn = lambda: SharedReplay(memory_size=1000000, batch_size=64, +# state_shape=(task.state_dim, ), action_shape=(task.action_dim, )) +# config.discount = 0.99 +# config.random_process_fn = \ +# lambda: OrnsteinUhlenbeckProcess(size=task.action_dim, theta=0.15, sigma=0.2, +# n_steps_annealing=100000) +# config.worker = DeterministicPolicyGradient +# config.num_workers = 6 +# config.min_memory_size = 50 +# config.target_network_mix = 0.001 +# config.test_interval = 500 +# config.test_repetitions = 1 +# config.gradient_clip = 20 +# config.logger = Logger('./log', logger) +# agent = AsyncAgent(config) +# agent.run() def ddpg_continuous(): config = Config() @@ -359,8 +264,7 @@ def categorical_dqn_cart_pole(): def categorical_dqn_pixel_atari(name): config = Config() config.history_length = 4 - config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False, - history_length=config.history_length) + config.task_fn = lambda: PixelAtari(name, frame_skip=4, history_length=config.history_length) action_dim = config.task_fn().action_dim config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.00025, eps=0.01 / 32) config.network_fn = lambda: CategoricalConvNet(config.history_length, action_dim, config.categorical_n_atoms, gpu=0) @@ -381,7 +285,7 @@ def categorical_dqn_pixel_atari(name): def n_step_dqn_cart_pole(): config = Config() - task_fn = lambda: ClassicalControl('CartPole-v0', max_steps=200) + task_fn = lambda **kwargs: ClassicalControl('CartPole-v0', max_steps=200) task = task_fn() config.num_workers = 5 config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers) @@ -397,11 +301,10 @@ def n_step_dqn_cart_pole(): def n_step_dqn_pixel_atari(name): config = Config() config.history_length = 4 - task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=True, - history_length=config.history_length) + task_fn = lambda **kwargs: PixelAtari(name, frame_skip=4, history_length=config.history_length) task = task_fn() config.num_workers = 8 - config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers) + config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers, tag=n_step_dqn_pixel_atari.__name__) config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01) config.network_fn = lambda: NatureConvNet(config.history_length, task.action_dim, gpu=0) config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1) @@ -433,8 +336,7 @@ def quantile_regression_dqn_cart_pole(): def quantile_regression_dqn_pixel_atari(name): config = Config() config.history_length = 4 - config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False, - history_length=config.history_length) + config.task_fn = lambda: PixelAtari(name, frame_skip=4, history_length=config.history_length) action_dim = config.task_fn().action_dim config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.00005, eps=0.01 / 32) config.network_fn = lambda: QuantileConvNet(config.history_length, action_dim, config.num_quantiles, gpu=0) @@ -460,28 +362,19 @@ if __name__ == '__main__': logger.setLevel(logging.INFO) # dqn_cart_pole() + # a2c_cart_pole() # categorical_dqn_cart_pole() # quantile_regression_dqn_cart_pole() - # async_cart_pole() - # a3c_cart_pole() - a2c_cart_pole() - # a3c_continuous() - # p3o_continuous() - # d3pg_continuous() # ddpg_continuous() # n_step_dqn_cart_pole() # dqn_pixel_atari('PongNoFrameskip-v4') - # categorical_dqn_pixel_atari('PongNoFrameskip-v4') - # quantile_regression_dqn_pixel_atari('PongNoFrameskip-v4') - # n_step_dqn_pixel_atari('PongNoFrameskip-v4') - # async_pixel_atari('PongNoFrameskip-v4') - # a3c_pixel_atari('PongNoFrameskip-v4') # a2c_pixel_atari('PongNoFrameskip-v4') + # categorical_dqn_pixel_atari('PongNoFrameskip-v4') + quantile_regression_dqn_pixel_atari('PongNoFrameskip-v4') + # n_step_dqn_pixel_atari('PongNoFrameskip-v4') # dqn_pixel_atari('BreakoutNoFrameskip-v4') - # async_pixel_atari('BreakoutNoFrameskip-v4') - # a3c_pixel_atari('BreakoutNoFrameskip-v4') # dqn_ram_atari('Pong-ramNoFrameskip-v4') diff --git a/network/base_network.py b/network/base_network.py index 075f86e..0c906f6 100644 --- a/network/base_network.py +++ b/network/base_network.py @@ -106,3 +106,78 @@ class QuantileNet(BasicNet): def predict(self, x, to_numpy=False): quantiles = self.forward(x) return quantiles.view((-1, self.n_actions, self.n_quantiles)) + +class GammaNet(BasicNet): + def predict(self, features, aux_features): + attention = self.compute_attention(features) + + aux_features = torch.stack(aux_features) + aux_features = aux_features * attention.t().unsqueeze(-1) + aux_features = aux_features.transpose(0, 1).contiguous().sum(1) + + phi = features + aux_features + pre_prob = self.fc_actor(phi) + prob = F.softmax(pre_prob, dim=1) + log_prob = F.log_softmax(pre_prob, dim=1) + value = self.fc_critic(phi) + return prob, log_prob, value + + def compute_attention(self, phi): + attention = self.fc_attention(phi) + attention = F.sigmoid(attention) + return attention + + def q(self, x): + return self.fc_q(x) + + def predict(self, features, aux_features): + aux_features.append(features) + phi = torch.cat(aux_features, dim=1) + + pre_prob = self.fc_actor(phi) + prob = F.softmax(pre_prob, dim=1) + log_prob = F.log_softmax(pre_prob, dim=1) + value = self.fc_critic(phi) + return prob, log_prob, value + + def feature(self, x): + return self.forward(x) + + +class GammaAttentionNet(BasicNet): + def predict(self, features, aux_features): + attention = self.compute_attention(features) + + aux_features = torch.stack(aux_features) + aux_features = aux_features * attention.t().unsqueeze(-1) + aux_features = aux_features.transpose(0, 1).contiguous().sum(1) + + phi = features + aux_features + pre_prob = self.fc_actor(phi) + prob = F.softmax(pre_prob, dim=1) + log_prob = F.log_softmax(pre_prob, dim=1) + value = self.fc_critic(phi) + return prob, log_prob, value + + def compute_attention(self, phi): + attention = self.fc_attention(phi) + # attention = F.relu(attention) + # attention = F.tanh(attention) + # attention = (attention + 1) / 0.5 + # attention = F.tanh(attention) + # attention = F.tanh(attention) + # attention = F.sigmoid(attention) + attention = F.softmax(attention, dim=1) + # max_attention = 10 + # cond = (attention < max_attention).float().detach() + # attention = attention * cond + max_attention * (1 - cond) + # cond = (attention > -max_attention).float().detach() + # attention = attention * cond + -max_attention * (1 - cond) + # self.attention = attention.data.cpu().numpy() + return attention + + def q(self, x): + return self.fc_q(x) + + def feature(self, x): + return self.forward(x) diff --git a/network/conv_network.py b/network/conv_network.py index 2530a7a..b1f2fc4 100644 --- a/network/conv_network.py +++ b/network/conv_network.py @@ -183,4 +183,60 @@ class QuantileConvNet(nn.Module, QuantileNet): y = y.view(y.size(0), -1) y = F.relu(self.fc4(y)) y = self.fc5(y) - return y \ No newline at end of file + return y + +class GammaConvNet(nn.Module, GammaNet): + def __init__(self, in_channels, action_dim, num_peers, gpu=-1): + super(GammaConvNet, self).__init__() + hidden_size = 512 + 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, 32, kernel_size=3, stride=1) + self.fc4 = nn.Linear(7 * 7 * 32, hidden_size) + + self.fc_actor = nn.Linear(hidden_size * num_peers, action_dim) + self.fc_critic = nn.Linear(hidden_size * num_peers, 1) + + self.fc_attention = nn.Linear(hidden_size, num_peers - 1) + self.fc_q = nn.Linear(hidden_size, action_dim) + + self.fc_actor_main = nn.Linear(hidden_size, action_dim) + self.fc_critic_main = nn.Linear(hidden_size, 1) + self.compute_attention = self.softmax_attention + BasicNet.__init__(self, gpu=gpu) + + def forward(self, x, update_lstm=True): + x = self.variable(x) + x = F.relu(self.conv1(x)) + x = F.relu(self.conv2(x)) + x = F.relu(self.conv3(x)) + x = x.view(x.size(0), -1) + phi = F.relu(self.fc4(x)) + return phi + +class GammaAttentionConvNet(nn.Module, GammaAttentionNet): + def __init__(self, in_channels, action_dim, num_peers, gpu=-1): + super(GammaAttentionConvNet, self).__init__() + hidden_size = 512 + 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, 32, kernel_size=3, stride=1) + self.fc4 = nn.Linear(7 * 7 * 32, hidden_size) + + self.fc_actor = nn.Linear(hidden_size, action_dim) + self.fc_critic = nn.Linear(hidden_size, 1) + + self.fc_attention = nn.Linear(hidden_size, num_peers - 1) + self.fc_q = nn.Linear(hidden_size, action_dim) + + BasicNet.__init__(self, gpu=gpu) + self.fc_attention.weight.data.zero_() + + def forward(self, x, update_lstm=True): + x = self.variable(x) + x = F.relu(self.conv1(x)) + x = F.relu(self.conv2(x)) + x = F.relu(self.conv3(x)) + x = x.view(x.size(0), -1) + phi = F.relu(self.fc4(x)) + return phi diff --git a/network/shallow_network.py b/network/shallow_network.py index 294a2f9..f1ba44f 100644 --- a/network/shallow_network.py +++ b/network/shallow_network.py @@ -40,7 +40,7 @@ class DuelingFCNet(nn.Module, DuelingNet): # Network for CartPole with actor critic class ActorCriticFCNet(nn.Module, ActorCriticNet): - def __init__(self, state_dim, action_dim): + def __init__(self, state_dim, action_dim, gpu=-1): super(ActorCriticFCNet, self).__init__() hidden_size1 = 64 hidden_size2 = 64 @@ -48,7 +48,7 @@ class ActorCriticFCNet(nn.Module, ActorCriticNet): self.fc2 = nn.Linear(hidden_size1, hidden_size2) self.fc_actor = nn.Linear(hidden_size2, action_dim) self.fc_critic = nn.Linear(hidden_size2, 1) - BasicNet.__init__(self, False) + BasicNet.__init__(self, gpu=gpu) def forward(self, x, update_LSTM=True): x = self.variable(x) @@ -91,4 +91,47 @@ class QuantileFCNet(nn.Module, QuantileNet): phi = F.relu(self.fc1(x)) phi = F.relu(self.fc2(phi)) quantiles = self.fc3(phi) - return quantiles \ No newline at end of file + return quantiles + +class GammaFCNet(nn.Module, GammaNet): + def __init__(self, state_dim, action_dim, num_peers, gpu=-1): + super(GammaFCNet, self).__init__() + hidden_size = 64 + self.fc1 = nn.Linear(state_dim, hidden_size) + self.fc2 = nn.Linear(hidden_size, hidden_size) + self.fc_actor = nn.Linear(hidden_size * num_peers, action_dim) + self.fc_critic = nn.Linear(hidden_size * num_peers, 1) + + self.fc_attention = nn.Linear(hidden_size, num_peers - 1) + self.fc_q = nn.Linear(hidden_size, action_dim) + + # self.fc_actor_main = nn.Linear(hidden_size, action_dim) + # self.fc_critic_main = nn.Linear(hidden_size, 1) + self.compute_attention = self.softmax_attention + BasicNet.__init__(self, gpu=gpu) + + def forward(self, x, update_lstm=True): + x = self.variable(x) + x = F.relu(self.fc1(x)) + x = F.relu(self.fc2(x)) + return x + +class GammaAttentionFCNet(nn.Module, GammaAttentionNet): + def __init__(self, state_dim, action_dim, num_peers, gpu=-1): + super(GammaAttentionFCNet, self).__init__() + hidden_size = 64 + self.fc1 = nn.Linear(state_dim, hidden_size) + self.fc2 = nn.Linear(hidden_size, hidden_size) + self.fc_actor = nn.Linear(hidden_size, action_dim) + self.fc_critic = nn.Linear(hidden_size, 1) + + self.fc_attention = nn.Linear(hidden_size, num_peers - 1) + self.fc_q = nn.Linear(hidden_size, action_dim) + BasicNet.__init__(self, gpu=gpu) + self.fc_attention.weight.data.zero_() + + def forward(self, x, update_lstm=True): + x = self.variable(x) + x = F.relu(self.fc1(x)) + x = F.relu(self.fc2(x)) + return x diff --git a/utils/config.py b/utils/config.py index 183bf7e..b250e2f 100644 --- a/utils/config.py +++ b/utils/config.py @@ -24,7 +24,7 @@ class Config: self.exploration_steps = 0 self.logger = None self.history_length = 1 - self.test_interval = 100 + self.test_interval = 0 self.test_repetitions = 50 self.double_q = False self.tag = 'vanilla' diff --git a/utils/misc.py b/utils/misc.py index 8dbdeda..41ec712 100644 --- a/utils/misc.py +++ b/utils/misc.py @@ -77,7 +77,17 @@ def run_iterations(agent): pickle.dump({'rewards': rewards, 'steps': steps}, f) agent.save('data/%s-%s-model-%s.bin' % (agent_name, config.tag, agent.task.name)) + if config.test_interval and iteration % config.test_interval == 0: + test_rewards, test_steps = agent.evaluate() + config.logger.info('total steps %d, test reward %f, test steps %d' % ( + agent.total_steps, test_rewards, test_steps + )) iteration += 1 + if config.max_steps and agent.total_steps >= config.max_steps: + agent.close() + break + + return steps, rewards def sync_grad(target_network, src_network): for param, src_param in zip(target_network.parameters(), src_network.parameters()): diff --git a/utils/plot.py b/utils/plot.py index 75b739e..25ed041 100644 --- a/utils/plot.py +++ b/utils/plot.py @@ -50,7 +50,7 @@ def plot_curves(xy_list, xaxis, title): plt.ylabel("Episode Rewards") plt.tight_layout() -def plot_results(dirs, num_timesteps, xaxis, task_name): +def plot_results(dirs, num_timesteps=1e8, xaxis=X_TIMESTEPS, task_name=''): tslist = [] for dir in dirs: ts = load_results(dir)