MCPcopy Create free account
hub / github.com/NVIDIA/TensorRT / get_resnet_model

Function get_resnet_model

tools/tensorflow-quantization/examples/resnet/utils.py:24–45  ·  view source on GitHub ↗

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")

Source from the content-addressed store, hash-verified

22
23
24def 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

Calls 1

get_tfkeras_modelFunction · 0.90