File size: 4,098 Bytes
f9c98cb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
107
108
import argparse
import logging
import time

from core import *
from managers import *
from utils import *
import torch

logging.basicConfig(level=logging.INFO)

parser = argparse.ArgumentParser(description='TransE model')

parser.add_argument("--experiment_name", type=str, default="RK_ComplEx",
                    help="A folder with this name would be created to dump saved models and log files")
parser.add_argument("--dataset", "-d", type=str, default="RK",
                    help="Dataset string")
parser.add_argument("--model", "-m", type=str, default="ComplEx",
                    help="Model to use")

parser.add_argument("--nEpochs", type=int, default=2000,
                    help="Learning rate of the optimizer")
parser.add_argument("--nBatches", type=int, default=25,
                    help="Batch size")
parser.add_argument("--eval_every", type=int, default=10,
                    help="Interval of epochs to evaluate the model?")
parser.add_argument("--save_every", type=int, default=50,
                    help="Interval of epochs to save a checkpoint of the model?")
parser.add_argument('--eval_mode', type=str, default="head",
                    help='Evaluate on head and/or tail prediction?')

parser.add_argument("--sample_size", type=int, default=0,
                    help="No. of negative samples to compare to for MRR/MR/Hit@10")
parser.add_argument("--patience", type=int, default=10,
                    help="Early stopping patience")
parser.add_argument("--margin", type=int, default=1,
                    help="The margin between positive and negative samples in the max-margin loss")
parser.add_argument("--p_norm", type=int, default=1,
                    help="The norm to use for the distance metric")
parser.add_argument("--optimizer", type=str, default="SGD",
                    help="Which optimizer to use? SGD/Adam")
parser.add_argument("--embedding_dim", type=int, default=100,
                    help="Entity and relations embedding size")
parser.add_argument("--lr", type=float, default=0.01,
                    help="Learning rate of the optimizer")
parser.add_argument("--momentum", type=float, default=0,
                    help="Momentum of the SGD optimizer")
parser.add_argument("--lmbda", type=float, default=0,
                    help="Regularization constant")

parser.add_argument("--debug", type=bool_flag, default=False,
                    help="Run the code in debug mode?")
parser.add_argument('--disable-cuda', action='store_true',
                    help='Disable CUDA')
parser.add_argument('--filter', action='store_true',
                    help='Filter the samples while evaluation')

params = parser.parse_args()

initialize_experiment(params)

params.device = None
if not params.disable_cuda and torch.cuda.is_available():
    params.device = torch.device('cuda')
else:
    params.device = torch.device('cpu')

# params.batch_size = int(len(train_data_sampler.data) / params.nBatches)

logging.info(params.device)

data_sampler = DataSampler(params)
model = initialize_model(params)
trainer = Trainer(model, data_sampler, params)
evaluator = Evaluator(model, data_sampler, params)

logging.info('Starting training...')

# tb_logger = Logger(params.exp_dir)

for e in range(params.nEpochs):
    res = 0
    tic = time.time()
    model.train()
    loss, auc = trainer.one_epoch()
    toc = time.time()

    # tb_logger.scalar_summary('loss', loss, e)

    logging.info('Epoch %d with loss: %f, AUC: %f in %f'
                 % (e, loss, auc, toc - tic))

    if (e + 1) % params.eval_every == 0:
        tic = time.time()
        model.eval()
        log_data = evaluator.get_log_data('test')
        toc = time.time()
        logging.info('Performance: %s in %f' % (str(log_data), (toc - tic)))

        # for tag, value in log_data.items():
        #     tb_logger.scalar_summary(tag, value, e + 1)

        to_continue = trainer.select_model(log_data)
        if not to_continue:
            break
    if (e + 1) % params.save_every == 0:
        torch.save(model, os.path.join(params.exp_dir, 'checkpoint.pth'))