-
Notifications
You must be signed in to change notification settings - Fork 14
Expand file tree
/
Copy pathpredict.py
More file actions
46 lines (37 loc) · 1.38 KB
/
Copy pathpredict.py
File metadata and controls
46 lines (37 loc) · 1.38 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
import csv
import numpy as np
from fpgnn.tool import set_predict_argument, get_scaler, load_args, load_data, load_model
from fpgnn.train import predict
from fpgnn.data import MoleDataSet
def predicting(args):
print('Load args.')
scaler = get_scaler(args.model_path)
print('scaler',scaler)
train_args = load_args(args.model_path)
for key,value in vars(train_args).items():
if not hasattr(args, key):
setattr(args, key, value)
print('Load data.')
test_data = load_data(args.predict_path,args)
print('Load model')
model = load_model(args.model_path,args.cuda)
test_pred = predict(model,test_data,args.batch_size,scaler)
assert len(test_data) == len(test_pred)
test_pred = np.array(test_pred)
test_pred = test_pred.tolist()
print('Write result.')
write_smile = test_data.smile()
with open(args.result_path, 'w',newline = '') as file:
writer = csv.writer(file)
line = ['Smiles']
line.extend(args.task_names)
writer.writerow(line)
#for i in range(fir_data_len):
for i in range(len(test_data)):
line = []
line.append(write_smile[i])
line.extend(test_pred[i])
writer.writerow(line)
if __name__=='__main__':
args = set_predict_argument()
predicting(args)