A create optimizer util to be backward compatability with new args.
(task: base_task.Task,
params: config_definitions.ExperimentConfig
)
| 242 | |
| 243 | |
| 244 | def create_optimizer(task: base_task.Task, |
| 245 | params: config_definitions.ExperimentConfig |
| 246 | ) -> tf_keras.optimizers.Optimizer: |
| 247 | """A create optimizer util to be backward compatability with new args.""" |
| 248 | if 'dp_config' in inspect.signature(task.create_optimizer).parameters: |
| 249 | dp_config = None |
| 250 | if hasattr(params.task, 'differential_privacy_config'): |
| 251 | dp_config = params.task.differential_privacy_config |
| 252 | optimizer = task.create_optimizer( |
| 253 | params.trainer.optimizer_config, params.runtime, |
| 254 | dp_config=dp_config) |
| 255 | else: |
| 256 | if hasattr(params.task, 'differential_privacy_config' |
| 257 | ) and params.task.differential_privacy_config is not None: |
| 258 | raise ValueError('Differential privacy config is specified but ' |
| 259 | 'task.create_optimizer api does not accept it.') |
| 260 | optimizer = task.create_optimizer( |
| 261 | params.trainer.optimizer_config, |
| 262 | params.runtime) |
| 263 | return optimizer |
| 264 | |
| 265 | |
| 266 | @gin.configurable |
no test coverage detected