Creates an TF optimizer from configurations. Args: optimizer_config: the parameters of the Optimization settings. runtime_config: the parameters of the runtime. dp_config: the parameter of differential privacy. Returns: A tf.optimizers.Optimizer object.
(cls, optimizer_config: OptimizationConfig,
runtime_config: Optional[RuntimeConfig] = None,
dp_config: Optional[DifferentialPrivacyConfig] = None)
| 71 | |
| 72 | @classmethod |
| 73 | def create_optimizer(cls, optimizer_config: OptimizationConfig, |
| 74 | runtime_config: Optional[RuntimeConfig] = None, |
| 75 | dp_config: Optional[DifferentialPrivacyConfig] = None): |
| 76 | """Creates an TF optimizer from configurations. |
| 77 | |
| 78 | Args: |
| 79 | optimizer_config: the parameters of the Optimization settings. |
| 80 | runtime_config: the parameters of the runtime. |
| 81 | dp_config: the parameter of differential privacy. |
| 82 | |
| 83 | Returns: |
| 84 | A tf.optimizers.Optimizer object. |
| 85 | """ |
| 86 | gradient_transformers = None |
| 87 | if dp_config is not None: |
| 88 | logging.info("Adding differential privacy transform with config %s.", |
| 89 | dp_config.as_dict()) |
| 90 | noise_stddev = dp_config.clipping_norm * dp_config.noise_multiplier |
| 91 | gradient_transformers = [ |
| 92 | functools.partial( |
| 93 | ops.clip_l2_norm, l2_norm_clip=dp_config.clipping_norm), |
| 94 | functools.partial( |
| 95 | ops.add_noise, noise_stddev=noise_stddev) |
| 96 | ] |
| 97 | |
| 98 | opt_factory = optimization.OptimizerFactory(optimizer_config) |
| 99 | optimizer = opt_factory.build_optimizer( |
| 100 | opt_factory.build_learning_rate(), |
| 101 | gradient_transformers=gradient_transformers |
| 102 | ) |
| 103 | # Configuring optimizer when loss_scale is set in runtime config. This helps |
| 104 | # avoiding overflow/underflow for float16 computations. |
| 105 | if runtime_config: |
| 106 | optimizer = performance.configure_optimizer( |
| 107 | optimizer, |
| 108 | use_float16=runtime_config.mixed_precision_dtype == "float16", |
| 109 | loss_scale=runtime_config.loss_scale) |
| 110 | |
| 111 | return optimizer |
| 112 | |
| 113 | def initialize(self, model: tf_keras.Model): |
| 114 | """[Optional] A callback function used as CheckpointManager's init_fn. |