-
Notifications
You must be signed in to change notification settings - Fork 10
Expand file tree
/
Copy pathplane_plot.py
More file actions
114 lines (92 loc) · 4.09 KB
/
Copy pathplane_plot.py
File metadata and controls
114 lines (92 loc) · 4.09 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
import argparse
import numpy as np
import matplotlib
import matplotlib.pyplot as plt
import matplotlib.colors as colors
import os
import seaborn as sns
parser = argparse.ArgumentParser(description='Plane visualization')
parser.add_argument('--dir', type=str, default='/tmp/curve/', metavar='DIR',
help='training directory (default: None)')
args = parser.parse_args()
file = np.load(os.path.join(args.dir, 'plane.npz'))
matplotlib.rc('text', usetex=True)
matplotlib.rc('text.latex', preamble=[r'\usepackage{sansmath}', r'\sansmath'])
matplotlib.rc('font', **{'family':'sans-serif','sans-serif':['DejaVu Sans']})
matplotlib.rc('xtick.major', pad=12)
matplotlib.rc('ytick.major', pad=12)
matplotlib.rc('grid', linewidth=0.8)
sns.set_style('whitegrid')
class LogNormalize(colors.Normalize):
def __init__(self, vmin=None, vmax=None, clip=None, log_alpha=None):
self.log_alpha = log_alpha
colors.Normalize.__init__(self, vmin, vmax, clip)
def __call__(self, value, clip=None):
log_v = np.ma.log(value - self.vmin)
log_v = np.ma.maximum(log_v, self.log_alpha)
return 0.9 * (log_v - self.log_alpha) / (np.log(self.vmax - self.vmin) - self.log_alpha)
def plane(grid, values, vmax=None, log_alpha=-5, N=7, cmap='jet_r'):
cmap = plt.get_cmap(cmap)
if vmax is None:
clipped = values.copy()
else:
clipped = np.minimum(values, vmax)
log_gamma = (np.log(clipped.max() - clipped.min()) - log_alpha) / N
levels = clipped.min() + np.exp(log_alpha + log_gamma * np.arange(N + 1))
levels[0] = clipped.min()
levels[-1] = clipped.max()
levels = np.concatenate((levels, [1e10]))
norm = LogNormalize(clipped.min() - 1e-8, clipped.max() + 1e-8, log_alpha=log_alpha)
contour = plt.contour(grid[:, :, 0], grid[:, :, 1], values, cmap=cmap, norm=norm,
linewidths=2.5,
zorder=1,
levels=levels)
contourf = plt.contourf(grid[:, :, 0], grid[:, :, 1], values, cmap=cmap, norm=norm,
levels=levels,
zorder=0,
alpha=0.55)
colorbar = plt.colorbar(format='%.2g')
labels = list(colorbar.ax.get_yticklabels())
labels[-1].set_text(r'$>\,$' + labels[-2].get_text())
colorbar.ax.set_yticklabels(labels)
return contour, contourf, colorbar
plt.figure(figsize=(12.4, 7))
contour, contourf, colorbar = plane(
file['grid'],
file['tr_loss'],
vmax=5.0,
log_alpha=-5.0,
N=7
)
bend_coordinates = file['bend_coordinates']
curve_coordinates = file['curve_coordinates']
plt.scatter(bend_coordinates[[0, 2], 0], bend_coordinates[[0, 2], 1], marker='o', c='k', s=120, zorder=2)
plt.scatter(bend_coordinates[1, 0], bend_coordinates[1, 1], marker='D', c='k', s=120, zorder=2)
plt.plot(curve_coordinates[:, 0], curve_coordinates[:, 1], linewidth=4, c='k', label='$w(t)$', zorder=4)
plt.plot(bend_coordinates[[0, 2], 0], bend_coordinates[[0, 2], 1], c='k', linestyle='--', dashes=(3, 4), linewidth=3, zorder=2)
plt.margins(0.0)
plt.yticks(fontsize=18)
plt.xticks(fontsize=18)
colorbar.ax.tick_params(labelsize=18)
plt.savefig(os.path.join(args.dir, 'train_loss_plane.pdf'), format='pdf', bbox_inches='tight')
plt.show()
plt.figure(figsize=(12.4, 7))
contour, contourf, colorbar = plane(
file['grid'],
file['te_err'],
vmax=40,
log_alpha=-1.0,
N=7
)
bend_coordinates = file['bend_coordinates']
curve_coordinates = file['curve_coordinates']
plt.scatter(bend_coordinates[[0, 2], 0], bend_coordinates[[0, 2], 1], marker='o', c='k', s=120, zorder=2)
plt.scatter(bend_coordinates[1, 0], bend_coordinates[1, 1], marker='D', c='k', s=120, zorder=2)
plt.plot(curve_coordinates[:, 0], curve_coordinates[:, 1], linewidth=4, c='k', label='$w(t)$', zorder=4)
plt.plot(bend_coordinates[[0, 2], 0], bend_coordinates[[0, 2], 1], c='k', linestyle='--', dashes=(3, 4), linewidth=3, zorder=2)
plt.margins(0.0)
plt.yticks(fontsize=18)
plt.xticks(fontsize=18)
colorbar.ax.tick_params(labelsize=18)
plt.savefig(os.path.join(args.dir, 'test_error_plane.pdf'), format='pdf', bbox_inches='tight')
plt.show()