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

Method test_configure_optimizer

official/core/base_trainer_test.py:299–327  ·  view source on GitHub ↗
(self, mixed_precision_dtype, loss_scale)

Source from the content-addressed store, hash-verified

297 loss_scale=[None, 'dynamic', 128, 256],
298 ))
299 def test_configure_optimizer(self, mixed_precision_dtype, loss_scale):
300 config = cfg.ExperimentConfig(
301 runtime=cfg.RuntimeConfig(
302 mixed_precision_dtype=mixed_precision_dtype, loss_scale=loss_scale),
303 trainer=cfg.TrainerConfig(
304 optimizer_config=cfg.OptimizationConfig({
305 'optimizer': {
306 'type': 'sgd'
307 },
308 'learning_rate': {
309 'type': 'constant'
310 },
311 })))
312 trainer = self.create_test_trainer(config)
313 if mixed_precision_dtype == 'float16':
314 self.assertIsInstance(trainer.optimizer,
315 tf_keras.mixed_precision.LossScaleOptimizer)
316 if loss_scale in (None, 'dynamic'):
317 self.assertTrue(trainer.optimizer.dynamic)
318 else:
319 self.assertFalse(trainer.optimizer.dynamic)
320 self.assertEqual(trainer.optimizer.initial_scale, loss_scale)
321 else:
322 self.assertIsInstance(
323 trainer.optimizer,
324 (tf_keras.optimizers.SGD, tf_keras.optimizers.legacy.SGD))
325
326 metrics = trainer.train(tf.convert_to_tensor(5, dtype=tf.int32))
327 self.assertIn('training_loss', metrics)
328
329 def test_export_best_ckpt(self):
330 config = cfg.ExperimentConfig(

Callers

nothing calls this directly

Calls 2

create_test_trainerMethod · 0.95
trainMethod · 0.45

Tested by

no test coverage detected