-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtraintestSplit.py
More file actions
40 lines (34 loc) · 1.21 KB
/
Copy pathtraintestSplit.py
File metadata and controls
40 lines (34 loc) · 1.21 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
import numpy as np
def split(frac_training,frac_valid, datafile, file_train, file_valid, file_test):
f_train = open(file_train, "w")
f_valid = open(file_valid, "w")
f_test = open(file_test, "w")
ntrain = 0
nvalid = 0
ntest = 0
with open(datafile, "r") as f:
for line in f:
setID = np.argmax(np.random.multinomial(1, [frac_training, frac_valid, 1-frac_training-frac_valid], size = 1))
if setID == 0:
f_train.write(line)
ntrain += 1
elif setID == 1:
f_valid.write(line)
nvalid += 1
elif setID == 2:
f_test.write(line)
ntest += 1
else:
print "error"
print ntrain
print nvalid
print ntest
if __name__ == "__main__":
frac_training = 0.7
frac_valid = 0.1
datafile = "data/reaction_NYTWaPoWSJ_K10"
split(frac_training = frac_training, frac_valid = frac_valid,
datafile = datafile,
file_train = datafile + "_" + str(frac_training)+"train",
file_valid = datafile + "_" + str(frac_valid)+"valid",
file_test = datafile + "_" + str(1- frac_training - frac_valid)+"test")