{"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":8100460,"sourceType":"datasetVersion","datasetId":4767531},{"sourceId":8100436,"sourceType":"datasetVersion","datasetId":4783604},{"sourceId":8087245,"sourceType":"datasetVersion","datasetId":4774039},{"sourceId":8105897,"sourceType":"datasetVersion","datasetId":4787525}],"dockerImageVersionId":30674,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n\nimport pandas as pd\nimport torch\nfrom PIL import Image\nfrom torch.backends import cudnn\nfrom torch.utils.data import DataLoader\nfrom torchvision.datasets import VisionDataset\nfrom torchvision.transforms import InterpolationMode, v2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-11T04:03:13.522467Z","iopub.execute_input":"2024-04-11T04:03:13.523163Z","iopub.status.idle":"2024-04-11T04:03:13.528673Z","shell.execute_reply.started":"2024-04-11T04:03:13.523122Z","shell.execute_reply":"2024-04-11T04:03:13.527584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Basic experiment settings\ntorch.manual_seed(3407)\ntorch.cuda.manual_seed(3407)\n\n\ncudnn.deterministic = True\ncudnn.benchmark = False\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)\n\n# Declare directories\ntest_dir = \"/kaggle/input/cassava-leaf-disease-classification/test_images/\"\n\n# Experiment parameters\n# img_size = 528\nimg_size = 384\nbatch_size = 16\nnum_workers = 4\nnum_classes = 5\ntta = True\n\n# Currently testing vision transformers so check its architecture\nmodel = torch.load(\"/kaggle/input/vit-v6/vit_v6.pt\", map_location=device)","metadata":{"execution":{"iopub.status.busy":"2024-04-11T04:03:13.752847Z","iopub.execute_input":"2024-04-11T04:03:13.753180Z","iopub.status.idle":"2024-04-11T04:03:14.103043Z","shell.execute_reply.started":"2024-04-11T04:03:13.753154Z","shell.execute_reply":"2024-04-11T04:03:14.101775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Test dataset without labels\nclass CassavaDataset(VisionDataset):\n    \"\"\"Custom dataset for the Cassava data.\n\n    Args:\n        data_dir: base directory to the images.\n        img_label_df: dataframe containing image-label pairs.\n        transforms: set of transforms to be used.\n    \"\"\"\n\n    def __init__(self, data_dir, transform=None, ttas=None, img_size=384):\n        super().__init__(root=data_dir)\n\n        self.transform = transform\n        self.images = os.listdir(data_dir)\n        self.ttas = ttas\n        self.cc = v2.CenterCrop((600, 600))\n\n    def __getitem__(self, idx):\n        filename = self.images[idx]\n        img = Image.open(os.path.join(self.root, filename))\n        img = self.cc(img)\n        \n        if self.ttas is not None and self.transform is not None:\n            img = [self.transform(t(img)) for t in self.ttas]\n        elif self.transform:\n            img = self.transform(img)\n\n        return img, filename\n\n\n    def __len__(self):\n        return len(self.images)","metadata":{"execution":{"iopub.status.busy":"2024-04-11T04:03:14.105381Z","iopub.execute_input":"2024-04-11T04:03:14.106263Z","iopub.status.idle":"2024-04-11T04:03:14.115989Z","shell.execute_reply.started":"2024-04-11T04:03:14.106228Z","shell.execute_reply":"2024-04-11T04:03:14.115040Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create the set of transforms for test time\ntest_transforms = v2.Compose(\n    [\n        v2.Resize((img_size, img_size), interpolation=InterpolationMode.BICUBIC),\n        v2.ToImage(),\n        v2.ToDtype(torch.float32, scale=True),\n        v2.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    ]\n)\n\nif tta:\n    ttas = [\n#         v2.RandomResizedCrop((img_size, img_size), (0.5, 1)),\n        v2.RandomRotation(180),\n        v2.RandomVerticalFlip(1),\n#         v2.RandomAffine(180),\n        v2.RandomPerspective(p=1),\n#         lambda x: x\n    ]\nelse:\n    ttas = None\n\n# Dataset and dataloader\ntest_dataset = CassavaDataset(test_dir, transform=test_transforms, ttas=ttas)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=batch_size,\n    shuffle=True,\n    num_workers=num_workers,\n    pin_memory=True,\n)\n\n# Adjust unnormalized outputs\nnormalizer = torch.nn.Softmax(dim=1)","metadata":{"execution":{"iopub.status.busy":"2024-04-11T04:03:14.164205Z","iopub.execute_input":"2024-04-11T04:03:14.164764Z","iopub.status.idle":"2024-04-11T04:03:14.196099Z","shell.execute_reply.started":"2024-04-11T04:03:14.164735Z","shell.execute_reply":"2024-04-11T04:03:14.195272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Perform inference\nall_names = []\nall_preds = []\n\nmodel.eval()\n\nwith torch.no_grad():\n    test_loss = 0\n\n    for batch_idx, (inputs, filenames) in enumerate(test_loader):\n        batch_size = len(filenames)\n        \n        if tta:\n            inputs, filenames = torch.cat(inputs, dim=0).to(device), list(filenames)\n\n            # Now do the input predictions\n            preds = normalizer(model(inputs))\n\n            # Find the average prediction by each model\n            batch_preds = torch.stack(torch.split(preds, len(filenames)), dim=0)\n            mean_preds = torch.mean(batch_preds, dim=0)\n            pred_labels = torch.argmax(mean_preds, 1).tolist()\n            \n        else:\n            # The same thing but without the troublesome stuff\n            inputs, filenames = inputs.to(device), list(filenames)\n            \n            preds = normalizer(model(inputs))\n            pred_labels = torch.argmax(preds, 1).tolist()\n            \n        all_names.extend(filenames)\n        all_preds.extend(pred_labels)","metadata":{"execution":{"iopub.status.busy":"2024-04-11T04:03:14.389433Z","iopub.execute_input":"2024-04-11T04:03:14.389822Z","iopub.status.idle":"2024-04-11T04:03:14.858149Z","shell.execute_reply.started":"2024-04-11T04:03:14.389790Z","shell.execute_reply":"2024-04-11T04:03:14.856879Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create submission\nmy_submission = pd.DataFrame({\"image_id\": all_names, \"label\": all_preds})\nmy_submission.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2024-04-11T04:03:14.860367Z","iopub.execute_input":"2024-04-11T04:03:14.860693Z","iopub.status.idle":"2024-04-11T04:03:14.868725Z","shell.execute_reply.started":"2024-04-11T04:03:14.860663Z","shell.execute_reply":"2024-04-11T04:03:14.867883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"my_submission","metadata":{"execution":{"iopub.status.busy":"2024-04-11T04:03:15.721992Z","iopub.execute_input":"2024-04-11T04:03:15.722389Z","iopub.status.idle":"2024-04-11T04:03:15.733078Z","shell.execute_reply.started":"2024-04-11T04:03:15.722363Z","shell.execute_reply":"2024-04-11T04:03:15.732052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}