diff --git a/component/__init__.py b/component/__init__.py index fe77529..b51b19c 100644 --- a/component/__init__.py +++ b/component/__init__.py @@ -2,4 +2,5 @@ from .atari_wrapper import * from .policy import * from .replay import * from .task import * -from .random_process import * \ No newline at end of file +from .random_process import * +from .bench import * \ No newline at end of file diff --git a/component/bench.py b/component/bench.py new file mode 100644 index 0000000..59e68fe --- /dev/null +++ b/component/bench.py @@ -0,0 +1,138 @@ +# from https://raw.githubusercontent.com/openai/baselines/master/baselines/bench/monitor.py + +__all__ = ['Monitor', 'get_monitor_files', 'load_results'] + +import gym +from gym.core import Wrapper +import time +from glob import glob +import csv +import os.path as osp +import json + +class Monitor(Wrapper): + EXT = "monitor.csv" + f = None + + def __init__(self, env, filename, allow_early_resets=True, reset_keywords=(), info_keywords=()): + Wrapper.__init__(self, env=env) + self.tstart = time.time() + if filename is None: + self.f = None + self.logger = None + else: + if not filename.endswith(Monitor.EXT): + if osp.isdir(filename): + filename = osp.join(filename, Monitor.EXT) + else: + filename = filename + "." + Monitor.EXT + self.f = open(filename, "wt") + self.f.write('#%s\n'%json.dumps({"t_start": self.tstart, 'env_id' : env.spec and env.spec.id})) + self.logger = csv.DictWriter(self.f, fieldnames=('r', 'l', 't')+reset_keywords+info_keywords) + self.logger.writeheader() + self.f.flush() + + self.reset_keywords = reset_keywords + self.info_keywords = info_keywords + self.allow_early_resets = allow_early_resets + self.rewards = None + self.needs_reset = True + self.episode_rewards = [] + self.episode_lengths = [] + self.episode_times = [] + self.total_steps = 0 + self.current_reset_info = {} # extra info about the current episode, that was passed in during reset() + + def reset(self, **kwargs): + if not self.allow_early_resets and not self.needs_reset: + raise RuntimeError("Tried to reset an environment before done. If you want to allow early resets, wrap your env with Monitor(env, path, allow_early_resets=True)") + self.rewards = [] + self.needs_reset = False + for k in self.reset_keywords: + v = kwargs.get(k) + if v is None: + raise ValueError('Expected you to pass kwarg %s into reset'%k) + self.current_reset_info[k] = v + return self.env.reset(**kwargs) + + def step(self, action): + if self.needs_reset: + raise RuntimeError("Tried to step environment that needs reset") + ob, rew, done, info = self.env.step(action) + self.rewards.append(rew) + if done: + self.needs_reset = True + eprew = sum(self.rewards) + eplen = len(self.rewards) + epinfo = {"r": round(eprew, 6), "l": eplen, "t": round(time.time() - self.tstart, 6)} + for k in self.info_keywords: + epinfo[k] = info[k] + self.episode_rewards.append(eprew) + self.episode_lengths.append(eplen) + self.episode_times.append(time.time() - self.tstart) + epinfo.update(self.current_reset_info) + if self.logger: + self.logger.writerow(epinfo) + self.f.flush() + info['episode'] = epinfo + self.total_steps += 1 + return (ob, rew, done, info) + + def close(self): + if self.f is not None: + self.f.close() + + def get_total_steps(self): + return self.total_steps + + def get_episode_rewards(self): + return self.episode_rewards + + def get_episode_lengths(self): + return self.episode_lengths + + def get_episode_times(self): + return self.episode_times + +class LoadMonitorResultsError(Exception): + pass + +def get_monitor_files(dir): + return glob(osp.join(dir, "*" + Monitor.EXT)) + +def load_results(dir): + import pandas + monitor_files = ( + glob(osp.join(dir, "*monitor.json")) + + glob(osp.join(dir, "*monitor.csv"))) # get both csv and (old) json files + if not monitor_files: + raise LoadMonitorResultsError("no monitor files of the form *%s found in %s" % (Monitor.EXT, dir)) + dfs = [] + headers = [] + for fname in monitor_files: + with open(fname, 'rt') as fh: + if fname.endswith('csv'): + firstline = fh.readline() + assert firstline[0] == '#' + header = json.loads(firstline[1:]) + df = pandas.read_csv(fh, index_col=None) + headers.append(header) + elif fname.endswith('json'): # Deprecated json format + episodes = [] + lines = fh.readlines() + header = json.loads(lines[0]) + headers.append(header) + for line in lines[1:]: + episode = json.loads(line) + episodes.append(episode) + df = pandas.DataFrame(episodes) + else: + assert 0, 'unreachable' + df['t'] += header['t_start'] + dfs.append(df) + df = pandas.concat(dfs) + df.sort_values('t', inplace=True) + df.reset_index(inplace=True) + df['t'] -= min(header['t_start'] for header in headers) + # df.headers = headers # HACK to preserve backwards compatibility + return df diff --git a/component/task.py b/component/task.py index 187231f..6e747a9 100644 --- a/component/task.py +++ b/component/task.py @@ -9,6 +9,8 @@ import numpy as np from .atari_wrapper import * import multiprocessing as mp import sys +from .bench import Monitor +from utils import * class BasicTask: def __init__(self, max_steps=sys.maxsize): @@ -32,6 +34,9 @@ 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) @@ -138,9 +143,11 @@ class Roboschool(BasicTask): def step(self, action): return BasicTask.step(self, np.clip(action, -1, 1)) -def sub_task(parent_pipe, pipe, task_fn): +def sub_task(parent_pipe, pipe, task_fn, filename=None): parent_pipe.close() task = task_fn() + if filename is not None: + task.set_monitor(filename) task.env.seed(np.random.randint(0, sys.maxsize)) while True: op, data = pipe.recv() @@ -155,12 +162,17 @@ def sub_task(parent_pipe, pipe, task_fn): assert False, 'Unknown Operation' class ParallelizedTask: - def __init__(self, task_fn, num_workers): + def __init__(self, task_fn, num_workers, tag='vanilla'): 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)] self.pipes, worker_pipes = zip(*[mp.Pipe() for _ in range(num_workers)]) - args = [(p, wp, task_fn) for p, wp in zip(self.pipes, worker_pipes)] + args = [(p, wp, task_fn, filename) + for p, wp, filename in zip(self.pipes, worker_pipes, filenames)] 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() diff --git a/main.py b/main.py index 4095200..657de7d 100644 --- a/main.py +++ b/main.py @@ -11,8 +11,9 @@ from utils import * import model.action_conditional_video_prediction as acvp def dqn_cart_pole(): + game = 'CartPole-v0' config = Config() - config.task_fn = lambda: ClassicalControl('CartPole-v0', max_steps=200) + config.task_fn = lambda: ClassicalControl(game, max_steps=200) config.optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001) config.network_fn = lambda: FCNet([4, 50, 200, 2]) # config.network_fn = lambda: DuelingFCNet([8, 50, 200, 2]) diff --git a/utils/plot.py b/utils/plot.py new file mode 100644 index 0000000..75b739e --- /dev/null +++ b/utils/plot.py @@ -0,0 +1,64 @@ +# from https://raw.githubusercontent.com/openai/baselines/master/baselines/results_plotter.py +__all__ = ['plot_results'] + +import numpy as np +import matplotlib.pyplot as plt +from component import load_results +plt.rcParams['svg.fonttype'] = 'none' + +X_TIMESTEPS = 'timesteps' +X_EPISODES = 'episodes' +X_WALLTIME = 'walltime_hrs' +POSSIBLE_X_AXES = [X_TIMESTEPS, X_EPISODES, X_WALLTIME] +EPISODES_WINDOW = 100 +COLORS = ['blue', 'green', 'red', 'cyan', 'magenta', 'yellow', 'black', 'purple', 'pink', + 'brown', 'orange', 'teal', 'coral', 'lightblue', 'lime', 'lavender', 'turquoise', + 'darkgreen', 'tan', 'salmon', 'gold', 'lightpurple', 'darkred', 'darkblue'] + +def rolling_window(a, window): + shape = a.shape[:-1] + (a.shape[-1] - window + 1, window) + strides = a.strides + (a.strides[-1],) + return np.lib.stride_tricks.as_strided(a, shape=shape, strides=strides) + +def window_func(x, y, window, func): + yw = rolling_window(y, window) + yw_func = func(yw, axis=-1) + return x[window-1:], yw_func + +def ts2xy(ts, xaxis): + if xaxis == X_TIMESTEPS: + x = np.cumsum(ts.l.values) + y = ts.r.values + elif xaxis == X_EPISODES: + x = np.arange(len(ts)) + y = ts.r.values + elif xaxis == X_WALLTIME: + x = ts.t.values / 3600. + y = ts.r.values + else: + raise NotImplementedError + return x, y + +def plot_curves(xy_list, xaxis, title): + for (i, (x, y)) in enumerate(xy_list): + color = COLORS[i] + # plt.scatter(x, y, s=2) + x, y_mean = window_func(x, y, EPISODES_WINDOW, np.mean) #So returns average of last EPISODE_WINDOW episodes + plt.plot(x, y_mean, color=color) + plt.title(title) + plt.xlabel(xaxis) + plt.ylabel("Episode Rewards") + plt.tight_layout() + +def plot_results(dirs, num_timesteps, xaxis, task_name): + tslist = [] + for dir in dirs: + ts = load_results(dir) + ts = ts[ts.l.cumsum() <= num_timesteps] + tslist.append(ts) + xy_list = [ts2xy(ts, xaxis) for ts in tslist] + plot_curves(xy_list, xaxis, task_name) + +if __name__ == '__main__': + plot_results(['../log/CartPole-v0-vanilla'], 10e6, X_TIMESTEPS, "CartPole") + plt.show()