updated printing

This commit is contained in:
Ben Bolte
2016-08-15 14:50:01 -07:00
parent 432a0a358e
commit 30b7afff9a
+105 -7
View File
@@ -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)