(opt, models)
| 84 | return modelG |
| 85 | |
| 86 | def create_optimizer(opt, models): |
| 87 | modelG, modelD, flowNet = models |
| 88 | optimizer_D_T = [] |
| 89 | if opt.fp16: |
| 90 | from apex import amp |
| 91 | for s in range(opt.n_scales_temporal): |
| 92 | optimizer_D_T.append(getattr(modelD, 'optimizer_D_T'+str(s))) |
| 93 | modelG, optimizer_G = amp.initialize(modelG, modelG.optimizer_G, opt_level='O1') |
| 94 | modelD, optimizers_D = amp.initialize(modelD, [modelD.optimizer_D] + optimizer_D_T, opt_level='O1') |
| 95 | optimizer_D, optimizer_D_T = optimizers_D[0], optimizers_D[1:] |
| 96 | modelG, modelD, flownet = wrap_model(opt, modelG, modelD, flowNet) |
| 97 | else: |
| 98 | optimizer_G = modelG.module.optimizer_G |
| 99 | optimizer_D = modelD.module.optimizer_D |
| 100 | for s in range(opt.n_scales_temporal): |
| 101 | optimizer_D_T.append(getattr(modelD.module, 'optimizer_D_T'+str(s))) |
| 102 | return modelG, modelD, flowNet, optimizer_G, optimizer_D, optimizer_D_T |
| 103 | |
| 104 | def init_params(opt, modelG, modelD, data_loader): |
| 105 | iter_path = os.path.join(opt.checkpoints_dir, opt.name, 'iter.txt') |
no test coverage detected