Main training loop. modified from Keras fit_generator
(name, g_train, d_train, sampler, generator, samples_per_epoch, nb_epoch,
z_dim=100, verbose=1, callbacks=[],
validation_data=None, nb_val_samples=None,
saver=None)
| 33 | |
| 34 | |
| 35 | def train_model(name, g_train, d_train, sampler, generator, samples_per_epoch, nb_epoch, |
| 36 | z_dim=100, verbose=1, callbacks=[], |
| 37 | validation_data=None, nb_val_samples=None, |
| 38 | saver=None): |
| 39 | """ |
| 40 | Main training loop. |
| 41 | modified from Keras fit_generator |
| 42 | """ |
| 43 | self = {} |
| 44 | epoch = 0 |
| 45 | counter = 0 |
| 46 | out_labels = ['g_loss', 'd_loss', 'd_loss_fake', 'd_loss_legit', 'time'] # self.metrics_names |
| 47 | callback_metrics = out_labels + ['val_' + n for n in out_labels] |
| 48 | |
| 49 | # prepare callbacks |
| 50 | history = cbks.History() |
| 51 | callbacks = [cbks.BaseLogger()] + callbacks + [history] |
| 52 | if verbose: |
| 53 | callbacks += [cbks.ProgbarLogger()] |
| 54 | callbacks = cbks.CallbackList(callbacks) |
| 55 | |
| 56 | callbacks._set_params({ |
| 57 | 'nb_epoch': nb_epoch, |
| 58 | 'nb_sample': samples_per_epoch, |
| 59 | 'verbose': verbose, |
| 60 | 'metrics': callback_metrics, |
| 61 | }) |
| 62 | callbacks.on_train_begin() |
| 63 | |
| 64 | while epoch < nb_epoch: |
| 65 | callbacks.on_epoch_begin(epoch) |
| 66 | samples_seen = 0 |
| 67 | batch_index = 0 |
| 68 | while samples_seen < samples_per_epoch: |
| 69 | z, x = next(generator) |
| 70 | # build batch logs |
| 71 | batch_logs = {} |
| 72 | if type(x) is list: |
| 73 | batch_size = len(x[0]) |
| 74 | elif type(x) is dict: |
| 75 | batch_size = len(list(x.values())[0]) |
| 76 | else: |
| 77 | batch_size = len(x) |
| 78 | batch_logs['batch'] = batch_index |
| 79 | batch_logs['size'] = batch_size |
| 80 | callbacks.on_batch_begin(batch_index, batch_logs) |
| 81 | |
| 82 | t1 = time.time() |
| 83 | d_losses = d_train(x, z, counter) |
| 84 | z, x = next(generator) |
| 85 | g_loss, samples, xs = g_train(x, z, counter) |
| 86 | outs = (g_loss, ) + d_losses + (time.time() - t1, ) |
| 87 | counter += 1 |
| 88 | |
| 89 | # save samples |
| 90 | if batch_index % 100 == 0: |
| 91 | join_image = np.zeros_like(np.concatenate([samples[:64], xs[:64]], axis=0)) |
| 92 | for j, (i1, i2) in enumerate(zip(samples[:64], xs[:64])): |
no test coverage detected