mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-10 11:40:58 +08:00
Regenerate images
This commit is contained in:
+5
-4
@@ -1,5 +1,5 @@
|
||||
import numpy as np
|
||||
from component import load_results
|
||||
import component
|
||||
|
||||
class Plotter:
|
||||
COLORS = ['blue', 'green', 'red', 'cyan', 'magenta', 'yellow', 'black', 'purple', 'pink',
|
||||
@@ -40,7 +40,7 @@ class Plotter:
|
||||
def load_results(self, dirs, max_timesteps=1e8, x_axis=X_TIMESTEPS, episode_window=100):
|
||||
tslist = []
|
||||
for dir in dirs:
|
||||
ts = load_results(dir)
|
||||
ts = component.load_monitor_log(dir)
|
||||
ts = ts[ts.l.cumsum() <= max_timesteps]
|
||||
tslist.append(ts)
|
||||
xy_list = [self.ts2xy(ts, x_axis) for ts in tslist]
|
||||
@@ -49,9 +49,10 @@ class Plotter:
|
||||
|
||||
def plot_results(self, dirs, max_timesteps=1e8, x_axis=X_TIMESTEPS, episode_window=100):
|
||||
import matplotlib.pyplot as plt
|
||||
plt.ticklabel_format(axis='x', style='sci', scilimits=(1, 1))
|
||||
xy_list = self.load_results(dirs, max_timesteps, x_axis, episode_window)
|
||||
for (i, (x, y, y_mean)) in enumerate(xy_list):
|
||||
for (i, (x, y, smoothed)) in enumerate(xy_list):
|
||||
color = Plotter.COLORS[i]
|
||||
plt.plot(x, y_mean, color=color)
|
||||
plt.plot(smoothed[0], smoothed[1], color=color)
|
||||
plt.xlabel(x_axis)
|
||||
plt.ylabel("Episode Rewards")
|
||||
|
||||
Reference in New Issue
Block a user