-
Notifications
You must be signed in to change notification settings - Fork 4
Expand file tree
/
Copy pathmain.py
More file actions
50 lines (43 loc) · 1.79 KB
/
Copy pathmain.py
File metadata and controls
50 lines (43 loc) · 1.79 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
from molecule_transformer_trainer import MoleculeTransformerTrainer
from model import MoleculeTransformer
import argparse
import sys
import glob
parser = argparse.ArgumentParser(description="Molecule transformer training")
parser.add_argument("-i", "--input", help="Input training SMILES file.")
parser.add_argument("-e", "--epochs", type=int, default=3)
parser.add_argument("-v", "--vocab", help="Vocabulary file")
parser.add_argument("--lossWeight", choices=['none','log','sqrt','raw'], default='none',
help="The type of class weights for the cross entropy loss.")
args = parser.parse_args()
if args.input == None:
parser.print_help()
sys.exit()
trainer = MoleculeTransformerTrainer(args.input,
class_weight=args.lossWeight,
vocab_file=args.vocab,
n_tokens=415)
print("Dataset split. . . .")
directory = MoleculeTransformerTrainer.split_file(args.input, line_num=1000000)
datasets = glob.glob(directory + "/*")
print("Splited datasets:" + str(datasets))
if args.vocab != None:
trainer.load_vocab()
print(trainer.smile_mol_tokenizer.vocab.stoi)
print(len(trainer.smile_mol_tokenizer.vocab.stoi))
print("Training processes . . . .")
build=False
for i in range(args.epochs):
for data in datasets:
print("Load the dataset {}".format(data))
trainer.gen_dataloader(data)
print("Training for splited file {}".format(data))
if not build:
trainer.build_model()
trainer.model_summary()
build = True
trainer.train(log='stdout')
trainer.save_model("model{}".format(i+1))
acc = trainer.evaluate_acc()
print("ACCURACY of epoch {} = {}".format(i+1, acc))
trainer.export_training_figure(i)