MCPcopy Create free account
hub / github.com/commaai/research / train_model

Function train_model

train_generative_model.py:35–123  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

33
34
35def 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])):

Callers 1

Calls 2

save_imagesFunction · 0.90
samplerFunction · 0.50

Tested by

no test coverage detected