r"""Return the scheduler object. Args: opt_opt (obj): Config for the specific optimization module (gen/dis). params (obj): Parameters to be trained by the parameters. Returns: (obj): Optimizer
(opt_opt, params)
| 186 | |
| 187 | |
| 188 | def get_optimizer_for_params(opt_opt, params): |
| 189 | r"""Return the scheduler object. |
| 190 | |
| 191 | Args: |
| 192 | opt_opt (obj): Config for the specific optimization module (gen/dis). |
| 193 | params (obj): Parameters to be trained by the parameters. |
| 194 | |
| 195 | Returns: |
| 196 | (obj): Optimizer |
| 197 | """ |
| 198 | # We will use fuse optimizers by default. |
| 199 | if opt_opt.type == 'adam': |
| 200 | opt = Adam(params, |
| 201 | lr=opt_opt.lr, |
| 202 | betas=(opt_opt.adam_beta1, opt_opt.adam_beta2)) |
| 203 | else: |
| 204 | raise NotImplementedError( |
| 205 | 'Optimizer {} is not yet implemented.'.format(opt_opt.type)) |
| 206 | return opt |
| 207 |