-
Notifications
You must be signed in to change notification settings - Fork 3
Expand file tree
/
Copy pathnnplot.py
More file actions
executable file
·78 lines (68 loc) · 2.03 KB
/
Copy pathnnplot.py
File metadata and controls
executable file
·78 lines (68 loc) · 2.03 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
#!/usr/bin/env python
#
# NN Plot
# Tim O'Shea (c) 2016
#
#
#
import matplotlib.pyplot as plt
import matplotlib.patches as patches
import matplotlib.lines as lines
# load the config def
import sys,json
fn = sys.argv[1]
print fn
cfg = json.loads(open(fn).read())
io = cfg["io"]
layers = cfg["layers"]
# ...
maxy = max(map(lambda x: x[3], layers))
print maxy
fig = plt.figure()
ax = fig.add_subplot(111)
for i,l in enumerate(layers):
ll=0.25 + i
w = 0.5
h = layers[i][3]
lr=(maxy-h)/2.0
print (ll,lr,w,h)
ax.add_patch(
patches.Rectangle( (ll, lr), w, h, fill=False)
)
ax.annotate(l[0], xy=(ll+w/2.0, lr+h/2.0),
ha='center', va='center',
rotation=90)
if(i < len(layers)-1):
h_n = layers[i+1][3]
lr_n=(maxy-h_n)/2.0
i_x = ll+w
o_x = ll+1
for j in range(0,layers[i][2]):
for k in range(0,layers[i+1][2]):
j_rel = (j*1.0/(layers[i][2]-1))
k_rel = (k*1.0/(layers[i+1][2]-1))
if(l[1]==None or abs(j_rel - k_rel)<l[1]):
i_ht = lr+h*j_rel
o_ht = lr_n+h_n*k_rel
ax.add_line(lines.Line2D([i_x,o_x], [i_ht,o_ht], color='black', linestyle='solid'))
if(i == 0):
for j in range(0,l[2]):
j_rel = (j*1.0/(l[2]-1))
ax.add_patch( patches.Circle([i_x-w,lr+h*j_rel], 0.025, color='blue', ec="none") )
if(i == len(layers)-2):
for j in range(0,layers[i+1][2]):
j_rel = (j*1.0/(layers[i+1][2]-1))
ax.add_patch( patches.Circle([o_x+w,lr_n+h_n*j_rel], 0.025, color='blue', ec="none") )
ax.annotate("Inputs (2x88)", xy=(0,maxy/2.0),
ha='center',va='center',
rotation=90)
ax.annotate("Outputs (2x88)", xy=(len(layers),maxy/2.0),
ha='center',va='center',
rotation=90)
ymarg = 0.05
xmarg = 0.05
plt.ylim(-ymarg, maxy+0.5+ymarg)
plt.xlim(-xmarg,len(layers)+xmarg)
plt.axis('off')
plt.savefig('test.png', bbix_inches='tight')
plt.show()