{"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":"none","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ✅ STEP 1: Install Required Libraries\n!pip install -q transformers timm albumentations\n\n# ✅ STEP 2: Imports\nimport os\nimport cv2\nimport timm\nimport torch\nimport random\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport albumentations as A\nimport matplotlib.pyplot as plt\n\nfrom tqdm import tqdm\nfrom sklearn.metrics import classification_report, confusion_matrix\nfrom sklearn.model_selection import StratifiedKFold\nfrom torch.utils.data import Dataset, DataLoader\nfrom transformers import get_cosine_schedule_with_warmup\nfrom torchvision import transforms as T\nfrom torch import nn\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(\"Running on:\", device)\n\n# ✅ STEP 3: Load Data\ndf = pd.read_csv('/kaggle/input/cassava-leaf-disease-classification/train.csv')\ndf['image_path'] = '/kaggle/input/cassava-leaf-disease-classification/train_images/' + df['image_id']\nlabels_map = {\n    0: 'Healthy',\n    1: 'Cassava Bacterial Blight',\n    2: 'Cassava Brown Streak Disease',\n    3: 'Cassava Green Mottle',\n    4: 'Cassava Mosaic Disease'\n}\ndf['label_name'] = df['label'].map(labels_map)\ndf.head()\n\n# ✅ STEP 4: Data Visualization\nsns.countplot(df['label'])\nplt.title(\"Class Distribution\")\nplt.show()\n\n# ✅ STEP 5: Dataset Class\nclass CassavaDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        row = self.df.loc[index]\n        img = cv2.imread(row.image_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        if self.transform:\n            img = self.transform(image=img)['image']\n        label = row.label\n        return img, label\n\n# ✅ STEP 6: Augmentation\ntrain_transforms = A.Compose([\n    A.RandomResizedCrop(224, 224),\n    A.HorizontalFlip(),\n    A.VerticalFlip(),\n    A.ShiftScaleRotate(),\n    A.Normalize(),\n    A.pytorch.ToTensorV2()\n])\n\nval_transforms = A.Compose([\n    A.Resize(224, 224),\n    A.Normalize(),\n    A.pytorch.ToTensorV2()\n])\n\n# ✅ STEP 7: Swin Transformer Model\nclass SwinClassifier(nn.Module):\n    def __init__(self, num_classes):\n        super(SwinClassifier, self).__init__()\n        self.model = timm.create_model('swin_base_patch4_window7_224', pretrained=True)\n        in_features = self.model.head.in_features\n        self.model.head = nn.Linear(in_features, num_classes)\n\n    def forward(self, x):\n        return self.model(x)\n\n# ✅ STEP 8: Training Utilities\ndef train_one_epoch(model, loader, optimizer, criterion):\n    model.train()\n    running_loss = 0.0\n    for imgs, labels in tqdm(loader):\n        imgs, labels = imgs.to(device), labels.to(device)\n        optimizer.zero_grad()\n        preds = model(imgs)\n        loss = criterion(preds, labels)\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item()\n    return running_loss / len(loader)\n\ndef validate(model, loader, criterion):\n    model.eval()\n    all_preds = []\n    all_labels = []\n    val_loss = 0.0\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs, labels = imgs.to(device), labels.to(device)\n            preds = model(imgs)\n            loss = criterion(preds, labels)\n            val_loss += loss.item()\n            all_preds.extend(torch.argmax(preds, 1).cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n    return val_loss / len(loader), all_preds, all_labels\n\n# ✅ STEP 9: Stratified K-Fold\nfolds = StratifiedKFold(n_splits=5, shuffle=True, random_state=42)\n\nfor fold, (train_idx, val_idx) in enumerate(folds.split(df, df['label'])):\n    print(f\"\\n📁 Fold {fold+1}\")\n    \n    train_df = df.iloc[train_idx]\n    val_df = df.iloc[val_idx]\n\n    train_dataset = CassavaDataset(train_df, transform=train_transforms)\n    val_dataset = CassavaDataset(val_df, transform=val_transforms)\n\n    train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)\n    val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False)\n\n    model = SwinClassifier(num_classes=5).to(device)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)\n    criterion = nn.CrossEntropyLoss()\n\n    best_val_loss = np.inf\n\n    for epoch in range(10):\n        print(f\"\\nEpoch {epoch+1}\")\n        train_loss = train_one_epoch(model, train_loader, optimizer, criterion)\n        val_loss, val_preds, val_labels = validate(model, val_loader, criterion)\n\n        print(f\"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}\")\n        print(classification_report(val_labels, val_preds))\n\n        if val_loss < best_val_loss:\n            torch.save(model.state_dict(), f\"swin_fold{fold+1}.pt\")\n            best_val_loss = val_loss\n\n    # Only one fold for now to keep runtime low\n    break\n\n# ✅ STEP 10: Confusion Matrix\nconf_mat = confusion_matrix(val_labels, val_preds)\nplt.figure(figsize=(8,6))\nsns.heatmap(conf_mat, annot=True, fmt='d', xticklabels=labels_map.values(), yticklabels=labels_map.values(), cmap='Blues')\nplt.title(\"Confusion Matrix\")\nplt.show()\n\n# ✅ STEP 11: Inference for Submission (Optional)\ntest_dir = '/kaggle/input/cassava-leaf-disease-classification/test_images/'\ntest_files = os.listdir(test_dir)\n\ntest_df = pd.DataFrame()\ntest_df['image_id'] = test_files\ntest_df['image_path'] = test_dir + test_df['image_id']\n\nclass TestDataset(Dataset):\n    def __init__(self, df, transform):\n        self.df = df\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        img = cv2.imread(self.df.iloc[index]['image_path'])\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img = self.transform(image=img)['image']\n        return img\n\ntest_ds = TestDataset(test_df, val_transforms)\ntest_loader = DataLoader(test_ds, batch_size=32, shuffle=False)\n\nmodel.load_state_dict(torch.load(\"swin_fold1.pt\"))\nmodel.eval()\n\nfinal_preds = []\nwith torch.no_grad():\n    for imgs in test_loader:\n        imgs = imgs.to(device)\n        outputs = model(imgs)\n        preds = torch.argmax(outputs, 1).cpu().numpy()\n        final_preds.extend(preds)\n\n# ✅ STEP 12: Save Submission\nsubmission = pd.read_csv('/kaggle/input/cassava-leaf-disease-classification/sample_submission.csv')\nsubmission['label'] = final_preds\nsubmission.to_csv('submission.csv', index=False)\n\nprint(\"✅ Submission saved.\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}