{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":21154,"databundleVersionId":1243559,"sourceType":"competition"},{"sourceId":629,"sourceType":"modelInstanceVersion","modelInstanceId":496},{"sourceId":630,"sourceType":"modelInstanceVersion","modelInstanceId":497}],"dockerImageVersionId":30381,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Vision Transformer (ViT) with TensorFlow \nStarter code to demonstrate how to use Kaggle Models","metadata":{"papermill":{"duration":0.006046,"end_time":"2023-02-03T17:34:16.514669","exception":false,"start_time":"2023-02-03T17:34:16.508623","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"This notebook is attached to a Kaggle Model (**ViT** B8) and in this notebook we will demonstrate how to work with two different variations of this same model: (A) ****ViT** B8 **Classification**; and (B) **ViT** B8**Feature Vector**. \n\nThe first model variation (the classification model) we will use directly for inference, and the second model variation (the feature vector) we will fine-tune against a new dataset in order to train it to perform a new task (transfer learning).\n\nVision Transformer (ViT) is a computer vision model that can be used to identify objects within images. \n\nFor more details, see [An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale\n](https://arxiv.org/pdf/2010.11929.pdf).\n\n![](https://i.imgur.com/521lIWZ.png)\n\nFigure 1 from [An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale\n](https://arxiv.org/pdf/2010.11929.pdf)\n","metadata":{"papermill":{"duration":0.004634,"end_time":"2023-02-03T17:34:16.524398","exception":false,"start_time":"2023-02-03T17:34:16.519764","status":"completed"},"tags":[]}},{"cell_type":"code","source":"\nfrom __future__ import absolute_import, division, print_function\n\nimport numpy as np \nimport pandas as pd \nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))","metadata":{"papermill":{"duration":0.034985,"end_time":"2023-02-03T17:34:16.564237","exception":false,"start_time":"2023-02-03T17:34:16.529252","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-28T11:57:35.308397Z","iopub.execute_input":"2024-03-28T11:57:35.308854Z","iopub.status.idle":"2024-03-28T11:57:35.339688Z","shell.execute_reply.started":"2024-03-28T11:57:35.308820Z","shell.execute_reply":"2024-03-28T11:57:35.338682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!python --version","metadata":{"execution":{"iopub.status.busy":"2024-03-28T11:57:35.342078Z","iopub.execute_input":"2024-03-28T11:57:35.342535Z","iopub.status.idle":"2024-03-28T11:57:36.360208Z","shell.execute_reply.started":"2024-03-28T11:57:35.342498Z","shell.execute_reply":"2024-03-28T11:57:36.359185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Vision Transformer (ViT) in Kaggle Models","metadata":{"papermill":{"duration":0.004846,"end_time":"2023-02-03T17:34:16.575002","exception":false,"start_time":"2023-02-03T17:34:16.570156","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"With the [Kaggle Models](https://www.kaggle.com/models) product, in addition to being able to compare many different  models to each other, you can also compare many different variations of each model to each other, until you find the variation that is most suitable for your specific purposes.\n\nThe main types of variations of Kaggle Models include: (1) small variations to the model architecture (i.e. small, medium, or large model variations -- e.g. EffNetV2-XL vs EffNetV2-L; (2) major variations to the model weights (i.e. trained for different amounts of time, or trained against different data -- e.g. EffNetV2-B0 vs EffNetV2-B1, or EffNetV2-Imagenet1k vs EffNetV2-Imagenet21k, or similar); and (3) small variations to the model architecture to make the model better for some specific purpose (i.e. to be used directly for inference or to be modified for some new purpose -- e.g. EffNetV2 classification vs EffNetV2 feature vector). \n\n\nFor the classification example, the model path can be either: (A) https://kaggle.com/models/kaggle/vision-transformer/frameworks/TensorFlow2/variations/vit-b8-fe/versions/1; or (B) '/kaggle/input/vision-transformer/tensorflow2/vit-b8-classification/1'.\n\n\nFor the feature vector example, the model path can be either: (A) https://kaggle.com/models/kaggle/vision-transformer/frameworks/TensorFlow2/variations/vit-b8-fe/versions/1; or (B) '/kaggle/input/vision-transformer/tensorflow2/vit-b8-fe/1'.\n","metadata":{"papermill":{"duration":0.004796,"end_time":"2023-02-03T17:34:16.584762","exception":false,"start_time":"2023-02-03T17:34:16.579966","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"We will be using TensorFlow 2.0 for both examples, using code that was adapted from the following tutorials:  \n - https://www.tensorflow.org/hub/tutorials/image_classification\n - https://www.tensorflow.org/hub/tutorials/tf2_image_retraining","metadata":{"papermill":{"duration":0.004776,"end_time":"2023-02-03T17:34:16.594518","exception":false,"start_time":"2023-02-03T17:34:16.589742","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# Using Kaggle Models for Inference","metadata":{"papermill":{"duration":0.004772,"end_time":"2023-02-03T17:34:16.604268","exception":false,"start_time":"2023-02-03T17:34:16.599496","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import tensorflow as tf\nimport tensorflow_hub as hub\n\nimport requests\nfrom PIL import Image\nfrom io import BytesIO\n\nimport matplotlib.pyplot as plt\nimport numpy as np\n\n\nimport math, re\nimport numpy as np\n\nfrom IPython.display import Image as IImage, display\n\nimport PIL\nfrom PIL import Image\nimport random\nimport requests\n\nprint(\"Tensorflow version \" + tf.__version__)","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"id":"N8H5ufxkc2mk","papermill":{"duration":5.123187,"end_time":"2023-02-03T17:34:21.732354","exception":false,"start_time":"2023-02-03T17:34:16.609167","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-28T11:57:36.361823Z","iopub.execute_input":"2024-03-28T11:57:36.362301Z","iopub.status.idle":"2024-03-28T11:57:36.370579Z","shell.execute_reply.started":"2024-03-28T11:57:36.362265Z","shell.execute_reply":"2024-03-28T11:57:36.369670Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver() \n    print('Running on TPU ', tpu.master())\nexcept ValueError:\n    tpu = None\n\nif tpu:\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.experimental.TPUStrategy(tpu)\nelse:\n    strategy = tf.distribute.get_strategy() \n\nprint(\"REPLICAS: \", strategy.num_replicas_in_sync)","metadata":{"execution":{"iopub.status.busy":"2024-03-28T11:57:36.371809Z","iopub.execute_input":"2024-03-28T11:57:36.372104Z","iopub.status.idle":"2024-03-28T11:57:36.382963Z","shell.execute_reply.started":"2024-03-28T11:57:36.372077Z","shell.execute_reply":"2024-03-28T11:57:36.381896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from kaggle_datasets import KaggleDatasets\n\nGCS_DS_PATH = '/kaggle/input/tpu-getting-started'\nprint(GCS_DS_PATH) # what do gcs paths look like?","metadata":{"execution":{"iopub.status.busy":"2024-03-28T11:57:36.386101Z","iopub.execute_input":"2024-03-28T11:57:36.386412Z","iopub.status.idle":"2024-03-28T11:57:36.394046Z","shell.execute_reply.started":"2024-03-28T11:57:36.386387Z","shell.execute_reply":"2024-03-28T11:57:36.393171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMAGE_SIZE = [224, 224]\nGCS_PATH = GCS_DS_PATH + '/tfrecords-jpeg-224x224'\nAUTO = tf.data.experimental.AUTOTUNE\n\nIMAGE_SIZE = [224,224] # at this size, a GPU will run out of memory. Use the TPU\nEPOCHS = 20\nBATCH_SIZE = 16 * strategy.num_replicas_in_sync#用于设置批量大小，并且根据 TensorFlow 中的分布式策略 (strategy.num_replicas_in_sync) 进行了调整\n\nNUM_TRAINING_IMAGES = 12753\nNUM_TEST_IMAGES = 7382\nSTEPS_PER_EPOCH = NUM_TRAINING_IMAGES // BATCH_SIZE\n\nTRAINING_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/train/*.tfrec')\nVALIDATION_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/val/*.tfrec')\nTEST_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/test/*.tfrec') \n\nCLASSES = ['pink primrose',    'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea',     'wild geranium',     'tiger lily',           'moon orchid',              'bird of paradise', 'monkshood',        'globe thistle',         # 00 - 09\n           'snapdragon',       \"colt's foot\",               'king protea',      'spear thistle', 'yellow iris',       'globe-flower',         'purple coneflower',        'peruvian lily',    'balloon flower',   'giant white arum lily', # 10 - 19\n           'fire lily',        'pincushion flower',         'fritillary',       'red ginger',    'grape hyacinth',    'corn poppy',           'prince of wales feathers', 'stemless gentian', 'artichoke',        'sweet william',         # 20 - 29\n           'carnation',        'garden phlox',              'love in the mist', 'cosmos',        'alpine sea holly',  'ruby-lipped cattleya', 'cape flower',              'great masterwort', 'siam tulip',       'lenten rose',           # 30 - 39\n           'barberton daisy',  'daffodil',                  'sword lily',       'poinsettia',    'bolero deep blue',  'wallflower',           'marigold',                 'buttercup',        'daisy',            'common dandelion',      # 40 - 49\n           'petunia',          'wild pansy',                'primula',          'sunflower',     'lilac hibiscus',    'bishop of llandaff',   'gaura',                    'geranium',         'orange dahlia',    'pink-yellow dahlia',    # 50 - 59\n           'cautleya spicata', 'japanese anemone',          'black-eyed susan', 'silverbush',    'californian poppy', 'osteospermum',         'spring crocus',            'iris',             'windflower',       'tree poppy',            # 60 - 69\n           'gazania',          'azalea',                    'water lily',       'rose',          'thorn apple',       'morning glory',        'passion flower',           'lotus',            'toad lily',        'anthurium',             # 70 - 79\n           'frangipani',       'clematis',                  'hibiscus',         'columbine',     'desert-rose',       'tree mallow',          'magnolia',                 'cyclamen ',        'watercress',       'canna lily',            # 80 - 89\n           'hippeastrum ',     'bee balm',                  'pink quill',       'foxglove',      'bougainvillea',     'camellia',             'mallow',                   'mexican petunia',  'bromelia',         'blanket flower',        # 90 - 99\n           'trumpet creeper',  'blackberry lily',           'common tulip',     'wild rose']                                                                                                                                               # 100 - 102\n\n\ndef decode_image(image_data):\n    image = tf.image.decode_jpeg(image_data, channels=3)\n    image = tf.cast(image, tf.float32) / 255.0  # convert image to floats in [0, 1] range\n    image = tf.reshape(image, [*IMAGE_SIZE, 3]) # explicit size needed for TPU\n    return image\n\ndef read_labeled_tfrecord(example):\n    LABELED_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string), # tf.string means bytestring\n        \"class\": tf.io.FixedLenFeature([], tf.int64),  # shape [] means single element\n    }\n    example = tf.io.parse_single_example(example, LABELED_TFREC_FORMAT)\n    image = decode_image(example['image'])\n    label = tf.cast(example['class'], tf.int32)\n    return image, label # returns a dataset of (image, label) pairs\n\ndef read_unlabeled_tfrecord(example):\n    UNLABELED_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string), # tf.string means bytestring\n        \"id\": tf.io.FixedLenFeature([], tf.string),  # shape [] means single element\n        # class is missing, this competitions's challenge is to predict flower classes for the test dataset\n    }\n    example = tf.io.parse_single_example(example, UNLABELED_TFREC_FORMAT)\n    image = decode_image(example['image'])\n    idnum = example['id']\n    return image, idnum # returns a dataset of image(s)\n\ndef load_dataset(filenames, labeled=True, ordered=False):\n    # Read from TFRecords. For optimal performance, reading from multiple files at once and\n    # disregarding data order. Order does not matter since we will be shuffling the data anyway.\n\n    ignore_order = tf.data.Options()\n    if not ordered:\n        ignore_order.experimental_deterministic = False # disable order, increase speed\n\n    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads=AUTO) # automatically interleaves reads from multiple files\n    dataset = dataset.with_options(ignore_order) # uses data as soon as it streams in, rather than in its original order\n    dataset = dataset.map(read_labeled_tfrecord if labeled else read_unlabeled_tfrecord, num_parallel_calls=AUTO)\n    # returns a dataset of (image, label) pairs if labeled=True or (image, id) pairs if labeled=False\n    return dataset","metadata":{"execution":{"iopub.status.busy":"2024-03-28T11:57:36.396079Z","iopub.execute_input":"2024-03-28T11:57:36.396455Z","iopub.status.idle":"2024-03-28T11:57:36.450050Z","shell.execute_reply.started":"2024-03-28T11:57:36.396420Z","shell.execute_reply":"2024-03-28T11:57:36.449315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def random_flip_left_right(image):\n    return tf.image.random_flip_left_right(image)\ndef random_contrast(image, minval=0.6, maxval=1.4):\n    r = tf.random.uniform([], minval=minval, maxval=maxval)\n    image = tf.image.adjust_contrast(image, contrast_factor=r)\n    return tf.cast(image, tf.uint8)\n# Randomly change an image brightness\ndef random_brightness(image, minval=0., maxval=.2):\n    r = tf.random.uniform([], minval=minval, maxval=maxval)\n    image = tf.image.adjust_brightness(image, delta=r)\n    return tf.cast(image, tf.uint8)\n# Randomly change an image saturation\ndef random_saturation(image, minval=0.4, maxval=2.):\n    r = tf.random.uniform((), minval=minval, maxval=maxval)\n    image = tf.image.adjust_saturation(image, saturation_factor=r)\n    return tf.cast(image, tf.uint8)\n# Randomly change an image hue.\ndef random_hue(image, minval=-0.04, maxval=0.08):\n    r = tf.random.uniform((), minval=minval, maxval=maxval)\n    image = tf.image.adjust_hue(image, delta=r)\n    return tf.cast(image, tf.uint8)\n# Distort an image by cropping it with a different aspect ratio.\ndef distorted_random_crop(image,\n                min_object_covered=0.1,\n                aspect_ratio_range=(3./4., 4./3.),\n                area_range=(0.06, 1.0),\n                max_attempts=100,\n                scope=None):\n\n    cropbox = tf.constant([0.0, 0.0, 1.0, 1.0], dtype=tf.float32, shape=[1, 1, 4])\n    sample_distorted_bounding_box = tf.image.sample_distorted_bounding_box(\n        tf.shape(image),\n        bounding_boxes=cropbox,\n        min_object_covered=min_object_covered,\n        aspect_ratio_range=aspect_ratio_range,\n        area_range=area_range,\n        max_attempts=max_attempts,\n        use_image_if_no_bounding_boxes=True)\n    bbox_begin, bbox_size, distort_bbox = sample_distorted_bounding_box\n\n    # Crop the image to the specified bounding box.\n    cropped_image = tf.slice(image, bbox_begin, bbox_size)\n    return cropped_image\n# Apply all transformations to an image.\n# That is a common image augmentation technique for image datasets, such as ImageNet.\ndef transform_image(image):\n    #image = distorted_random_crop(image)\n    image = random_flip_left_right(image)\n    image = random_contrast(image)\n    image = random_brightness(image)\n    #image = random_hue(image)\n    image = random_saturation(image)\n    return image\n# Resize transformed image to a 224x224px square image, ready for training.\ndef resize_image(image):\n    image = tf.image.resize(image, size=(224, 224), preserve_aspect_ratio=False)\n    image = tf.cast(image, tf.uint8)\n    return image","metadata":{"execution":{"iopub.status.busy":"2024-03-28T11:57:36.451529Z","iopub.execute_input":"2024-03-28T11:57:36.451805Z","iopub.status.idle":"2024-03-28T11:57:36.469342Z","shell.execute_reply.started":"2024-03-28T11:57:36.451780Z","shell.execute_reply":"2024-03-28T11:57:36.468078Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def data_augment(image, label):\n    # Thanks to the dataset.prefetch(AUTO)\n    # statement in the next function (below), this happens essentially\n    # for free on TPU. Data pipeline code is executed on the \"CPU\"\n    # part of the TPU while the TPU itself is computing gradients.\n    # Display fully pre-processed image.\n    #image = np.array(image)\n    transformed_img = transform_image(image)\n    \n    return image, label   \n\ndef get_training_dataset():\n    dataset = load_dataset(TRAINING_FILENAMES, labeled=True)\n    dataset = dataset.map(data_augment, num_parallel_calls=AUTO)\n    dataset = dataset.repeat() # the training dataset must repeat for several epochs\n    dataset = dataset.shuffle(2048)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(AUTO) # prefetch next batch while training (autotune prefetch buffer size)\n    return dataset\n\ndef get_validation_dataset(ordered=False):\n    dataset = load_dataset(VALIDATION_FILENAMES, labeled=True, ordered=ordered)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.cache()\n    dataset = dataset.prefetch(AUTO)\n    return dataset\n\ndef get_test_dataset(ordered=False):\n    dataset = load_dataset(TEST_FILENAMES, labeled=False, ordered=ordered)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(AUTO)\n    return dataset\n\ndef count_data_items(filenames):\n    # the number of data items is written in the name of the .tfrec\n    # files, i.e. flowers00-230.tfrec = 230 data items\n    n = [int(re.compile(r\"-([0-9]*)\\.\").search(filename).group(1)) for filename in filenames]\n    return np.sum(n)\n\nNUM_TRAINING_IMAGES = count_data_items(TRAINING_FILENAMES)\nNUM_VALIDATION_IMAGES = count_data_items(VALIDATION_FILENAMES)\nNUM_TEST_IMAGES = count_data_items(TEST_FILENAMES)\nprint('Dataset: {} training images, {} validation images, {} unlabeled test images'.format(NUM_TRAINING_IMAGES, NUM_VALIDATION_IMAGES, NUM_TEST_IMAGES))\ntraining_dataset = get_training_dataset()\nvalidation_dataset = get_validation_dataset()\ntest_dataset = get_test_dataset()","metadata":{"execution":{"iopub.status.busy":"2024-03-28T11:57:36.470959Z","iopub.execute_input":"2024-03-28T11:57:36.471318Z","iopub.status.idle":"2024-03-28T11:57:36.887599Z","shell.execute_reply.started":"2024-03-28T11:57:36.471283Z","shell.execute_reply":"2024-03-28T11:57:36.886614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# Define the batch size. This will be 16 with TPU off and 128 (=16*8) with TPU on\nBATCH_SIZE = 16 * strategy.num_replicas_in_sync\n\nds_train = get_training_dataset()\nds_valid = get_validation_dataset()\nds_test = get_test_dataset()\n\nprint(\"Training:\", ds_train)\nprint (\"Validation:\", ds_valid)\nprint(\"Test:\", ds_test)","metadata":{"execution":{"iopub.status.busy":"2024-03-28T11:57:36.889082Z","iopub.execute_input":"2024-03-28T11:57:36.889418Z","iopub.status.idle":"2024-03-28T11:57:36.986400Z","shell.execute_reply.started":"2024-03-28T11:57:36.889388Z","shell.execute_reply":"2024-03-28T11:57:36.985439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.set_printoptions(threshold=15, linewidth=80)\n\nprint(\"Training data shapes:\")\nfor image, label in ds_train.take(3):\n    print(image.numpy().shape, label.numpy().shape)\n#print(image.numpy())\nprint(\"Training data label examples:\", label.numpy())","metadata":{"execution":{"iopub.status.busy":"2024-03-28T11:57:36.987699Z","iopub.execute_input":"2024-03-28T11:57:36.988361Z","iopub.status.idle":"2024-03-28T11:57:38.080577Z","shell.execute_reply.started":"2024-03-28T11:57:36.988305Z","shell.execute_reply":"2024-03-28T11:57:38.079613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nprint(\"Test data shapes:\")\nfor image, idnum in ds_test.take(3):\n    print(image.numpy().shape, idnum.numpy().shape)\nprint(\"Test data IDs:\", idnum.numpy().astype('U')) # U=unicode string","metadata":{"execution":{"iopub.status.busy":"2024-03-28T11:57:38.081719Z","iopub.execute_input":"2024-03-28T11:57:38.082008Z","iopub.status.idle":"2024-03-28T11:57:38.152367Z","shell.execute_reply.started":"2024-03-28T11:57:38.081982Z","shell.execute_reply":"2024-03-28T11:57:38.151199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from matplotlib import pyplot as plt\n\ndef batch_to_numpy_images_and_labels(data):\n    images, labels = data\n    numpy_images = images.numpy()\n    numpy_labels = labels.numpy()\n    if numpy_labels.dtype == object: # binary string in this case,\n                                     # these are image ID strings\n        numpy_labels = [None for _ in enumerate(numpy_images)]\n    # If no labels, only image IDs, return None for labels (this is\n    # the case for test data)\n    return numpy_images, numpy_labels\n\ndef title_from_label_and_target(label, correct_label):\n    if correct_label is None:\n        return CLASSES[label], True\n    correct = (label == correct_label)\n    return \"{} [{}{}{}]\".format(CLASSES[label], 'OK' if correct else 'NO', u\"\\u2192\" if not correct else '',\n                                CLASSES[correct_label] if not correct else ''), correct\n\ndef display_one_flower(image, title, subplot, red=False, titlesize=16):\n    plt.subplot(*subplot)\n    plt.axis('off')\n    plt.imshow(image)\n    if len(title) > 0:\n        plt.title(title, fontsize=int(titlesize) if not red else int(titlesize/1.2), color='red' if red else 'black', fontdict={'verticalalignment':'center'}, pad=int(titlesize/1.5))\n    return (subplot[0], subplot[1], subplot[2]+1)\n    \ndef display_batch_of_images(databatch, predictions=None):\n    \"\"\"This will work with:\n    display_batch_of_images(images)\n    display_batch_of_images(images, predictions)\n    display_batch_of_images((images, labels))\n    display_batch_of_images((images, labels), predictions)\n    \"\"\"\n    # data\n    images, labels = batch_to_numpy_images_and_labels(databatch)\n    if labels is None:\n        labels = [None for _ in enumerate(images)]\n        \n    # auto-squaring: this will drop data that does not fit into square\n    # or square-ish rectangle\n    rows = int(math.sqrt(len(images)))\n    cols = len(images)//rows\n        \n    # size and spacing\n    FIGSIZE = 13.0\n    SPACING = 0.1\n    subplot=(rows,cols,1)\n    if rows < cols:\n        plt.figure(figsize=(FIGSIZE,FIGSIZE/cols*rows))\n    else:\n        plt.figure(figsize=(FIGSIZE/rows*cols,FIGSIZE))\n    \n    # display\n    for i, (image, label) in enumerate(zip(images[:rows*cols], labels[:rows*cols])):\n        title = '' if label is None else CLASSES[label]\n        correct = True\n        if predictions is not None:\n            title, correct = title_from_label_and_target(predictions[i], label)\n        dynamic_titlesize = FIGSIZE*SPACING/max(rows,cols)*40+3 # magic formula tested to work from 1x1 to 10x10 images\n        subplot = display_one_flower(image, title, subplot, not correct, titlesize=dynamic_titlesize)\n    \n    #layout\n    plt.tight_layout()\n    if label is None and predictions is None:\n        plt.subplots_adjust(wspace=0, hspace=0)\n    else:\n        plt.subplots_adjust(wspace=SPACING, hspace=SPACING)\n    plt.show()\n\n\ndef display_training_curves(training, validation, title, subplot):\n    if subplot%10==1: # set up the subplots on the first call\n        plt.subplots(figsize=(10,10), facecolor='#F0F0F0')\n        plt.tight_layout()\n    ax = plt.subplot(subplot)\n    ax.set_facecolor('#F8F8F8')\n    ax.plot(training)\n    ax.plot(validation)\n    ax.set_title('model '+ title)\n    ax.set_ylabel(title)\n    #ax.set_ylim(0.28,1.05)\n    ax.set_xlabel('epoch')\n    ax.legend(['train', 'valid.'])","metadata":{"execution":{"iopub.status.busy":"2024-03-28T11:57:38.154134Z","iopub.execute_input":"2024-03-28T11:57:38.154569Z","iopub.status.idle":"2024-03-28T11:57:38.176883Z","shell.execute_reply.started":"2024-03-28T11:57:38.154529Z","shell.execute_reply":"2024-03-28T11:57:38.175920Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#@title Helper functions for loading image (hidden)\n\noriginal_image_cache = {}\n\ndef preprocess_image(image):\n    image = np.array(image)\n  # reshape into shape [batch_size, height, width, num_channels]\n    img_reshaped = tf.reshape(image, [1, image.shape[0], image.shape[1], image.shape[2]])\n  # Use `convert_image_dtype` to convert to floats in the [0,1] range.\n    image = tf.image.convert_image_dtype(img_reshaped, tf.float32)\n    return image\n\n\ndef load_image(image, image_size=256, dynamic_size=False, max_dynamic_size=512):\n \n    img = preprocess_image(image)\n    original_image_cache[image] = img\n  # Load and convert to float32 numpy array, add batch dimension, and normalize to range [0, 1].\n    img_raw = img\n    if tf.reduce_max(img) > 1.0:\n        img = img / 255.\n    if len(img.shape) == 3:\n        img = tf.stack([img, img, img], axis=-1)\n    if not dynamic_size:\n        img = tf.image.resize_with_pad(img, image_size, image_size)\n    elif img.shape[1] > max_dynamic_size or img.shape[2] > max_dynamic_size:\n        img = tf.image.resize_with_pad(img, max_dynamic_size, max_dynamic_size)\n    return img, img_raw\n\ndef show_image(image, title=''):\n    image_size = image.shape[1]\n    w = (image_size * 6) // 320\n    plt.figure(figsize=(w, w))\n    plt.imshow(image[0], aspect='equal')\n    plt.axis('off')\n    plt.title(title)\n    plt.show()\n\nimage_size = 224\ndynamic_size = False\nmax_dynamic_size = 512","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"cellView":"form","id":"oKvj6lY6kZx8","papermill":{"duration":0.021359,"end_time":"2023-02-03T17:34:21.759489","exception":false,"start_time":"2023-02-03T17:34:21.73813","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-28T11:57:38.178151Z","iopub.execute_input":"2024-03-28T11:57:38.178546Z","iopub.status.idle":"2024-03-28T11:57:38.191352Z","shell.execute_reply.started":"2024-03-28T11:57:38.178517Z","shell.execute_reply":"2024-03-28T11:57:38.190342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Select an Image Classification model","metadata":{"papermill":{"duration":0.005204,"end_time":"2023-02-03T17:34:21.769909","exception":false,"start_time":"2023-02-03T17:34:21.764705","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Select an Image Classification model\n\nmodel_name = \"vit-b8\"\n\nmodel_handle_map = {\n#   \"vit-b8\": \"/kaggle/input/vision-transformer/tensorflow2/vit-b8-classification/1\",\n#   \"evit-b8\": \"/kaggle/input/vision-transformer/tensorflow2/vit-b8-classification/1\",\n#   \"vit-b8\": \"/kaggle/input/efficientnet-v2/tensorflow2/imagenet1k-b0-classification/versions/1\",\n#   \"vit-b8\": \"https://kaggle.com/models/kaggle/vision-transformer/frameworks/TensorFlow2/variations/vit-b8-classification/\",\n    \"vit-b8\": \"https://www.kaggle.com/models/spsayakpaul/vision-transformer/frameworks/TensorFlow2/variations/vit-b16-classification/versions/1\",\n#   \"vit-b8\": \"https://tfhub.dev/sayakpaul/vit_b8_classification/1\",\n}\n\n\nmodel_image_size_map = {\n  \"vit-b8\": 224,\n}\n\nmodel_handle = model_handle_map[model_name]\n\nprint(f\"Selected model: {model_name} : {model_handle}\")","metadata":{"id":"iQ3aamrBfs-c","papermill":{"duration":0.0146,"end_time":"2023-02-03T17:34:21.789682","exception":false,"start_time":"2023-02-03T17:34:21.775082","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-28T11:57:38.196236Z","iopub.execute_input":"2024-03-28T11:57:38.197038Z","iopub.status.idle":"2024-03-28T11:57:38.204652Z","shell.execute_reply.started":"2024-03-28T11:57:38.197008Z","shell.execute_reply":"2024-03-28T11:57:38.203804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"max_dynamic_size = 512\nif model_name in model_image_size_map:\n    image_size = model_image_size_map[model_name]\n    dynamic_size = False\n    print(f\"Images will be converted to {image_size}x{image_size}\")\nelse:\n    dynamic_size = True\n    print(f\"Images will be capped to a max size of {max_dynamic_size}x{max_dynamic_size}\")\n\n#labels_file = \"https://storage.googleapis.com/download.tensorflow.org/data/ImageNetLabels.txt\"\n\n#download labels and creates a maps\ndownloaded_file = '/kaggle/input/tpu-getting-started/sample_submission.csv'\n\nclasses = []\n\nwith open(downloaded_file) as f:\n    labels = f.readlines()\n    classes = [l.strip() for l in labels]","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"papermill":{"duration":0.133836,"end_time":"2023-02-03T17:34:21.928637","exception":false,"start_time":"2023-02-03T17:34:21.794801","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-28T11:57:38.205701Z","iopub.execute_input":"2024-03-28T11:57:38.206013Z","iopub.status.idle":"2024-03-28T11:57:38.221195Z","shell.execute_reply.started":"2024-03-28T11:57:38.205986Z","shell.execute_reply":"2024-03-28T11:57:38.220284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Select an Input Image","metadata":{"papermill":{"duration":0.005205,"end_time":"2023-02-03T17:34:21.940223","exception":false,"start_time":"2023-02-03T17:34:21.935018","status":"completed"},"tags":[]}},{"cell_type":"code","source":"classifier = hub.load(model_handle)\n\ninput_shape = image.shape\nwarmup_input = tf.random.uniform(input_shape, 0, 1.0)\n%time warmup_logits = classifier(warmup_input).numpy()","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"id":"LRAccT3UhRga","papermill":{"duration":11.773977,"end_time":"2023-02-03T17:34:37.396201","exception":false,"start_time":"2023-02-03T17:34:25.622224","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-28T11:57:38.222356Z","iopub.execute_input":"2024-03-28T11:57:38.222640Z","iopub.status.idle":"2024-03-28T11:57:44.875811Z","shell.execute_reply.started":"2024-03-28T11:57:38.222614Z","shell.execute_reply":"2024-03-28T11:57:44.874771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Everything is ready for inference. Here you can see the top 5 results from the model for the selected image.","metadata":{"id":"e7vkdUqpBkfE","papermill":{"duration":0.006461,"end_time":"2023-02-03T17:34:37.409513","exception":false,"start_time":"2023-02-03T17:34:37.403052","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Run model on image\n%time probabilities = tf.nn.softmax(classifier(image)).numpy()\n\ntop_5 = tf.argsort(probabilities, axis=-1, direction=\"DESCENDING\")[0][:5].numpy()\nnp_classes = np.array(classes)\n\n# Some models include an additional 'background' class in the predictions, so\n# we must account for this when reading the class labels.\nincludes_background_class = probabilities.shape[1] == 1001\n\nfor i, item in enumerate(top_5):\n    class_index = item if includes_background_class else item + 1\n    line = f'({i+1}) {class_index:4} - {classes[class_index]}: {probabilities[0][top_5][i]}'\n    print(line)\n\nshow_image(image, '')","metadata":{"id":"I0QNHg3bk-G1","papermill":{"duration":0.164086,"end_time":"2023-02-03T17:34:37.580062","exception":false,"start_time":"2023-02-03T17:34:37.415976","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-28T11:57:44.877157Z","iopub.execute_input":"2024-03-28T11:57:44.877457Z","iopub.status.idle":"2024-03-28T11:57:45.232443Z","shell.execute_reply.started":"2024-03-28T11:57:44.877430Z","shell.execute_reply":"2024-03-28T11:57:45.231555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Using Kaggle Models for Transfer Learning","metadata":{"papermill":{"duration":0.008788,"end_time":"2023-02-03T17:34:37.608075","exception":false,"start_time":"2023-02-03T17:34:37.599287","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"Select a model","metadata":{"papermill":{"duration":0.007291,"end_time":"2023-02-03T17:34:37.622886","exception":false,"start_time":"2023-02-03T17:34:37.615595","status":"completed"},"tags":[]}},{"cell_type":"code","source":"model_name = \"vit-b8\"\n\nmodel_handle_map = {\n#  \"vit-b8\": \"/kaggle/input/vision-transformer/tensorflow2/vit-b8-fe/1\",\n#   \"evit-b8\": \"/kaggle/input/vision-transformer/tensorflow2/vit-b8-fe/1\",\n#   \"vit-b8\": \"/kaggle/input/vision-transformer/tensorflow2/vit-b8-fe/versions/1\",\n#   \"vit-b8\": \"https://kaggle.com/models/kaggle/vision-transformer/frameworks/TensorFlow2/variations/vit-b8-fe/\",\n   \"vit-b8\": \"https://www.kaggle.com/models/spsayakpaul/vision-transformer/frameworks/TensorFlow2/variations/vit-b16-classification/versions/1\",\n#   \"vit-b8\": \"https://tfhub.dev/sayakpaul/vit_b8_fe/1\",\n}\n\nmodel_image_size_map = {\n  \"vit-b8\": 224,\n}\n\nmodel_handle = model_handle_map.get(model_name)\npixels = model_image_size_map.get(model_name, 224)\n\nprint(f\"Selected model: {model_name} : {model_handle}\")\n\nIMAGE_SIZE = (pixels, pixels)\nprint(f\"Input size {IMAGE_SIZE}\")\n\nBATCH_SIZE = 16#@param {type:\"integer\"}","metadata":{"papermill":{"duration":0.017714,"end_time":"2023-02-03T17:34:37.648195","exception":false,"start_time":"2023-02-03T17:34:37.630481","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-28T11:57:45.233757Z","iopub.execute_input":"2024-03-28T11:57:45.234110Z","iopub.status.idle":"2024-03-28T11:57:45.241508Z","shell.execute_reply.started":"2024-03-28T11:57:45.234079Z","shell.execute_reply":"2024-03-28T11:57:45.240550Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Select a dataset to fine-tune the model against","metadata":{"papermill":{"duration":0.007561,"end_time":"2023-02-03T17:34:37.665457","exception":false,"start_time":"2023-02-03T17:34:37.657896","status":"completed"},"tags":[]}},{"cell_type":"code","source":"\n\nnormalization_layer = tf.keras.layers.Rescaling(1. / 255)\npreprocessing_model = tf.keras.Sequential([normalization_layer])\ntrain_ds = training_dataset.map(lambda images, labels:\n                        (preprocessing_model(images), labels))\n\nval_ds = validation_dataset.map(lambda images, labels:\n                    (normalization_layer(images), labels))","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"papermill":{"duration":8.82833,"end_time":"2023-02-03T17:34:46.501641","exception":false,"start_time":"2023-02-03T17:34:37.673311","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-28T11:57:45.242762Z","iopub.execute_input":"2024-03-28T11:57:45.243044Z","iopub.status.idle":"2024-03-28T11:57:45.322662Z","shell.execute_reply.started":"2024-03-28T11:57:45.243017Z","shell.execute_reply":"2024-03-28T11:57:45.319094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Defining the model.\n\nAll it takes is to put a linear classifier on top of the `feature_extractor_layer` with the Hub module.\n\nFor speed, we start out with a non-trainable `feature_extractor_layer`, but you can also enable fine-tuning for greater accuracy.","metadata":{"papermill":{"duration":0.01354,"end_time":"2023-02-03T17:34:46.530896","exception":false,"start_time":"2023-02-03T17:34:46.517356","status":"completed"},"tags":[]}},{"cell_type":"code","source":"do_fine_tuning = False \n\nprint(\"Building model with\", model_handle)\nmodel = tf.keras.Sequential([\n    # Explicitly define the input shape so the model can be properly\n    # loaded by the TFLiteConverter\n    \n    hub.KerasLayer(model_handle, trainable=do_fine_tuning),\n   \n])\nmodel.build((None,)+IMAGE_SIZE+(3,))\nmodel.summary()","metadata":{"papermill":{"duration":6.44893,"end_time":"2023-02-03T17:34:52.993834","exception":false,"start_time":"2023-02-03T17:34:46.544904","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-28T12:15:12.801998Z","iopub.execute_input":"2024-03-28T12:15:12.802964Z","iopub.status.idle":"2024-03-28T12:15:16.604882Z","shell.execute_reply.started":"2024-03-28T12:15:12.802923Z","shell.execute_reply":"2024-03-28T12:15:16.603902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Training the model","metadata":{"papermill":{"duration":0.014175,"end_time":"2023-02-03T17:34:53.023003","exception":false,"start_time":"2023-02-03T17:34:53.008828","status":"completed"},"tags":[]}},{"cell_type":"code","source":"model.compile(\n optimizer='adam',\n    loss = 'sparse_categorical_crossentropy',\n    metrics=['sparse_categorical_accuracy'],)\n\nNUM_TRAINING_IMAGES = 12753\nNUM_TEST_IMAGES = 7382\nNUM_VALIDATION_IMAGES=3712\nSTEPS_PER_EPOCH = NUM_TRAINING_IMAGES // BATCH_SIZE\nvalidation_steps = NUM_VALIDATION_IMAGES // BATCH_SIZE\nhist = model.fit(\n    training_dataset,\n    epochs=5, steps_per_epoch=STEPS_PER_EPOCH,\n    validation_data=validation_dataset,\n    validation_steps=validation_steps\n    ).history\n","metadata":{"papermill":{"duration":49.111774,"end_time":"2023-02-03T17:35:42.148519","exception":false,"start_time":"2023-02-03T17:34:53.036745","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-28T11:57:49.067220Z","iopub.status.idle":"2024-03-28T11:57:49.067566Z","shell.execute_reply.started":"2024-03-28T11:57:49.067390Z","shell.execute_reply":"2024-03-28T11:57:49.067406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure()\nplt.ylabel(\"Loss (training and validation)\")\nplt.xlabel(\"Training Steps\")\nplt.ylim([0,2])\nplt.plot(hist[\"loss\"])\nplt.plot(hist[\"val_loss\"])\n\nplt.figure()\nplt.ylabel(\"Accuracy (training and validation)\")\nplt.xlabel(\"Training Steps\")\nplt.ylim([0,1])\nplt.plot(hist[\"sparse_categorical_accuracy\"])\nplt.plot(hist[\"val_sparse_categorical_accuracy\"])","metadata":{"_kg_hide-input":true,"papermill":{"duration":0.429965,"end_time":"2023-02-03T17:35:42.618246","exception":false,"start_time":"2023-02-03T17:35:42.188281","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-28T11:57:49.069004Z","iopub.status.idle":"2024-03-28T11:57:49.069356Z","shell.execute_reply.started":"2024-03-28T11:57:49.069181Z","shell.execute_reply":"2024-03-28T11:57:49.069197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = get_test_dataset(ordered=True)\n\nprint('Computing predictions...')\ntest_images_ds = test_ds.map(lambda image, idnum: image)\nprobabilities = model.predict(test_images_ds)\npredictions = np.argmax(probabilities, axis=-1)\nprint(predictions)","metadata":{"_kg_hide-input":true,"papermill":{"duration":1.066882,"end_time":"2023-02-03T17:35:43.724923","exception":false,"start_time":"2023-02-03T17:35:42.658041","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-28T11:57:49.070779Z","iopub.status.idle":"2024-03-28T11:57:49.071157Z","shell.execute_reply.started":"2024-03-28T11:57:49.070951Z","shell.execute_reply":"2024-03-28T11:57:49.070968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Generating submission.csv file...')\n\n# Get image ids from test set and convert to unicode\ntest_ids_ds = test_dataset.map(lambda image, idnum: idnum).unbatch()\ntest_ids = next(iter(test_ids_ds.batch(NUM_TEST_IMAGES))).numpy().astype('U')\n\n# Write the submission file\nnp.savetxt(\n    'submission.csv',\n    np.rec.fromarrays([test_ids, predictions]),\n    fmt=['%s', '%d'],\n    delimiter=',',\n    header='id,label',\n    comments='',\n)\n\n# Look at the first few predictions\n!head submission.csv","metadata":{"execution":{"iopub.status.busy":"2024-03-28T11:57:49.072832Z","iopub.status.idle":"2024-03-28T11:57:49.073350Z","shell.execute_reply.started":"2024-03-28T11:57:49.073060Z","shell.execute_reply":"2024-03-28T11:57:49.073084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Credit:\n\nThis notebook was both a combination and an adaptation of the following tutorials from the TensorFlow team (modified to be relevant to the Kaggle Models product):\n - https://www.tensorflow.org/hub/tutorials/image_classification\n - https://www.tensorflow.org/hub/tutorials/tf2_image_retraining","metadata":{"papermill":{"duration":0.040592,"end_time":"2023-02-03T17:35:43.807495","exception":false,"start_time":"2023-02-03T17:35:43.766903","status":"completed"},"tags":[]}},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.040371,"end_time":"2023-02-03T17:35:43.888871","exception":false,"start_time":"2023-02-03T17:35:43.8485","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}