{"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":1709588,"sourceType":"datasetVersion","datasetId":1013614},{"sourceId":8009334,"sourceType":"datasetVersion","datasetId":4717506},{"sourceId":8018537,"sourceType":"datasetVersion","datasetId":4724628},{"sourceId":8018622,"sourceType":"datasetVersion","datasetId":4724693},{"sourceId":8231648,"sourceType":"datasetVersion","datasetId":4752193}],"dockerImageVersionId":30674,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# !pip install vision_transformer_pytorch\nimport sys\npackage_path = '../input/vision-transformer-pytorch/VisionTransformer-Pytorch'\nsys.path.append(package_path)","metadata":{"execution":{"iopub.status.busy":"2024-04-26T01:36:18.594968Z","iopub.execute_input":"2024-04-26T01:36:18.595596Z","iopub.status.idle":"2024-04-26T01:36:18.602174Z","shell.execute_reply.started":"2024-04-26T01:36:18.595563Z","shell.execute_reply":"2024-04-26T01:36:18.601423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom vision_transformer_pytorch import VisionTransformer\n\nclass CassavaLeafTestDataset(Dataset):\n    def __init__(self, tfrecord_files,transform=None):\n        self.tfrecord_files = tfrecord_files\n        self.transform = transform\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 = image.numpy()\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.experimental.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        if self.transform:\n            tf = self.transform(image = image)\n            image = tf['image']\n        image_id = self.image_ids[idx]\n        return image, image_id\n\n    \nmodel = VisionTransformer.from_name('ViT-B_16', num_classes=5)\nmodel.load_state_dict(torch.load('/kaggle/input/update/model_test.pt'))\n\ntransform = A.Compose([\n#             albu.CenterCrop(512, 512, p=1.),\n            A.Resize(384, 384),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                max_pixel_value=255.0, \n                p=1.0),\n            ToTensorV2()], p=1.)\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nTEST_FILENAMES = tf.io.gfile.glob('/kaggle/input/cassava-leaf-disease-classification' + '/test_tfrecords/ld_test*.tfrec')\ntfrecord_files = TEST_FILENAMES\ntest_dataset = CassavaLeafTestDataset(tfrecord_files, transform=transform)\ntest_dataloader = DataLoader(test_dataset, batch_size=4, shuffle=False)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-26T01:36:20.470601Z","iopub.execute_input":"2024-04-26T01:36:20.471333Z","iopub.status.idle":"2024-04-26T01:36:48.346820Z","shell.execute_reply.started":"2024-04-26T01:36:20.471303Z","shell.execute_reply":"2024-04-26T01:36:48.345979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions=[]\n\nfor imgs in test_dataloader:\n    imgs = imgs[0]\n    imgs =imgs.to(device)\n    with torch.no_grad():\n        model=model.to(device)\n        outputs = model(imgs)\n        _, predicted = torch.max(outputs, dim=1)\n        predicted=predicted.to('cpu')\n        predictions.append(predicted)","metadata":{"execution":{"iopub.status.busy":"2024-04-26T01:36:52.424717Z","iopub.execute_input":"2024-04-26T01:36:52.425520Z","iopub.status.idle":"2024-04-26T01:36:53.413071Z","shell.execute_reply.started":"2024-04-26T01:36:52.425489Z","shell.execute_reply":"2024-04-26T01:36:53.412223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(\"../input/cassava-leaf-disease-classification/sample_submission.csv\")\ntest_df['label'] = np.concatenate(predictions)","metadata":{"execution":{"iopub.status.busy":"2024-04-26T01:36:55.456118Z","iopub.execute_input":"2024-04-26T01:36:55.456521Z","iopub.status.idle":"2024-04-26T01:36:55.485140Z","shell.execute_reply.started":"2024-04-26T01:36:55.456491Z","shell.execute_reply":"2024-04-26T01:36:55.484199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-04-26T01:37:00.054010Z","iopub.execute_input":"2024-04-26T01:37:00.054856Z","iopub.status.idle":"2024-04-26T01:37:00.060347Z","shell.execute_reply.started":"2024-04-26T01:37:00.054825Z","shell.execute_reply":"2024-04-26T01:37:00.059269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head submission.csv","metadata":{"execution":{"iopub.status.busy":"2024-04-26T01:37:01.474825Z","iopub.execute_input":"2024-04-26T01:37:01.475973Z","iopub.status.idle":"2024-04-26T01:37:02.491714Z","shell.execute_reply.started":"2024-04-26T01:37:01.475925Z","shell.execute_reply":"2024-04-26T01:37:02.489994Z"},"trusted":true},"execution_count":null,"outputs":[]}]}