MCPcopy Create free account
hub / github.com/NVIDIA/vid2vid / train

Function train

train.py:14–128  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

12from util.visualizer import Visualizer
13
14def train():
15 opt = TrainOptions().parse()
16 if opt.debug:
17 opt.display_freq = 1
18 opt.print_freq = 1
19 opt.nThreads = 1
20
21 ### initialize dataset
22 data_loader = CreateDataLoader(opt)
23 dataset = data_loader.load_data()
24 dataset_size = len(data_loader)
25 print('#training videos = %d' % dataset_size)
26
27 ### initialize models
28 models = create_model(opt)
29 modelG, modelD, flowNet, optimizer_G, optimizer_D, optimizer_D_T = create_optimizer(opt, models)
30
31 ### set parameters
32 n_gpus, tG, tD, tDB, s_scales, t_scales, input_nc, output_nc, \
33 start_epoch, epoch_iter, print_freq, total_steps, iter_path = init_params(opt, modelG, modelD, data_loader)
34 visualizer = Visualizer(opt)
35
36 ### real training starts here
37 for epoch in range(start_epoch, opt.niter + opt.niter_decay + 1):
38 epoch_start_time = time.time()
39 for idx, data in enumerate(dataset, start=epoch_iter):
40 if total_steps % print_freq == 0:
41 iter_start_time = time.time()
42 total_steps += opt.batchSize
43 epoch_iter += opt.batchSize
44
45 # whether to collect output images
46 save_fake = total_steps % opt.display_freq == 0
47 n_frames_total, n_frames_load, t_len = data_loader.dataset.init_data_params(data, n_gpus, tG)
48 fake_B_prev_last, frames_all = data_loader.dataset.init_data(t_scales)
49
50 for i in range(0, n_frames_total, n_frames_load):
51 input_A, input_B, inst_A = data_loader.dataset.prepare_data(data, i, input_nc, output_nc)
52
53 ###################################### Forward Pass ##########################
54 ####### generator
55 fake_B, fake_B_raw, flow, weight, real_A, real_Bp, fake_B_last = modelG(input_A, input_B, inst_A, fake_B_prev_last)
56
57 ####### discriminator
58 ### individual frame discriminator
59 real_B_prev, real_B = real_Bp[:, :-1], real_Bp[:, 1:] # the collection of previous and current real frames
60 flow_ref, conf_ref = flowNet(real_B, real_B_prev) # reference flows and confidences
61 fake_B_prev = modelG.module.compute_fake_B_prev(real_B_prev, fake_B_prev_last, fake_B)
62 fake_B_prev_last = fake_B_last
63
64 losses = modelD(0, reshape([real_B, fake_B, fake_B_raw, real_A, real_B_prev, fake_B_prev, flow, weight, flow_ref, conf_ref]))
65 losses = [ torch.mean(x) if x is not None else 0 for x in losses ]
66 loss_dict = dict(zip(modelD.module.loss_names, losses))
67
68 ### temporal discriminator
69 # get skipped frames for each temporal scale
70 frames_all, frames_skipped = modelD.module.get_all_skipped_frames(frames_all, \
71 real_B, fake_B, flow_ref, conf_ref, t_scales, tD, n_frames_load, i, flowNet)

Callers 1

train.pyFile · 0.70

Calls 15

print_current_errorsMethod · 0.95
plot_current_errorsMethod · 0.95
vis_printMethod · 0.95
TrainOptionsClass · 0.90
CreateDataLoaderFunction · 0.90
create_modelFunction · 0.90
create_optimizerFunction · 0.90
init_paramsFunction · 0.90
VisualizerClass · 0.90
save_modelsFunction · 0.90
update_modelsFunction · 0.90

Tested by

no test coverage detected