{"cells":[{"metadata":{"_uuid":"0832f86d-60dd-4ca3-8d14-095e91fc4d05","_cell_guid":"964e4107-6f1d-49f3-84d4-35b19d380110","trusted":true},"cell_type":"code","source":"from kaggle_datasets import KaggleDatasets\n\nimport tensorflow as tf\nimport numpy as np\nimport pandas as pd\nfrom matplotlib import pyplot as plt\nimport seaborn as sns\nsns.set()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"be67e1af-cd19-4712-8ae7-eba2aef22828","_cell_guid":"7e75974c-4236-4304-be43-a5bd81cd9f01","trusted":true},"cell_type":"code","source":"# Detect hardware, return appropriate distribution strategy\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver()  # TPU detection. No parameters necessary if TPU_NAME environment variable is set. On Kaggle this is always the case.\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)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"17030b0b-f99d-46d3-9151-7b1000095fb4","_cell_guid":"cb114b02-a5d9-43a3-a91e-d2c1306844a0","trusted":true},"cell_type":"code","source":"# Competition data access\n# TPUs read data directly from Google Cloud Storage (GCS). \n# This Kaggle utility will copy the dataset to a GCS bucket\n# co-located with the TPU.\nGCS_DS_PATH = KaggleDatasets().get_gcs_path()\nGCS_DS_PATH","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"7aed880f-4971-495c-bbd0-0ff4644cb33e","_cell_guid":"9452f489-e61a-461e-b6ae-008b10f9c5d6","trusted":true},"cell_type":"code","source":"# GCS_PATH = os.path.join(GCS_DS_PATH + '/tfrecords-jpeg-512x512')\nGCS_PATH = GCS_DS_PATH + '/tfrecords-jpeg-512x512'\n_TRAINING_FILENAMES = tf.io.gfile.glob( \\\n    GCS_PATH + '/train/*.tfrec')\n_VALIDATION_FILENAMES = tf.io.gfile.glob( \\\n    GCS_PATH + '/val/*.tfrec')\n_TEST_FILENAMES = tf.io.gfile.glob( \\\n    GCS_PATH + '/test/*.tfrec')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"IMAGE_SIZE = [512, 512]\nEPOCHS = 10\nBATCH_SIZE = 20","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"7825f94f-1fc7-42bd-aeb6-7ee03f89627b","_cell_guid":"c3c937c1-1188-4c7b-bb3f-3b39f92cbcf7","trusted":true},"cell_type":"code","source":"def plotBatch(dataset):\n    images, labels = next(\n        iter(dataset.unbatch().batch(BATCH_SIZE)))\n    images = images.numpy()\n    labels = labels.numpy()\n\n    cols = 5\n    rows = -((-len(labels)) // cols)\n\n    fig, axes = plt.subplots(rows, cols)\n    for row in range(rows):\n        for col in range(cols):\n            idx = row * cols + col\n            ax = axes[row, col]\n            ax.imshow(images[idx])\n            ax.axis(\"off\")\n            label = labels[idx]\n","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"1e71e362-c926-415e-be97-4290f3c426eb","_cell_guid":"4dcdf154-6720-42a7-b9fe-d222243c2d23","trusted":true},"cell_type":"code","source":"def decode_image(image_data):\n    image = tf.image.decode_jpeg(image_data, channels=3)\n    image = tf.cast(image, tf.float32) / 255.0\n    image = tf.reshape(image, [*IMAGE_SIZE, 3])\n    return image","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"c9c34d39-4869-4cd7-90ff-ea024dc6c669","_cell_guid":"c8ef4b8f-887e-48cd-9b94-2f8fd7da843e","trusted":true},"cell_type":"code","source":"def read_labeled_tfrecord(example):\n    format = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"class\": tf.io.FixedLenFeature([], tf.int64),\n    }\n\n    example = tf.io.parse_single_example(\n        example, format)\n    image = decode_image(example['image'])\n    label = tf.cast(example['class'], tf.int32)\n\n    return image, label\n\n\ndef read_unlabeled_tfrecord(example):\n    format = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"id\": tf.io.FixedLenFeature([], tf.string),\n    }\n\n    example = tf.io.parse_single_example(\n        example, format)\n    image = decode_image(example['image'])\n    id = example['id']\n\n    return image, id","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"768afb9a-6b08-4454-b613-caee0fe2a1bb","_cell_guid":"4ab373cd-8703-44b1-8eaf-2c37a1e42c85","trusted":true},"cell_type":"code","source":"def getTrainData():\n    dataset = tf.data.TFRecordDataset(\n        _TRAINING_FILENAMES)\n    dataset = dataset.map(read_labeled_tfrecord)\n    dataset = dataset.cache()\n    dataset = dataset.repeat(EPOCHS)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(1)\n\n    return dataset","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"2bb56bde-4927-4527-be9a-1806ffa12c94","_cell_guid":"3b3f30c8-5d03-4fbf-82de-7fb7915981aa","trusted":true},"cell_type":"code","source":"def getValidationData():\n    dataset = tf.data.TFRecordDataset(\n        _VALIDATION_FILENAMES)\n    dataset = dataset.map(read_labeled_tfrecord)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(1)\n\n    return dataset","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d5b87c44-5127-434a-b233-64d82722eb93","_cell_guid":"35f5c6ec-c599-4c03-9818-b0b87825d399","trusted":true},"cell_type":"code","source":"def getTestData():\n    dataset = tf.data.TFRecordDataset(\n        _TEST_FILENAMES)\n    dataset = dataset.map(read_unlabeled_tfrecord)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(1)\n\n    return dataset","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"82527bc7-3bad-4349-9269-72600c3ca1d1","_cell_guid":"997a36de-9cf8-436a-825d-db6f499774fe","trusted":true},"cell_type":"code","source":"train_dataset = getTrainData()\ntest_dataset = getTestData()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"8a4a7e99-71f8-441b-abec-0cbebe9aa1cb","_cell_guid":"62744cad-4eae-45c8-89f6-beb784102330","trusted":true},"cell_type":"code","source":"plotBatch(train_dataset)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"plotBatch(test_dataset)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"850ec7ac-6f87-42a9-9e46-811544f16bc7","_cell_guid":"e950d628-2168-49eb-a075-55723079db09","trusted":true},"cell_type":"code","source":"def define_model(input_shape, n_classes):\n    inp = tf.keras.layers.Input(shape=input_shape)\n    vgg16 = tf.keras.applications.VGG16(include_top=False)\n    for layer in vgg16.layers:\n        layer.trainable = False\n    vgg16Out = vgg16(inp)\n    avgPool = tf.keras.layers.GlobalAveragePooling2D()(vgg16Out)\n    dense1 = tf.keras.layers.Dense(1024, activation=tf.nn.relu)(avgPool)\n    out = tf.keras.layers.Dense(n_classes)(dense1)\n\n    model = tf.keras.Model(inp, out)\n    \n    return model","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"b8b4f7c6-f125-47e3-9662-277f7611214f","_cell_guid":"34d1c19b-4c3e-4846-9f12-ee2664befc88","trusted":true},"cell_type":"code","source":"model = define_model((*IMAGE_SIZE, 3), 104)\nloss = tf.losses.SparseCategoricalCrossentropy()\noptimizer = tf.keras.optimizers.Adam(0.001)\naccuracy = tf.metrics.Accuracy()\nstep = tf.Variable(1, name=\"global_step\")","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"e7dc4016-23f1-4e40-8fac-02f5c8d0b4ce","_cell_guid":"5203a7dd-9a11-4999-b0e1-3eba11c95cea","trusted":true},"cell_type":"code","source":"@tf.function\ndef train_step(features, labels):\n    with tf.GradientTape() as tape:\n        logits = model(features)\n        loss_value = loss(labels, logits) \n    \n    gradients = tape.gradient(loss_value,\n                    model.trainable_variables)\n    optimizer.apply_gradients(\n        zip(gradients, model.trainable_variables))\n    step.assign_add(1)\n    accuracy_value = accuracy(labels,\n                        tf.argmax(logits, -1))\n\n    return loss_value, accuracy_value","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"49e5ed89-9bb0-4210-a1e6-802474c91508","_cell_guid":"f44a8459-490a-4b03-ab25-90a8786a0562","trusted":true},"cell_type":"code","source":"# @tf.function\ndef loop(inputs):\n    for features, labels in inputs:\n        loss_value, accuracy_value = train_step(\n            features, labels)\n        if step.numpy() % 10 == 0:\n            tf.print(\"step: {} loss: {} acc: {}\".format(\n                 step.numpy(), loss_value.numpy(), \n                    accuracy_value.numpy()))","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"7c65ba75-3897-48d6-8e48-01491a2fbbd7","_cell_guid":"a942b128-ec1e-4bcf-a95a-1c75d8ac9e95","trusted":true},"cell_type":"code","source":"loop(train_dataset)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"988e6083-c9bd-4ebe-974c-86f6c8d8467a","_cell_guid":"2a77a27d-f04e-4730-b6b9-c99cb43360c2","trusted":true},"cell_type":"code","source":"test_dataset = getTestData()\nimages = test_dataset.map(lambda img, id: img)\npreds = list()\nfor imgs in images:\n    logits = model(imgs)\n    logits = logits.numpy()\n    preds.extend(logits)\npreds = np.array(preds)\npreds = np.argmax(preds, axis=-1)\nprint(preds)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"6fa4b8ff-6992-40e0-99dc-fe7426f16eb7","_cell_guid":"3a1a72f3-9543-42d9-b073-082aeaa55189","trusted":true},"cell_type":"code","source":"np.unique(preds)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"7e3f2781-c552-4057-9b01-831430576c5e","_cell_guid":"c2fd00c1-cab2-436b-baf4-d5d574e8a3c7","trusted":true},"cell_type":"code","source":"ids = test_dataset.map(lambda image, id: id)\nids = list(ids.unbatch())\nids = tf.convert_to_tensor(ids)\nids = ids.numpy().astype(\"U\")\nids","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"ea0bcd30-c172-450d-bfb8-56d389b44f90","_cell_guid":"48c8870f-95c1-43a3-9fab-15362d5f7f37","trusted":true},"cell_type":"code","source":"submission_df = pd.DataFrame(dict(id=ids, label=preds))\nsubmission_df.to_csv(\"submission.csv\", index=False)","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}