- Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtraining_script.py
More file actions
Latest commit
106 lines (81 loc) · 3.44 KB
/
Copy pathtraining_script.py
File metadata and controls
106 lines (81 loc) · 3.44 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
importdata_loader
importRNN
importipdb
importmatplotlib.pyplotasplt
importargparse
importos, pickle, utils
if__name__=="__main__":
parser=argparse.ArgumentParser()
parser.add_argument("--mSize", type=int, default=5)
parser.add_argument("--eSize", type=int, default=-1)
parser.add_argument("--nbEpoch", type=int, default=10)
parser.add_argument("--dataFolder", default="data")
parser.add_argument("--hSize", type=int, default=100)
parser.add_argument("--embSize", type=int, default=100)
parser.add_argument("--lr", type=float, default=0.02)
parser.add_argument("--savef", default="saving")
parser.add_argument("--rnnClass", default="RNN")
args=parser.parse_args()
# Parameters and stuff
v_size=10000000
m_size=args.mSize# minibatch size
epoch_size=args.eSize# The number of minibatch in an epoch
nbEpoch=args.nbEpoch# The number of epoch to do
folder=args.dataFolder
h_size=args.hSize
embSize=args.embSize
lr=args.lr
save_folder=args.savef
rnnClass=args.rnnClass
data_iterator_type=data_loader.predict_next_iterator
ifrnnClass=="DAE":
data_iterator_type=data_loader.predict_noisy_self
rnn_class=RNN.RNN
exec"rnn_class = RNN.{}".format(rnnClass)
d=data_loader.data_crawler(folder=folder, maxCount=v_size)
#All the datasets
trainset=data_loader.data_iterator(data=d.all_data[0], e_size=epoch_size,
m_size=m_size, vocab=d.vocab,
wordMapping=d.wordMapping)
validset=data_loader.data_iterator(data=d.all_data[1], e_size=epoch_size,
m_size=m_size, vocab=d.vocab,
wordMapping=d.wordMapping)
testset=data_loader.data_iterator(data=d.all_data[2], e_size=epoch_size,
m_size=m_size, vocab=d.vocab,
wordMapping=d.wordMapping)
trainset=data_loader.predict_noisy_self(trainset)
validset=data_loader.predict_noisy_self(validset)
testset=data_loader.predict_noisy_self(testset)
r_layer=rnn_class(h_size=h_size, e_size=embSize, v_size=d.nbWords, name="DAE_1")
mlp=RNN.MLP(v_size=d.nbWords, lr=lr)
mlp.layers.append(r_layer)
# Training
trainingLosses, validLosses=mlp.train(nbEpoch, trainset, validset, d, save_folder)
#Getting some prediction, for fun.
noS=1
fori, sentenceinzip(range(1), trainset):
#ipdb.set_trace()
try:
printtrainset.iterator.switchRep(sentence[0][:,0,:])
pred=mlp.predict(sentence)[0]
printtrainset.iterator.switchRep(pred)
except:
ipdb.set_trace()
mlp2=RNN.MLP(v_size=d.nbWords, lr=lr)
mlp2.load("saving/rnn.pkl")
ipdb.set_trace()
testPer=mlp.getPerplexity(testset)
testLoss=mlp.getLoss(testset)
print"The perplexity is: {}, the loss is: {}".format(testPer, testLoss)
# Showing the loss, for fun.
#utils.save_everything(save_folder, mlp, d)
pickle.dump([trainingLosses, validLosses],
open(os.path.join(save_folder,"debug_loss"), 'w'))
print"It's working!!"
#plt.plot(trainingLosses)
#plt.plot(validLosses)
#plt.ylabel("Loss")
#plt.savefig("Losses.png")
#rnn, metadata = load_everything("saving")
#print rnn.getLoss(testset)
#plt.show()