Create a Keras model with the given hyperparameters. Args: hparams: A dict mapping hyperparameters in `HPARAMS` to values. seed: A hashable object to be used as a random seed (e.g., to construct dropout layers in the model). Returns: A compiled Keras model.
(hparams, seed)
| 106 | |
| 107 | |
| 108 | def model_fn(hparams, seed): |
| 109 | """Create a Keras model with the given hyperparameters. |
| 110 | |
| 111 | Args: |
| 112 | hparams: A dict mapping hyperparameters in `HPARAMS` to values. |
| 113 | seed: A hashable object to be used as a random seed (e.g., to |
| 114 | construct dropout layers in the model). |
| 115 | |
| 116 | Returns: |
| 117 | A compiled Keras model. |
| 118 | """ |
| 119 | rng = random.Random(seed) |
| 120 | |
| 121 | model = tf.keras.models.Sequential() |
| 122 | model.add(tf.keras.layers.Input(INPUT_SHAPE)) |
| 123 | model.add(tf.keras.layers.Reshape(INPUT_SHAPE + (1,))) # grayscale channel |
| 124 | |
| 125 | # Add convolutional layers. |
| 126 | conv_filters = 8 |
| 127 | for _ in range(hparams[HP_CONV_LAYERS]): |
| 128 | model.add( |
| 129 | tf.keras.layers.Conv2D( |
| 130 | filters=conv_filters, |
| 131 | kernel_size=hparams[HP_CONV_KERNEL_SIZE], |
| 132 | padding="same", |
| 133 | activation="relu", |
| 134 | ) |
| 135 | ) |
| 136 | model.add(tf.keras.layers.MaxPool2D(pool_size=2, padding="same")) |
| 137 | conv_filters *= 2 |
| 138 | |
| 139 | model.add(tf.keras.layers.Flatten()) |
| 140 | model.add( |
| 141 | tf.keras.layers.Dropout( |
| 142 | hparams[HP_DROPOUT], seed=rng.randrange(1 << 32) |
| 143 | ) |
| 144 | ) |
| 145 | |
| 146 | # Add fully connected layers. |
| 147 | dense_neurons = 32 |
| 148 | for _ in range(hparams[HP_DENSE_LAYERS]): |
| 149 | model.add(tf.keras.layers.Dense(dense_neurons, activation="relu")) |
| 150 | dense_neurons *= 2 |
| 151 | |
| 152 | # Add the final output layer. |
| 153 | model.add(tf.keras.layers.Dense(OUTPUT_CLASSES, activation="softmax")) |
| 154 | |
| 155 | model.compile( |
| 156 | loss="sparse_categorical_crossentropy", |
| 157 | optimizer=hparams[HP_OPTIMIZER], |
| 158 | metrics=["accuracy"], |
| 159 | ) |
| 160 | return model |
| 161 | |
| 162 | |
| 163 | def run(data, base_logdir, session_id, hparams): |