Created
December 17, 2020 12:39
-
-
Save deepak-karkala/c6cd15f96f9fc26d4ac14287f44e8d65 to your computer and use it in GitHub Desktop.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| # Hyperparameter tuning using Tensorflow and Tensor board | |
| # Testing performance with varioud optimizers | |
| HP_OPTIMIZER = hp.HParam('optimizer', hp.Discrete(['adam', 'sgd', 'rmsprop'])) | |
| METRIC_ACCURACY = 'accuracy' | |
| with tf.summary.create_file_writer('logs/hparam_tuning').as_default(): | |
| hp.hparams_config( | |
| hparams=[HP_OPTIMIZER], | |
| metrics=[hp.Metric(METRIC_ACCURACY, display_name='Accuracy')], | |
| ) | |
| # Record accuaracy for each combination of hyperparameters | |
| def run(run_dir, hparams): | |
| with tf.summary.create_file_writer(run_dir).as_default(): | |
| hp.hparams(hparams) # record the values used in this trial | |
| accuracy = train_test_model(hparams) | |
| tf.summary.scalar(METRIC_ACCURACY, accuracy, step=1) | |
| # Fit model with each combination of hyperparameters | |
| def train_test_model(hparams): | |
| model = build_model() | |
| model.fit(train_dataset, epochs=EPOCHS, | |
| steps_per_epoch=STEPS_PER_EPOCH) | |
| _, accuracy = model.evaluate(val_dataset) | |
| return accuracy | |
| # Run model with different combinations of hyperparameters | |
| session_num = 0 | |
| for optimizer in HP_OPTIMIZER.domain.values: | |
| hparams = { | |
| HP_OPTIMIZER: optimizer, | |
| } | |
| run_name = "run-%d" % session_num | |
| print('--- Starting trial: %s' % run_name) | |
| print({h.name: hparams[h] for h in hparams}) | |
| run('logs/hparam_tuning/' + run_name, hparams) | |
| session_num += 1 |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment