MCPcopy Create free account
hub / github.com/albertpumarola/GANimation / Train

Class Train

train.py:10–137  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

8
9
10class Train:
11 def __init__(self):
12 self._opt = TrainOptions().parse()
13 data_loader_train = CustomDatasetDataLoader(self._opt, is_for_train=True)
14 data_loader_test = CustomDatasetDataLoader(self._opt, is_for_train=False)
15
16 self._dataset_train = data_loader_train.load_data()
17 self._dataset_test = data_loader_test.load_data()
18
19 self._dataset_train_size = len(data_loader_train)
20 self._dataset_test_size = len(data_loader_test)
21 print('#train images = %d' % self._dataset_train_size)
22 print('#test images = %d' % self._dataset_test_size)
23
24 self._model = ModelsFactory.get_by_name(self._opt.model, self._opt)
25 self._tb_visualizer = TBVisualizer(self._opt)
26
27 self._train()
28
29 def _train(self):
30 self._total_steps = self._opt.load_epoch * self._dataset_train_size
31 self._iters_per_epoch = self._dataset_train_size / self._opt.batch_size
32 self._last_display_time = None
33 self._last_save_latest_time = None
34 self._last_print_time = time.time()
35
36 for i_epoch in range(self._opt.load_epoch + 1, self._opt.nepochs_no_decay + self._opt.nepochs_decay + 1):
37 epoch_start_time = time.time()
38
39 # train epoch
40 self._train_epoch(i_epoch)
41
42 # save model
43 print('saving the model at the end of epoch %d, iters %d' % (i_epoch, self._total_steps))
44 self._model.save(i_epoch)
45
46 # print epoch info
47 time_epoch = time.time() - epoch_start_time
48 print('End of epoch %d / %d \t Time Taken: %d sec (%d min or %d h)' %
49 (i_epoch, self._opt.nepochs_no_decay + self._opt.nepochs_decay, time_epoch,
50 time_epoch / 60, time_epoch / 3600))
51
52 # update learning rate
53 if i_epoch > self._opt.nepochs_no_decay:
54 self._model.update_learning_rate()
55
56 def _train_epoch(self, i_epoch):
57 epoch_iter = 0
58 self._model.set_train()
59 for i_train_batch, train_batch in enumerate(self._dataset_train):
60 iter_start_time = time.time()
61
62 # display flags
63 do_visuals = self._last_display_time is None or time.time() - self._last_display_time > self._opt.display_freq_s
64 do_print_terminal = time.time() - self._last_print_time > self._opt.print_freq_s or do_visuals
65
66 # train model
67 self._model.set_input(train_batch)

Callers 1

train.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected