mirror of
https://github.com/wassname/keras-language-modeling.git
synced 2026-09-12 12:33:54 +08:00
updated printing
This commit is contained in:
+105
-7
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user