Returns a cloned optimizer with the provided optimizer.config or config.
(optimizer, config=None, worker_name=None)
| 386 | |
| 387 | |
| 388 | def _clone_optimizer(optimizer, config=None, worker_name=None): |
| 389 | """Returns a cloned optimizer with the provided optimizer.config or config.""" |
| 390 | if not isinstance(optimizer, keras_optimizers.Optimizer): |
| 391 | # In the first call to tpu_model(model), Keras may not have wrapped the TF |
| 392 | # optimizer in the TFOptimizer helper, e.g., the given model isn't compiled |
| 393 | # or optimizer isn't set, and later generated tpu_model compiles with a TF |
| 394 | # optimizer. |
| 395 | return optimizer |
| 396 | |
| 397 | if isinstance(optimizer, keras_optimizers.TFOptimizer): |
| 398 | return keras_optimizers.TFOptimizer(optimizer.optimizer) |
| 399 | |
| 400 | if config is None: |
| 401 | config = optimizer.get_config() |
| 402 | logging.info('Cloning %s %s', optimizer.__class__.__name__, config) |
| 403 | with ops.device( |
| 404 | '%s/device:CPU:0' % ('/job:%s' % worker_name if worker_name else '')): |
| 405 | # Explicitly put optimizer parameter variables on TPU worker. |
| 406 | return optimizer.__class__.from_config(config) |
| 407 | |
| 408 | |
| 409 | class TPURewriteContext(object): |
no test coverage detected