Creates a native tf.keras ResNet model. Args: resnet_depth (str): ResNet depth. Options=[50 (default), 101, 152]. resnet_version (str): ResNet version. Options=[v1 (default), v2]. Returns: model (tf.keras.Model): model corresponding to 'resnet_depth' and 'resne
(resnet_depth: str = "50", resnet_version: str = "v1")
| 22 | |
| 23 | |
| 24 | def get_resnet_model(resnet_depth: str = "50", resnet_version: str = "v1") -> tf.keras.Model: |
| 25 | """ |
| 26 | Creates a native tf.keras ResNet model. |
| 27 | |
| 28 | Args: |
| 29 | resnet_depth (str): ResNet depth. Options=[50 (default), 101, 152]. |
| 30 | resnet_version (str): ResNet version. Options=[v1 (default), v2]. |
| 31 | |
| 32 | Returns: |
| 33 | model (tf.keras.Model): model corresponding to 'resnet_depth' and 'resnet_version'. |
| 34 | """ |
| 35 | |
| 36 | shape = ( |
| 37 | _DEFAULT_IMAGE_SIZE["resnet_{}".format(resnet_version)], |
| 38 | _DEFAULT_IMAGE_SIZE["resnet_{}".format(resnet_version)], |
| 39 | _NUM_CHANNELS, |
| 40 | ) |
| 41 | |
| 42 | model_name = "resnet_" + resnet_depth + resnet_version |
| 43 | model = get_tfkeras_model(model_name=model_name, shape=shape) |
| 44 | |
| 45 | return model |