{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"},{"sourceId":8009334,"sourceType":"datasetVersion","datasetId":4717506},{"sourceId":8018537,"sourceType":"datasetVersion","datasetId":4724628},{"sourceId":8018622,"sourceType":"datasetVersion","datasetId":4724693},{"sourceId":8229905,"sourceType":"datasetVersion","datasetId":4752193}],"dockerImageVersionId":30674,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import math, re, os\nimport tensorflow as tf\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom kaggle_datasets import KaggleDatasets\nfrom tensorflow import keras\nfrom functools import partial\nfrom sklearn.model_selection import train_test_split\nimport tensorflow as tf\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom transformers import ViTImageProcessor\nfrom PIL import Image\nfrom transformers import ViTModel\nimport torchvision.transforms as transforms\n\nclass CassavaLeafTestDataset(Dataset):\n    def __init__(self, tfrecord_files, image_processor):\n        self.tfrecord_files = tfrecord_files\n        self.image_processor = image_processor\n        self.images , self.image_ids = self.load_and_preprocess(tfrecord_files)\n\n    def decode_image(self, image):\n        image = tf.image.decode_jpeg(image, channels=3)\n        image = tf.cast(image, tf.float32) / 255.0\n        image = tf.image.resize(image, [224, 224])\n        #image = tf.image.stateless_random_flip_left_right(image, seed=(2,3))\n        #image = tf.image.random_brightness(image, 0.1)\n        return image\n\n    def read_tfrecord(self, serialized_example):\n        tfrecord_format = {\n            \"image\": tf.io.FixedLenFeature([], tf.string),\n            \"image_name\": tf.io.FixedLenFeature([], tf.string),\n        }\n        return tf.io.parse_single_example(serialized_example, tfrecord_format)\n\n    def load_and_preprocess(self, tfrecord_files):\n        raw_dataset = tf.data.TFRecordDataset(tfrecord_files, num_parallel_reads=tf.data.AUTOTUNE)\n        parsed_dataset = raw_dataset.map(self.read_tfrecord)\n\n        images = []\n        image_ids = []\n        for parsed_record in parsed_dataset:\n            image = self.decode_image(parsed_record['image'])\n            idnum = parsed_record['image_name'].numpy().decode('utf-8')\n            images.append(image)\n            image_ids.append(idnum)\n        return images, image_ids\n\n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        image = self.images[idx]\n        inputs = self.image_processor(images=image, return_tensors=\"pt\", do_resize=False, do_rescale=False)\n        pixel_values = inputs[\"pixel_values\"].squeeze(0)\n        image_id = self.image_ids[idx]\n        return pixel_values, image_id\n\n    \nclass ViTForImageClassification(torch.nn.Module):\n    def __init__(self, num_labels=5):\n        super(ViTForImageClassification, self).__init__()\n        self.vit = ViTModel.from_pretrained('/kaggle/input/google-vit/google_vit', local_files_only=True)\n        self.dropout = torch.nn.Dropout(0.1)\n        self.classifier = torch.nn.Linear(self.vit.config.hidden_size, num_labels)\n\n    def forward(self, pixel_values, labels=None):\n        outputs = self.vit(pixel_values=pixel_values)\n        output = self.dropout(outputs.last_hidden_state[:,0])\n        logits = self.classifier(output)\n        \n        loss = None\n        if labels is not None:\n            loss_fct = nn.CrossEntropyLoss()\n            loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1))\n        if loss is not None:\n            return logits, loss\n        else:\n            return logits, None\n\n# Load the model and its weights\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel_path = '/kaggle/input/update/new_v6.pth'  # Specify the path to your .pth file , currently v2 or v5 ishould give the best results\nmodel = ViTForImageClassification().to(device)\nmodel.load_state_dict(torch.load(model_path, map_location=device))\nmodel.eval()\n\n#GCS_PATH = KaggleDatasets().get_gcs_path('cassava-leaf-disease-classification')\nTEST_FILENAMES = tf.io.gfile.glob('/kaggle/input/cassava-leaf-disease-classification' + '/test_tfrecords/ld_test*.tfrec')\n\nimage_processor = ViTImageProcessor.from_pretrained('/kaggle/input/google-vit/google_vit', local_files_only=True)\ntfrecord_files = TEST_FILENAMES\ntest_dataset = CassavaLeafTestDataset(tfrecord_files, image_processor=image_processor)\ntest_dataloader = DataLoader(test_dataset, batch_size=10, shuffle=False)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-25T19:50:35.173644Z","iopub.execute_input":"2024-04-25T19:50:35.173955Z","iopub.status.idle":"2024-04-25T19:51:15.084124Z","shell.execute_reply.started":"2024-04-25T19:50:35.173926Z","shell.execute_reply":"2024-04-25T19:51:15.083300Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = []\nimage_ids = []\n\n# Predict\nwith torch.no_grad():\n    for pixel_values, ids in test_dataloader:\n        pixel_values = pixel_values.to(device)  # Ensure pixel values are on the correct device\n        logits, _ = model(pixel_values, None)\n        preds = torch.argmax(logits, dim=1).cpu().numpy()\n        \n        predictions.extend(preds)\n        image_ids.extend(ids)\n\n# Prepare the submission DataFrame\nsubmission_df = pd.DataFrame({\n    'image_id': image_ids,\n    'label': predictions\n})\nsubmission_df.to_csv('/kaggle/working/submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-04-23T12:13:33.096851Z","iopub.execute_input":"2024-04-23T12:13:33.097203Z","iopub.status.idle":"2024-04-23T12:13:33.888422Z","shell.execute_reply.started":"2024-04-23T12:13:33.097175Z","shell.execute_reply":"2024-04-23T12:13:33.887409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head submission.csv","metadata":{"execution":{"iopub.status.busy":"2024-04-23T12:13:36.420258Z","iopub.execute_input":"2024-04-23T12:13:36.420955Z","iopub.status.idle":"2024-04-23T12:13:37.419619Z","shell.execute_reply.started":"2024-04-23T12:13:36.420921Z","shell.execute_reply":"2024-04-23T12:13:37.418349Z"},"trusted":true},"execution_count":null,"outputs":[]}]}