()
| 12 | from util.visualizer import Visualizer |
| 13 | |
| 14 | def 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) |
no test coverage detected