diff --git a/insurance_qa_eval.py b/insurance_qa_eval.py index d08d2bc..efb1341 100644 --- a/insurance_qa_eval.py +++ b/insurance_qa_eval.py @@ -4,16 +4,114 @@ import os import sys import random -from time import strftime, gmtime +from time import strftime, gmtime, time import pickle import json +from keras.callbacks import ProgbarLogger +from keras.utils.generic_utils import Progbar from scipy.stats import rankdata random.seed(42) +class ProgbarError(Progbar): + def update(self, current, values=[], force=False): + ''' + @param current: index of current step + @param values: list of tuples (name, value_for_last_step). + The progress bar will display averages for these values. + @param force: force visual progress update + ''' + for k, v in values: + if k not in self.sum_values: + self.sum_values[k] = [v * (current - self.seen_so_far), current - self.seen_so_far] + self.unique_values.append(k) + else: + self.sum_values[k][0] += v * (current - self.seen_so_far) + self.sum_values[k][1] += (current - self.seen_so_far) + self.seen_so_far = current + + now = time() + if self.verbose == 1: + if not force and (now - self.last_update) < self.interval: + return + + prev_total_width = self.total_width + sys.stderr.write("\b" * prev_total_width) + sys.stderr.write("\r") + + numdigits = int(np.floor(np.log10(self.target))) + 1 + barstr = '%%%dd/%%%dd [' % (numdigits, numdigits) + bar = barstr % (current, self.target) + prog = float(current) / self.target + prog_width = int(self.width * prog) + if prog_width > 0: + bar += ('=' * (prog_width-1)) + if current < self.target: + bar += '>' + else: + bar += '=' + bar += ('.' * (self.width - prog_width)) + bar += ']' + sys.stderr.write(bar) + self.total_width = len(bar) + + if current: + time_per_unit = (now - self.start) / current + else: + time_per_unit = 0 + eta = time_per_unit * (self.target - current) + info = '' + if current < self.target: + info += ' - ETA: %ds' % eta + else: + info += ' - %ds' % (now - self.start) + for k in self.unique_values: + info += ' - %s:' % k + if type(self.sum_values[k]) is list: + avg = self.sum_values[k][0] / max(1, self.sum_values[k][1]) + if abs(avg) > 1e-3: + info += ' %.4f' % avg + else: + info += ' %.4e' % avg + else: + info += ' %s' % self.sum_values[k] + + self.total_width += len(info) + if prev_total_width > self.total_width: + info += ((prev_total_width - self.total_width) * " ") + + sys.stderr.write(info) + sys.stderr.flush() + + if current >= self.target: + sys.stderr.write("\n") + + if self.verbose == 2: + if current >= self.target: + info = '%ds' % (now - self.start) + for k in self.unique_values: + info += ' - %s:' % k + avg = self.sum_values[k][0] / max(1, self.sum_values[k][1]) + if avg > 1e-3: + info += ' %.4f' % avg + else: + info += ' %.4e' % avg + sys.stderr.write(info + "\n") + + self.last_update = now + + +class ProgbarErrorLogger(ProgbarLogger): + def on_epoch_begin(self, epoch, logs={}): + if self.verbose: + print('Epoch %d/%d' % (epoch + 1, self.nb_epoch), file=sys.stderr) + self.progbar = ProgbarError(target=self.params['nb_sample']) + self.seen = 0 + + class Evaluator: def __init__(self, conf, model, optimizer=None): try: @@ -87,10 +185,11 @@ class Evaluator: ##### Training ##### - def print_time(self): - print(strftime('%Y-%m-%d %H:%M:%S :: ', gmtime()), end='') + def get_time(self): + return strftime('%Y-%m-%d %H:%M:%S', gmtime()) def train(self): + print('Began training at %s' % self.get_time()) batch_size = self.params['batch_size'] nb_epoch = self.params['nb_epoch'] validation_split = self.params['validation_split'] @@ -123,10 +222,10 @@ class Evaluator: # bad_answers = self.pada(get_bad_samples(indices, top_50)) bad_answers = self.pada(random.sample(self.answers.values(), len(good_answers))) - print('Epoch %d :: ' % i, end='') - self.print_time() hist = self.model.fit([questions, good_answers, bad_answers], nb_epoch=1, batch_size=batch_size, - validation_split=validation_split) + validation_split=validation_split, verbose=0, callbacks=[ProgbarErrorLogger()]) + print('%s -- Epoch %d' % (self.get_time(), i), end='') + print(' Loss = {}'.format(hist.history['val_loss'][0])) if hist.history['val_loss'][0] < val_loss['loss']: val_loss = {'loss': hist.history['val_loss'][0], 'epoch': i} @@ -153,7 +252,6 @@ class Evaluator: def get_score(self, verbose=False): for name, data in self.eval_sets().items(): - self.print_time() print('----- %s -----' % name) random.shuffle(data)