MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / _create_optimizer

Method _create_optimizer

modelzoo/dbmtl/train.py:259–294  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

257
258 # define optimizer and generate train_op
259 def _create_optimizer(self):
260 self.global_step = tf.train.get_or_create_global_step()
261 if self.tf or self._optimizer_type == 'adam':
262 dnn_optimizer = tf.train.AdamOptimizer(
263 learning_rate=self._learning_rate,
264 beta1=0.9,
265 beta2=0.999,
266 epsilon=1e-8)
267 elif self._optimizer_type == 'adagrad':
268 dnn_optimizer = tf.train.AdagradOptimizer(
269 learning_rate=self._learning_rate,
270 initial_accumulator_value=0.1,
271 use_locking=False)
272 elif self._optimizer_type == 'adamasync':
273 dnn_optimizer = tf.train.AdamAsyncOptimizer(
274 learning_rate=self._learning_rate,
275 beta1=0.9,
276 beta2=0.999,
277 epsilon=1e-8)
278 elif self._optimizer_type == 'adagraddecay':
279 dnn_optimizer = tf.train.AdagradDecayOptimizer(
280 learning_rate=self._learning_rate,
281 global_step=self.global_step)
282 else:
283 raise ValueError("Optimizer type error.")
284
285 train_ops = []
286 update_ops = tf.get_collection(tf.GraphKeys.UPDATE_OPS)
287 with tf.control_dependencies(update_ops):
288 train_ops.append(
289 dnn_optimizer.minimize(self.loss,
290 var_list=tf.get_collection(
291 tf.GraphKeys.TRAINABLE_VARIABLES,
292 scope='dnn'),
293 global_step=self.global_step))
294 self.train_op = tf.group(*train_ops)
295
296 def _create_metrics(self):
297 self.acc, self.acc_op = tf.metrics.accuracy(labels=self._label,

Callers 1

__init__Method · 0.95

Calls 6

AdamOptimizerMethod · 0.80
get_collectionMethod · 0.45
control_dependenciesMethod · 0.45
appendMethod · 0.45
minimizeMethod · 0.45
groupMethod · 0.45

Tested by

no test coverage detected