· 8 years ago · Mar 12, 2018, 12:44 AM
1""" Class for tracking experiments in a local sqlite database """
2
3class SqliteExperiment:
4 def __init__(self, hparams, metrics, experiment_id=None):
5 self.experiment_id = experiment_id or str(uuid4())
6 self.hparams = hparams
7 self.metrics = metrics
8 self.metric_names = ['experiment_id', 'measured_at'] + [n for n, t in metrics]
9 self.log_every = int(os.environ.get('LOG_EVERY', 10000))
10 self.last_log = None
11 self.last_epoch = None
12 self.db = sqlite3.connect('experiments.sqlite')
13 self.ensure_tables()
14
15 def ensure_tables(self):
16 """ Create tables for metrics and hyper params if they don't exist """
17 self.db.execute('''
18CREATE TABLE IF NOT EXISTS hparams (
19 experiment_id text primary key,
20 {}
21)
22 '''.format(self.to_sql_column_defs(self.hparams).strip(',')))
23
24 self.db.execute('''
25CREATE TABLE IF NOT EXISTS metrics (
26 experiment_id text,
27 measured_at int,
28 {}
29)
30 '''.format(self.to_sql_column_defs(self.metrics).strip(',')))
31 self.db.commit()
32
33 @classmethod
34 def to_sqlite_col_type(cls, col_type):
35 return {
36 int: 'integer',
37 float: 'real',
38 str: 'text',
39 bool: 'integer',
40 }[col_type]
41
42 def to_sql_column_defs(self, spec):
43 return ',\n'.join([
44 '{} {}'.format(col_name, self.to_sqlite_col_type(col_type))
45 for col_name, col_type in spec
46 ]) + ','
47
48 def log_hparams(self, hparams):
49 hparam_values = [self.experiment_id] + [hparams[name] for name, _type in self.hparams]
50 for idx, hparam in enumerate(hparam_values):
51 if isinstance(hparam, list):
52 hparam_values[idx] = ','.join(map(str, hparam))
53 self.db.execute('''
54insert into hparams values ({})
55 '''.format(', '.join(['?'] * len(hparam_values))), hparam_values)
56 self.db.commit()
57
58 def should_log(self, epoch, step):
59 should_log = False
60 if (self.last_epoch is None or self.last_log is None) \
61 or (self.last_epoch < epoch) \
62 or ((self.last_log + self.log_every) < step):
63 self.last_epoch = epoch
64 self.last_log = step
65 should_log = True
66 return should_log
67
68 def log_metrics(self, epoch, step, metrics, force=False):
69 if not force and not self.should_log(epoch, step):
70 return
71 metric_values = [self.experiment_id, time.time()] + \
72 [metrics.get(name) for name, _type in self.metrics]
73 self.db.execute(
74 '''
75insert into metrics ({}) values ({})
76 '''.format(', '.join(self.metric_names), ', '.join(['?'] * len(metric_values))),
77 metric_values)
78 self.db.commit()
79
80
81# Example Usage:
82################
83
84sle = SqliteExperiment(
85 [('vocab_size', int), ('msg_len', int), ('context_dim', int),
86 ('embed_dim', int), ('batch_size', int)],
87 [('loss', float), ('dev_loss', float), ('epoch', int),
88 ('acc', float), ('dev_acc', float)],
89 os.environ.get('EXPERIMENT_ID'))
90
91sle.log_hparams({'vocab_size': 16384, 'msg_len': 100, 'context_dim': 100,
92 'embed_dim': 200, 'batch_size': 512})
93
94def epoch_callback(loss, dev_loss, epoch, acc, dev_acc):
95 sle.log_metrics(epoch, train_x.shape[0],
96 {'loss': loss, 'dev_loss': dev_loss, 'acc': acc, 'dev_acc': dev_acc})
97
98my_model.train(train_x, train_y, callback=epoch_callback)