Add baselines' plotter

This commit is contained in:
Shangtong Zhang
2018-04-02 15:22:56 -06:00
parent e7e94d715b
commit e427e8f73f
5 changed files with 221 additions and 5 deletions
+2 -1
View File
@@ -2,4 +2,5 @@ from .atari_wrapper import *
from .policy import *
from .replay import *
from .task import *
from .random_process import *
from .random_process import *
from .bench import *
+138
View File
@@ -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
+15 -3
View File
@@ -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()
+2 -1
View File
@@ -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])
+64
View File
@@ -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()