MCPcopy Create free account
hub / github.com/tensorflow/models / create_optimizer

Method create_optimizer

official/core/base_task.py:73–111  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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.

Callers 10

_get_pretrain_modelFunction · 0.45
_get_classifier_modelFunction · 0.45
_get_squad_modelFunction · 0.45
mainFunction · 0.45
mainFunction · 0.45
mainFunction · 0.45
create_optimizerFunction · 0.45
create_test_trainerMethod · 0.45

Calls 4

build_optimizerMethod · 0.95
build_learning_rateMethod · 0.95
infoMethod · 0.80
as_dictMethod · 0.45

Tested by 3

create_test_trainerMethod · 0.36