mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Add baselines' plotter
This commit is contained in:
@@ -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 *
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user