{"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":8018622,"sourceType":"datasetVersion","datasetId":4724693},{"sourceId":8233730,"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/pretrain/pretrain', 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_path1 = '/kaggle/input/update/new_v2.pth' \nmodel_path2 = '/kaggle/input/update/new_v3.pth'\nmodel_path3 = '/kaggle/input/update/new_v4.pth'\nmodel_path4 = '/kaggle/input/update/new_v5.pth'\nmodel_path5 = '/kaggle/input/update/new_v6.pth'# Specify the path to your .pth file\nmodel1 = ViTForImageClassification().to(device)\nmodel1.load_state_dict(torch.load(model_path1, map_location=device))\nmodel2 = ViTForImageClassification().to(device)\nmodel2.load_state_dict(torch.load(model_path2, map_location=device))\nmodel3 = ViTForImageClassification().to(device)\nmodel3.load_state_dict(torch.load(model_path3, map_location=device))\nmodel4 = ViTForImageClassification().to(device)\nmodel4.load_state_dict(torch.load(model_path4, map_location=device))\nmodel5 = ViTForImageClassification().to(device)\nmodel5.load_state_dict(torch.load(model_path5, map_location=device))\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/pretrain/pretrain', local_files_only=True)\ntfrecord_files = TEST_FILENAMES\ntest_dataset = CassavaLeafTestDataset(tfrecord_files, image_processor=image_processor)\ntest_dataloader = DataLoader(test_dataset, batch_size=16, shuffle=False)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-26T09:46:10.368989Z","iopub.execute_input":"2024-04-26T09:46:10.369291Z","iopub.status.idle":"2024-04-26T09:46:49.097829Z","shell.execute_reply.started":"2024-04-26T09:46:10.369263Z","shell.execute_reply":"2024-04-26T09:46:49.096974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = []\nimage_ids = []\nmodels = [model1,model2,model3,model4,model5]\n# Predict\n# with 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)\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\n        # Initialize a tensor to accumulate logits from all models\n        accumulated_logits = None\n\n        for model in models:\n            logits, _ = model(pixel_values, None)\n            if accumulated_logits is None:\n                accumulated_logits = logits\n            else:\n                accumulated_logits += logits\n\n        # Average the accumulated logits\n        averaged_logits = accumulated_logits / len(models)\n\n        # Convert logits to actual predictions\n        preds = torch.argmax(averaged_logits, dim=1).cpu().numpy()\n\n        predictions.extend(preds)\n        image_ids.extend(ids)\n        \n        \n\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-26T09:47:12.963883Z","iopub.execute_input":"2024-04-26T09:47:12.964283Z","iopub.status.idle":"2024-04-26T09:47:13.664646Z","shell.execute_reply.started":"2024-04-26T09:47:12.964252Z","shell.execute_reply":"2024-04-26T09:47:13.663800Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head submission.csv","metadata":{"execution":{"iopub.status.busy":"2024-04-26T09:47:16.303515Z","iopub.execute_input":"2024-04-26T09:47:16.303878Z","iopub.status.idle":"2024-04-26T09:47:17.358116Z","shell.execute_reply.started":"2024-04-26T09:47:16.303850Z","shell.execute_reply":"2024-04-26T09:47:17.356995Z"},"trusted":true},"execution_count":null,"outputs":[]}]}