{"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"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"},{"sourceId":656417,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":496135,"modelId":511540}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.amp import autocast\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm.notebook import tqdm\nimport timm\n\nclass CFG:\n    img_size = 384\n    batch_size = 64\n    num_workers = 4\n    device = 'cuda'\n    test_dir = '/kaggle/input/cassava-leaf-disease-classification/test_images'\n    model_paths = [\n        '/kaggle/input/cassava-convnext-tiny/pytorch/default/1/best_fold0.pth',\n        '/kaggle/input/cassava-convnext-tiny/pytorch/default/1/best_fold1.pth',\n        '/kaggle/input/cassava-convnext-tiny/pytorch/default/1/best_fold2.pth',\n        '/kaggle/input/cassava-convnext-tiny/pytorch/default/1/best_fold3.pth',\n        '/kaggle/input/cassava-convnext-tiny/pytorch/default/1/best_fold4.pth',\n    ]\n\ntest_tfms = A.Compose([\n    A.Resize(CFG.img_size, CFG.img_size),\n    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    ToTensorV2()\n])\n\nclass TestDataset(Dataset):\n    def __init__(self, folder):\n        self.paths = sorted([str(p) for p in Path(folder).glob(\"*.jpg\")])\n    def __len__(self): return len(self.paths)\n    def __getitem__(self, idx):\n        img = cv2.imread(self.paths[idx])\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img = test_tfms(image=img)['image']\n        return img, os.path.basename(self.paths[idx])\n\nclass CassavaModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = timm.create_model('convnext_tiny', pretrained=False, num_classes=5)\n    def forward(self, x):\n        return self.backbone(x)\n\n@torch.no_grad()\ndef inference():\n    dataset = TestDataset(CFG.test_dir)\n    loader = DataLoader(dataset, batch_size=CFG.batch_size,\n                        shuffle=False, num_workers=CFG.num_workers, pin_memory=True)\n\n    ensemble_preds = None\n\n    for fold, path in enumerate(CFG.model_paths):\n        print(f\"Loading fold {fold} → {os.path.basename(path)}\")\n        model = CassavaModel().to(CFG.device)\n        \n        state = torch.load(path, map_location=CFG.device)\n        model.load_state_dict(state)        \n        model.eval()\n\n        fold_preds = []\n        for imgs, _ in tqdm(loader, leave=False, desc=f\"Fold {fold} TTA\"):\n            imgs = imgs.to(CFG.device)\n            with autocast(device_type='cuda'):\n                p1 = torch.softmax(model(imgs), dim=1)\n                p2 = torch.softmax(model(torch.flip(imgs, dims=[3])), dim=1)  \n            fold_preds.append(((p1 + p2) / 2).cpu().numpy())\n\n        fold_preds = np.concatenate(fold_preds)\n        ensemble_preds = fold_preds if ensemble_preds is None else ensemble_preds + fold_preds\n        \n        del model, state\n        torch.cuda.empty_cache()\n\n    final_labels = np.argmax(ensemble_preds / len(CFG.model_paths), axis=1)\n\n    sub = pd.DataFrame({\n        'image_id': [os.path.basename(p) for p in dataset.paths],\n        'label': final_labels\n    })\n    sub = sub.sort_values('image_id').reset_index(drop=True)\n    sub.to_csv('submission.csv', index=False)\n    \n    print(f\"\\n{len(sub)} predictions\")\n    print(sub.head())\n    return sub\n\ninference()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}