{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Set below variable to whatever accelerator you use\n\nGet a TPU if you can, otherwise you'll be waiting a while","metadata":{}},{"cell_type":"code","source":"ACCELERATOR = 'tpu-vm' # 'tpu-vm', 'tpu', 'gpu'","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:56:22.837926Z","iopub.execute_input":"2023-03-28T05:56:22.838786Z","iopub.status.idle":"2023-03-28T05:56:23.410104Z","shell.execute_reply.started":"2023-03-28T05:56:22.838749Z","shell.execute_reply":"2023-03-28T05:56:23.409193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# These are your independent variables","metadata":{}},{"cell_type":"code","source":"USE_AUGMENTATION = True\nUSE_EXTERNAL_DATA = True\nUSE_CLASS_WEIGHTS = True\nUSE_10_CROP_TESTING = True\n\nRESNET_DEPTH = 50 # 50, 101, 152","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:56:23.411657Z","iopub.execute_input":"2023-03-28T05:56:23.411948Z","iopub.status.idle":"2023-03-28T05:56:23.568352Z","shell.execute_reply.started":"2023-03-28T05:56:23.41192Z","shell.execute_reply":"2023-03-28T05:56:23.567334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Image Sizes","metadata":{}},{"cell_type":"markdown","source":"### Higher resolution likely gives higher accuracy with the expense of longer training time","metadata":{}},{"cell_type":"markdown","source":"### Also high resolution images require more memory","metadata":{}},{"cell_type":"markdown","source":"#### Crop Size is the size used by random cropping","metadata":{}},{"cell_type":"code","source":"IMAGE_SIZE = 512\nCROP_SIZE = 331","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:56:25.783632Z","iopub.execute_input":"2023-03-28T05:56:25.784253Z","iopub.status.idle":"2023-03-28T05:56:25.788579Z","shell.execute_reply.started":"2023-03-28T05:56:25.784211Z","shell.execute_reply":"2023-03-28T05:56:25.787708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# These relate to the training process.\nThere's probably some golden values that make training faster","metadata":{}},{"cell_type":"markdown","source":"### Depending on other variables, 16 may be rather low for epochs.\n### But those take time, P100 takes ~3 min per epoch with 331x331 images","metadata":{}},{"cell_type":"code","source":"if ACCELERATOR[0:3] == 'tpu':\n    BATCH_SIZE = 128\nelse:\n    BATCH_SIZE = 16\n\nSTEPS_PER_EXECUTION = 16\nEPOCHS = 16\nLR_PATIENCE = 3\nINITIAL_LR = 0.0001","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:56:31.126458Z","iopub.execute_input":"2023-03-28T05:56:31.126857Z","iopub.status.idle":"2023-03-28T05:56:41.690047Z","shell.execute_reply.started":"2023-03-28T05:56:31.126822Z","shell.execute_reply":"2023-03-28T05:56:41.688953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### If using a TPU VM, you need to install certain versions of packages","metadata":{}},{"cell_type":"code","source":"# For TPU VM\nif ACCELERATOR == 'tpu-vm':\n    !pip install /lib/wheels/tensorflow-2.9.1-cp38-cp38-linux_x86_64.whl\n    !pip install scikit-learn","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:56:41.691719Z","iopub.execute_input":"2023-03-28T05:56:41.692017Z","iopub.status.idle":"2023-03-28T05:56:49.544171Z","shell.execute_reply.started":"2023-03-28T05:56:41.69199Z","shell.execute_reply":"2023-03-28T05:56:49.54316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Normal imports","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nfrom kaggle_datasets import KaggleDatasets\nimport numpy as np\nfrom sklearn.utils import class_weight\n\nprint(\"Tensorflow version \" + tf.__version__)","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:56:57.253633Z","iopub.execute_input":"2023-03-28T05:56:57.254017Z","iopub.status.idle":"2023-03-28T05:57:00.57764Z","shell.execute_reply.started":"2023-03-28T05:56:57.253984Z","shell.execute_reply":"2023-03-28T05:57:00.57668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Get computation strategy for acceleration","metadata":{}},{"cell_type":"code","source":"# For TPU VM\nif ACCELERATOR == 'tpu-vm':\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect(tpu=\"local\") # \"local\" for 1VM TPU\n    strategy = tf.distribute.TPUStrategy(tpu)","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:57:05.625872Z","iopub.execute_input":"2023-03-28T05:57:05.627226Z","iopub.status.idle":"2023-03-28T05:57:37.295823Z","shell.execute_reply.started":"2023-03-28T05:57:05.627177Z","shell.execute_reply":"2023-03-28T05:57:37.294857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# For TPU\nif ACCELERATOR == 'tpu':\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver()\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.experimental.TPUStrategy(tpu)","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:57:37.297413Z","iopub.execute_input":"2023-03-28T05:57:37.297824Z","iopub.status.idle":"2023-03-28T05:57:54.476445Z","shell.execute_reply.started":"2023-03-28T05:57:37.297791Z","shell.execute_reply":"2023-03-28T05:57:54.475449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# For GPU\nif ACCELERATOR == 'gpu':\n    strategy = tf.distribute.get_strategy()\n    tpu = False","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:57:54.478085Z","iopub.execute_input":"2023-03-28T05:57:54.478445Z","iopub.status.idle":"2023-03-28T05:57:54.704958Z","shell.execute_reply.started":"2023-03-28T05:57:54.478413Z","shell.execute_reply":"2023-03-28T05:57:54.703922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Get Data","metadata":{}},{"cell_type":"code","source":"COMPETITION_DATA_PATH = KaggleDatasets().get_gcs_path(\"tpu-getting-started\")\nEXTERNAL_DATA_PATH = KaggleDatasets().get_gcs_path(\"tf-flower-photo-tfrec\")","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:57:57.829087Z","iopub.execute_input":"2023-03-28T05:57:57.829934Z","iopub.status.idle":"2023-03-28T05:57:58.466337Z","shell.execute_reply.started":"2023-03-28T05:57:57.829897Z","shell.execute_reply":"2023-03-28T05:57:58.465211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMAGE_SIZE_PATH = f'/tfrecords-jpeg-{IMAGE_SIZE}x{IMAGE_SIZE}'\nCROP_SIZE_PATH = f'/tfrecords-jpeg-{CROP_SIZE}x{CROP_SIZE}'","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:57:58.467945Z","iopub.execute_input":"2023-03-28T05:57:58.468249Z","iopub.status.idle":"2023-03-28T05:58:22.008769Z","shell.execute_reply.started":"2023-03-28T05:57:58.468223Z","shell.execute_reply":"2023-03-28T05:58:22.007662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_FILENAMES = tf.io.gfile.glob(COMPETITION_DATA_PATH + IMAGE_SIZE_PATH + '/train/*.tfrec')\nVAL_FILENAMES = tf.io.gfile.glob(COMPETITION_DATA_PATH + IMAGE_SIZE_PATH + '/val/*.tfrec')\nTEST_FILENAMES = tf.io.gfile.glob(COMPETITION_DATA_PATH + IMAGE_SIZE_PATH + '/test/*.tfrec')\n\nif USE_EXTERNAL_DATA:\n    IMAGENET_FILES = tf.io.gfile.glob(EXTERNAL_DATA_PATH + '/imagenet_no_test' + IMAGE_SIZE_PATH + '/*.tfrec')\n    INATURELIST_FILES = tf.io.gfile.glob(EXTERNAL_DATA_PATH + '/inaturalist_no_test' + IMAGE_SIZE_PATH + '/*.tfrec')\n    OPENIMAGE_FILES = tf.io.gfile.glob(EXTERNAL_DATA_PATH + '/openimage_no_test' + IMAGE_SIZE_PATH + '/*.tfrec')\n    OXFORD_FILES = tf.io.gfile.glob(EXTERNAL_DATA_PATH + '/oxford_102_no_test' + IMAGE_SIZE_PATH + '/*.tfrec')\n    TENSORFLOW_FILES = tf.io.gfile.glob(EXTERNAL_DATA_PATH + '/tf_flowers_no_test' + IMAGE_SIZE_PATH + '/*.tfrec')\n\n    TRAIN_FILENAMES = TRAIN_FILENAMES + IMAGENET_FILES + INATURELIST_FILES + OXFORD_FILES + TENSORFLOW_FILES","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:58:22.009923Z","iopub.execute_input":"2023-03-28T05:58:22.010192Z","iopub.status.idle":"2023-03-28T05:58:22.412121Z","shell.execute_reply.started":"2023-03-28T05:58:22.010166Z","shell.execute_reply":"2023-03-28T05:58:22.410859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import re\n\ndef count_files(filenames):\n    return sum([int(re.compile('(\\d+)\\.tfrec').search(i).groups()[0]) for i in filenames])","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:58:22.414289Z","iopub.execute_input":"2023-03-28T05:58:22.414642Z","iopub.status.idle":"2023-03-28T05:58:22.420335Z","shell.execute_reply.started":"2023-03-28T05:58:22.414611Z","shell.execute_reply":"2023-03-28T05:58:22.419421Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"count_files(TRAIN_FILENAMES)","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:58:22.421426Z","iopub.execute_input":"2023-03-28T05:58:22.421692Z","iopub.status.idle":"2023-03-28T05:58:27.8656Z","shell.execute_reply.started":"2023-03-28T05:58:22.421668Z","shell.execute_reply":"2023-03-28T05:58:27.864378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"count_files(VAL_FILENAMES)","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:58:27.867032Z","iopub.execute_input":"2023-03-28T05:58:27.867406Z","iopub.status.idle":"2023-03-28T05:58:29.931369Z","shell.execute_reply.started":"2023-03-28T05:58:27.867379Z","shell.execute_reply":"2023-03-28T05:58:29.930247Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"count_files(TEST_FILENAMES)","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:58:29.932653Z","iopub.execute_input":"2023-03-28T05:58:29.93297Z","iopub.status.idle":"2023-03-28T05:58:31.426108Z","shell.execute_reply.started":"2023-03-28T05:58:29.932943Z","shell.execute_reply":"2023-03-28T05:58:31.424867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_BATCHES = count_files(TRAIN_FILENAMES) // BATCH_SIZE\nVAL_BATCHES = count_files(VAL_FILENAMES) // BATCH_SIZE\nTEST_BATCHES = count_files(TEST_FILENAMES) // BATCH_SIZE","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:59:13.330329Z","iopub.execute_input":"2023-03-28T05:59:13.330824Z","iopub.status.idle":"2023-03-28T05:59:13.336856Z","shell.execute_reply.started":"2023-03-28T05:59:13.330784Z","shell.execute_reply":"2023-03-28T05:59:13.335878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Parse and Preprocess Data","metadata":{}},{"cell_type":"code","source":"def parse(labeled=True):\n    features = {\n        'image': tf.io.FixedLenFeature([], tf.string),\n        'class': tf.io.FixedLenFeature([], tf.int64),\n    } if labeled else {\n        'image': tf.io.FixedLenFeature([], tf.string),\n        'id': tf.io.FixedLenFeature([], tf.string)\n    }\n    def parse_func(raw):\n        data = tf.io.parse_example(raw, features)\n        image = tf.image.decode_jpeg(data['image'], channels=3)\n        #image = tf.image.resize(image, (IMAGE_SIZE, IMAGE_SIZE))\n        label = data['class' if labeled else 'id']\n        return image, label\n    return parse_func","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:59:14.302533Z","iopub.execute_input":"2023-03-28T05:59:14.303508Z","iopub.status.idle":"2023-03-28T05:59:17.6813Z","shell.execute_reply.started":"2023-03-28T05:59:14.303473Z","shell.execute_reply":"2023-03-28T05:59:17.67997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labeled_parse = parse()\nunlabeled_parse = parse(labeled=False)","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:59:17.683094Z","iopub.execute_input":"2023-03-28T05:59:17.68342Z","iopub.status.idle":"2023-03-28T05:59:17.813642Z","shell.execute_reply.started":"2023-03-28T05:59:17.683373Z","shell.execute_reply":"2023-03-28T05:59:17.812515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = tf.data.TFRecordDataset(TRAIN_FILENAMES)\nds = ds.shuffle(count_files(TRAIN_FILENAMES))\nds = ds.map(labeled_parse)\n\nds = ds.map(lambda image, label: (tf.image.resize(image, (IMAGE_SIZE, IMAGE_SIZE)), label))\n\ndef augment(image, label):\n    image = tf.image.random_flip_left_right(image)\n    image = tf.image.random_crop(image, (CROP_SIZE, CROP_SIZE, 3))\n    return image, label\n\nif USE_AUGMENTATION:\n    ds = ds.map(augment)\n\nds = ds.batch(BATCH_SIZE, drop_remainder=True)\nds = ds.prefetch(tf.data.AUTOTUNE)","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:59:17.814958Z","iopub.execute_input":"2023-03-28T05:59:17.815252Z","iopub.status.idle":"2023-03-28T05:59:18.473792Z","shell.execute_reply.started":"2023-03-28T05:59:17.815226Z","shell.execute_reply":"2023-03-28T05:59:18.472474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if USE_CLASS_WEIGHTS:\n    ds_freq = dict(zip(range(104), class_weight.compute_class_weight('balanced', classes=range(104), y=[l.numpy() for l in tf.data.TFRecordDataset(TRAIN_FILENAMES).map(labeled_parse).map(lambda image, label: label)])))","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:59:18.476192Z","iopub.execute_input":"2023-03-28T05:59:18.47651Z","iopub.status.idle":"2023-03-28T05:59:18.482672Z","shell.execute_reply.started":"2023-03-28T05:59:18.476482Z","shell.execute_reply":"2023-03-28T05:59:18.481686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vds = tf.data.TFRecordDataset(VAL_FILENAMES)\nvds = vds.map(labeled_parse)\n# if USE_AUGMENTATION:\n#     vds = vds.map(lambda image, label: (tf.image.resize(image, (CROP_SIZE, CROP_SIZE)), label))\nvds = vds.batch(BATCH_SIZE, drop_remainder=True)\nvds = vds.prefetch(tf.data.AUTOTUNE)\nif tpu:\n    vds = vds.cache()","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:59:18.483743Z","iopub.execute_input":"2023-03-28T05:59:18.484003Z","iopub.status.idle":"2023-03-28T05:59:18.727465Z","shell.execute_reply.started":"2023-03-28T05:59:18.483977Z","shell.execute_reply":"2023-03-28T05:59:18.726375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tds = tf.data.TFRecordDataset(TEST_FILENAMES)\ntds = tds.map(unlabeled_parse)\ntds = tds.map(lambda image, label: (tf.image.resize(image, (IMAGE_SIZE, IMAGE_SIZE)), label))","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:59:18.728656Z","iopub.execute_input":"2023-03-28T05:59:18.72893Z","iopub.status.idle":"2023-03-28T05:59:18.851066Z","shell.execute_reply.started":"2023-03-28T05:59:18.728902Z","shell.execute_reply":"2023-03-28T05:59:18.849808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create Model","metadata":{}},{"cell_type":"code","source":"from tensorflow.keras.models import Model\nfrom tensorflow.keras.layers import Input, Flatten, Dense, Conv2D, MaxPool2D, GlobalMaxPool2D, GlobalAveragePooling2D, Dropout","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:59:18.852487Z","iopub.execute_input":"2023-03-28T05:59:18.852799Z","iopub.status.idle":"2023-03-28T05:59:18.977515Z","shell.execute_reply.started":"2023-03-28T05:59:18.852769Z","shell.execute_reply":"2023-03-28T05:59:18.976439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with strategy.scope():\n    \n    x = Input(shape=(None, None, 3))\n    \n    if RESNET_DEPTH == 50:\n        resnet = tf.keras.applications.ResNet50(include_top=False, input_shape=[None, None, 3], )#weights=None)\n    elif RESNET_DEPTH == 101:\n        resnet = tf.keras.applications.ResNet101(include_top=False, input_shape=[None, None, 3])\n    elif RESNET_DEPTH == 152:\n        resnet = tf.keras.applications.ResNet152(include_top=False, input_shape=[None, None, 3])\n\n    y = tf.keras.applications.resnet50.preprocess_input(x)\n    features = resnet(y)\n    y = GlobalAveragePooling2D()(features)\n    hidden = Dense(2048, activation='relu')\n    classifier = Dense(104, activation='softmax')\n    y = classifier(Dropout(0.5)(hidden(y)))\n    extract = classifier(hidden(features))\n\n    model = Model(inputs=x, outputs=y)\n    extraction = Model(inputs=x, outputs=extract)\n    \n    if USE_10_CROP_TESTING:\n\n        x = Input(shape=(IMAGE_SIZE, IMAGE_SIZE, 3))\n        c = tf.image.central_crop(x, CROP_SIZE / IMAGE_SIZE)\n        c = tf.image.resize(c, (CROP_SIZE, CROP_SIZE))\n        tl = tf.image.crop_to_bounding_box(x, 0, 0, CROP_SIZE, CROP_SIZE)\n        tr = tf.image.crop_to_bounding_box(x, 0, IMAGE_SIZE - CROP_SIZE, CROP_SIZE, CROP_SIZE)\n        bl = tf.image.crop_to_bounding_box(x, IMAGE_SIZE - CROP_SIZE, 0, CROP_SIZE, CROP_SIZE)\n        br = tf.image.crop_to_bounding_box(x, IMAGE_SIZE - CROP_SIZE, IMAGE_SIZE - CROP_SIZE, CROP_SIZE, CROP_SIZE)\n\n        cf = tf.image.flip_left_right(c)\n        tlf = tf.image.flip_left_right(tl)\n        trf = tf.image.flip_left_right(tr)\n        blf = tf.image.flip_left_right(bl)\n        brf = tf.image.flip_left_right(br)\n\n        y = tf.reduce_mean(tf.stack([model(c), model(tl), model(tr), model(bl), model(br), model(cf), model(tlf), model(trf), model(blf), model(brf)], axis=1), axis=1)\n        test_model = Model(inputs=x, outputs=y)\n    else:\n        test_model = model","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:59:18.978785Z","iopub.execute_input":"2023-03-28T05:59:18.97919Z","iopub.status.idle":"2023-03-28T05:59:38.291403Z","shell.execute_reply.started":"2023-03-28T05:59:18.979161Z","shell.execute_reply":"2023-03-28T05:59:38.289924Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.summary()","metadata":{"execution":{"iopub.status.busy":"2023-03-28T05:59:58.436169Z","iopub.execute_input":"2023-03-28T05:59:58.43738Z","iopub.status.idle":"2023-03-28T05:59:58.47573Z","shell.execute_reply.started":"2023-03-28T05:59:58.437333Z","shell.execute_reply":"2023-03-28T05:59:58.474585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_model.summary()","metadata":{"execution":{"iopub.status.busy":"2023-03-28T06:00:00.320588Z","iopub.execute_input":"2023-03-28T06:00:00.321852Z","iopub.status.idle":"2023-03-28T06:00:00.364037Z","shell.execute_reply.started":"2023-03-28T06:00:00.321803Z","shell.execute_reply":"2023-03-28T06:00:00.3627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with strategy.scope():\n    model.compile(steps_per_execution=STEPS_PER_EXECUTION, optimizer=tf.keras.optimizers.Adam(learning_rate=INITIAL_LR), loss='sparse_categorical_crossentropy', metrics=['sparse_categorical_accuracy'])\n    learning_callback = tf.keras.callbacks.ReduceLROnPlateau(verbose=1, patience=LR_PATIENCE)","metadata":{"execution":{"iopub.status.busy":"2023-03-28T06:00:00.94692Z","iopub.execute_input":"2023-03-28T06:00:00.948207Z","iopub.status.idle":"2023-03-28T06:00:01.070246Z","shell.execute_reply.started":"2023-03-28T06:00:00.948158Z","shell.execute_reply":"2023-03-28T06:00:01.068799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if USE_CLASS_WEIGHTS:\n    historical = model.fit(ds.repeat(), class_weight=ds_freq, epochs=EPOCHS, validation_data=vds, callbacks=[learning_callback], steps_per_epoch=TRAIN_BATCHES, validation_steps=VAL_BATCHES)\nelse:\n    historical = model.fit(ds.repeat(), epochs=EPOCHS, validation_data=vds, callbacks=[learning_callback], steps_per_epoch=TRAIN_BATCHES, validation_steps=VAL_BATCHES)","metadata":{"execution":{"iopub.status.busy":"2023-03-28T06:00:13.850001Z","iopub.execute_input":"2023-03-28T06:00:13.85096Z","iopub.status.idle":"2023-03-28T06:02:06.829301Z","shell.execute_reply.started":"2023-03-28T06:00:13.850914Z","shell.execute_reply":"2023-03-28T06:02:06.827672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndef plot_historical(historical):\n    loss_values = historical.history['loss']\n    val_loss_values = historical.history['val_loss']\n    epochs = range(1, len(loss_values)+1)\n\n    plt.plot(epochs, loss_values, label='Training Loss')\n    plt.plot(epochs, val_loss_values, label='Validation Loss')\n\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.legend()\n\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-03-28T06:02:08.737267Z","iopub.execute_input":"2023-03-28T06:02:08.738124Z","iopub.status.idle":"2023-03-28T06:02:09.204525Z","shell.execute_reply.started":"2023-03-28T06:02:08.738079Z","shell.execute_reply":"2023-03-28T06:02:09.202927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_historical(historical)","metadata":{"execution":{"iopub.status.busy":"2023-03-28T06:02:09.763339Z","iopub.execute_input":"2023-03-28T06:02:09.764589Z","iopub.status.idle":"2023-03-28T06:02:10.003689Z","shell.execute_reply.started":"2023-03-28T06:02:09.764541Z","shell.execute_reply":"2023-03-28T06:02:10.0022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_images(dataset, grid_size=[4,4]):\n\n    ds = dataset.unbatch()\n    images = []\n    for i in ds.take(grid_size[0] * grid_size[1]):\n        images.append(i[0])\n\n    fig, axes = plt.subplots(nrows=grid_size[0], ncols=grid_size[1], figsize=(10,10))\n    fig.subplots_adjust(left=0, right=1, bottom=0, top=1, hspace=0, wspace=0)\n\n    for ax, img in zip(axes.flat, images):\n        ax.imshow(img, cmap='gray')\n        ax.set_aspect('equal')\n        ax.axis('off')\n        #ax.text(192 / 2, 192, label, ha='center', va='bottom', color='white', fontsize=20, bbox={'facecolor':'blue', 'alpha':0.5})\n\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-03-28T06:02:11.945641Z","iopub.execute_input":"2023-03-28T06:02:11.946891Z","iopub.status.idle":"2023-03-28T06:02:11.955207Z","shell.execute_reply.started":"2023-03-28T06:02:11.946844Z","shell.execute_reply":"2023-03-28T06:02:11.95394Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = tf.image.resize(ds.unbatch().shuffle(128).take(1).get_single_element()[0], (CROP_SIZE, CROP_SIZE))\n\nplt.imshow(img / 255)","metadata":{"execution":{"iopub.status.busy":"2023-03-28T06:02:13.386153Z","iopub.execute_input":"2023-03-28T06:02:13.386624Z","iopub.status.idle":"2023-03-28T06:02:17.796806Z","shell.execute_reply.started":"2023-03-28T06:02:13.386587Z","shell.execute_reply":"2023-03-28T06:02:17.795366Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ex = extraction(tf.reshape(img, (1, CROP_SIZE, CROP_SIZE, 3)))[0]\ni = tf.argmax(model(tf.reshape(img, (1, CROP_SIZE, CROP_SIZE, 3)))[0]).numpy()\n\noverlay = ex[:,:,i]\n\noverlay_img = tf.expand_dims(overlay, axis=-1)\noverlay_img = tf.repeat(overlay_img, 3, axis=-1)\n\nfig, ax = plt.subplots()\nax.imshow(img / 255)\n#ax.imshow(tf.image.resize(overlay_img, (192, 192), method='nearest'), alpha=0.3)\nax.imshow(overlay, cmap='jet', alpha=0.5, interpolation='bilinear', extent=(0, img.shape[1], img.shape[0], 0))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-03-28T06:02:19.496026Z","iopub.execute_input":"2023-03-28T06:02:19.496514Z","iopub.status.idle":"2023-03-28T06:02:22.078918Z","shell.execute_reply.started":"2023-03-28T06:02:19.496475Z","shell.execute_reply":"2023-03-28T06:02:22.077573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# Get the output and predicted class of the model\nwith tf.GradientTape() as tape:\n    tape.watch(img)\n    img2 = tf.reshape(img, (1, CROP_SIZE, CROP_SIZE, 3))\n    outputs = model(img2)[0]\n    predicted_class = tf.argmax(outputs)\n    output = outputs[predicted_class]\n\n# Calculate the gradients of the predicted class with respect to the image\ngrads = tape.gradient(output, img)\n\n# Calculate the guided gradients by removing negative gradients\nguided_grads = tf.cast(img > 0, 'float32') * tf.cast(grads > 0, 'float32') * grads\n\n# Calculate the average attention map by taking the mean of the guided gradients\nheatmap = tf.reduce_mean(guided_grads, axis=-1)\n\n# define the gamma value for contrast enhancement\ngamma = 0.6\n\n# normalize the heatmap between 0 and 1\nheatmap = tf.maximum(heatmap, 0) / tf.reduce_max(heatmap)\n\n# apply gamma correction to enhance the visualization\nheatmap = tf.pow(heatmap, gamma)\n\n# Plot the heatmap and overlay it on the original image\nplt.imshow(img2[0] / 255)\nplt.imshow(heatmap, alpha=0.5, cmap='jet')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-03-28T06:02:25.0927Z","iopub.execute_input":"2023-03-28T06:02:25.093843Z","iopub.status.idle":"2023-03-28T06:02:26.770209Z","shell.execute_reply.started":"2023-03-28T06:02:25.093799Z","shell.execute_reply":"2023-03-28T06:02:26.768874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.save('history.npy',historical.history)","metadata":{"execution":{"iopub.status.busy":"2023-03-28T06:02:29.285556Z","iopub.execute_input":"2023-03-28T06:02:29.286466Z","iopub.status.idle":"2023-03-28T06:02:29.292324Z","shell.execute_reply.started":"2023-03-28T06:02:29.286424Z","shell.execute_reply":"2023-03-28T06:02:29.291021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.save('model.h5')","metadata":{"execution":{"iopub.status.busy":"2023-03-28T06:02:30.149787Z","iopub.execute_input":"2023-03-28T06:02:30.150623Z","iopub.status.idle":"2023-03-28T06:02:33.313405Z","shell.execute_reply.started":"2023-03-28T06:02:30.150571Z","shell.execute_reply":"2023-03-28T06:02:33.311893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Computing predictions...')\ntest_images_ds = tds.map(lambda image, idnum: image).batch(BATCH_SIZE)\nprobabilities = test_model.predict(test_images_ds, steps=TEST_BATCHES+1)\npredictions = np.argmax(probabilities, axis=-1)\nprint(predictions)\n\nprint('Generating submission.csv file...')\ntest_ids_ds = tds.map(lambda image, idnum: idnum)\ntest_ids = next(iter(test_ids_ds.batch(count_files(TEST_FILENAMES)))).numpy().astype('U') # all in one batch\nnp.savetxt('submission.csv', np.rec.fromarrays([test_ids, predictions]), fmt=['%s', '%d'], delimiter=',', header='id,label', comments='')","metadata":{"execution":{"iopub.status.busy":"2023-03-28T06:02:34.92123Z","iopub.execute_input":"2023-03-28T06:02:34.922346Z","iopub.status.idle":"2023-03-28T06:03:18.230335Z","shell.execute_reply.started":"2023-03-28T06:02:34.922307Z","shell.execute_reply":"2023-03-28T06:03:18.229052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}