(self)
| 94 | |
| 95 | |
| 96 | def setup_optimizers(self): |
| 97 | train_opt = self.opt['train'] |
| 98 | optim_params = [] |
| 99 | for k, v in self.net_g.named_parameters(): |
| 100 | if v.requires_grad: |
| 101 | optim_params.append(v) |
| 102 | |
| 103 | optim_type = train_opt['optim_g'].pop('type') |
| 104 | self.optimizer_g = self.get_optimizer(optim_type, optim_params, **train_opt['optim_g']) |
| 105 | self.optimizers.append(self.optimizer_g) |
| 106 | |
| 107 | optim_params = [] |
| 108 | for k, v in self.net_mask.named_parameters(): |
| 109 | if v.requires_grad: |
| 110 | optim_params.append(v) |
| 111 | else: |
| 112 | logger = get_root_logger() |
| 113 | logger.warning(f"Params {k} will not be optimized.") |
| 114 | |
| 115 | optim_type = train_opt["optim_mask"].pop("type") |
| 116 | self.optimizer_mask = self.get_optimizer( |
| 117 | optim_type, optim_params, **train_opt["optim_mask"] |
| 118 | ) |
| 119 | |
| 120 | self.optimizers.append(self.optimizer_mask) |
| 121 | |
| 122 | def setup_loss_functions(self): |
| 123 | train_opt = self.opt['train'] |
no test coverage detected