{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":25563,"databundleVersionId":2094376}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Apple Leaf Disease Classification","metadata":{}},{"cell_type":"code","source":"# ── Cell 1: Imports ──────────────────────────────────────────────────────────\nimport os, random, warnings\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport cv2\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import f1_score, roc_auc_score, classification_report\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\n\nwarnings.filterwarnings('ignore')\nprint('Imports OK')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T12:19:05.314369Z","iopub.execute_input":"2026-04-11T12:19:05.315141Z","iopub.status.idle":"2026-04-11T12:19:16.495297Z","shell.execute_reply.started":"2026-04-11T12:19:05.315111Z","shell.execute_reply":"2026-04-11T12:19:16.494662Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Cell 2: Config ───────────────────────────────────────────────────────────\nclass CFG:\n    # Paths — 2021 dataset only\n    CSV_PATH  = '/kaggle/input/competitions/plant-pathology-2021-fgvc8/train.csv'\n    IMG_DIR   = '/kaggle/input/competitions/plant-pathology-2021-fgvc8/train_images'\n    SAVE_PATH = 'best_model_v3.pth'\n\n    IMG_SIZE    = 512\n    VAL_SPLIT   = 0.2\n    NUM_CLASSES = 4\n    EPOCHS      = 50\n    BATCH_SIZE  = 16\n    NUM_WORKERS = 2\n    SEED        = 42\n\n    LR              = 1e-4\n    WEIGHT_DECAY    = 1e-2\n    LABEL_SMOOTHING = 0.1\n    MIXUP_ALPHA     = 0.4\n    MULTI_DISEASE_BOOST  = 2.0   # sampler boost for multiple_diseases\n    MULTI_DISEASE_THRESH = 0.50  # prediction threshold\n\n    CLASSES = ['healthy', 'scab', 'rust', 'multiple_diseases']\n\n    # 2021 → your 4-class mapping (drop everything else)\n    LABEL_MAP = {\n        'healthy': 'healthy',\n        'scab':    'scab',\n        'rust':    'rust',\n        'complex': 'multiple_diseases',\n    }\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\ndef seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_everything(CFG.SEED)\nprint(f'Device: {DEVICE}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T12:24:47.606451Z","iopub.execute_input":"2026-04-11T12:24:47.606764Z","iopub.status.idle":"2026-04-11T12:24:47.616765Z","shell.execute_reply.started":"2026-04-11T12:24:47.606735Z","shell.execute_reply":"2026-04-11T12:24:47.616026Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Cell 3: Load & Filter 2021 Dataset ───────────────────────────────────────\ndf = pd.read_csv(CFG.CSV_PATH)\nprint(f'Raw 2021 dataset: {len(df)} rows')\nprint(df['labels'].value_counts(), '\\n')\n\n# Keep only the 4 clean single-label classes\ndf = df[df['labels'].isin(CFG.LABEL_MAP.keys())].copy()\ndf['label_name'] = df['labels'].map(CFG.LABEL_MAP)\ndf['label']      = df['label_name'].map({c: i for i, c in enumerate(CFG.CLASSES)})\ndf = df.reset_index(drop=True)\n\nprint(f'Filtered dataset: {len(df)} rows')\nprint(df['label_name'].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T12:24:52.448168Z","iopub.execute_input":"2026-04-11T12:24:52.448948Z","iopub.status.idle":"2026-04-11T12:24:52.486330Z","shell.execute_reply.started":"2026-04-11T12:24:52.448918Z","shell.execute_reply":"2026-04-11T12:24:52.485555Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Cell 4: Train / Val Split ─────────────────────────────────────────────────\ntrain_df, val_df = train_test_split(\n    df,\n    test_size    = CFG.VAL_SPLIT,\n    stratify     = df['label'],\n    random_state = CFG.SEED,\n)\ntrain_df = train_df.reset_index(drop=True)\nval_df   = val_df.reset_index(drop=True)\n\nprint(f'Train: {len(train_df)}  |  Val: {len(val_df)}')\nprint('Train class distribution:')\nprint(train_df['label_name'].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T12:24:54.863710Z","iopub.execute_input":"2026-04-11T12:24:54.864196Z","iopub.status.idle":"2026-04-11T12:24:54.879606Z","shell.execute_reply.started":"2026-04-11T12:24:54.864167Z","shell.execute_reply":"2026-04-11T12:24:54.878869Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Cell 5: Transforms ───────────────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_transform = A.Compose([\n    A.LongestMaxSize(max_size=CFG.IMG_SIZE),\n    A.PadIfNeeded(\n        min_height=CFG.IMG_SIZE, min_width=CFG.IMG_SIZE,\n        border_mode=0, fill=0, position='random', p=1.0,\n    ),\n    A.RandomCrop(height=CFG.IMG_SIZE, width=CFG.IMG_SIZE),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.Sharpen(alpha=(0.2, 0.5), lightness=(0.5, 1.0), p=0.3),\n    A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.3),\n    A.ElasticTransform(alpha=120, sigma=120 * 0.05, p=0.2),\n    A.Affine(translate_percent=0.1, scale=(0.9, 1.1), rotate=(-30, 30), p=0.5),\n    A.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1, p=0.5),\n    A.RandomBrightnessContrast(p=0.3),\n    A.GaussNoise(p=0.2),\n    A.MotionBlur(blur_limit=5, p=0.2),\n    A.CoarseDropout(\n        num_holes_range=(2, 8),\n        hole_height_range=(CFG.IMG_SIZE // 8, CFG.IMG_SIZE // 4),\n        hole_width_range=(CFG.IMG_SIZE // 8, CFG.IMG_SIZE // 4),\n        fill=0, p=0.5,\n    ),\n    A.Normalize(mean=MEAN, std=STD),\n    ToTensorV2(),\n])\n\nval_transform = A.Compose([\n    A.LongestMaxSize(max_size=CFG.IMG_SIZE),\n    A.PadIfNeeded(\n        min_height=CFG.IMG_SIZE, min_width=CFG.IMG_SIZE,\n        border_mode=0, fill=0, position='center', p=1.0,\n    ),\n    A.CenterCrop(height=CFG.IMG_SIZE, width=CFG.IMG_SIZE),\n    A.Normalize(mean=MEAN, std=STD),\n    ToTensorV2(),\n])\n\nprint('Transforms OK')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T12:24:57.329942Z","iopub.execute_input":"2026-04-11T12:24:57.330264Z","iopub.status.idle":"2026-04-11T12:24:57.349019Z","shell.execute_reply.started":"2026-04-11T12:24:57.330207Z","shell.execute_reply":"2026-04-11T12:24:57.348424Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Cell 6: Dataset ───────────────────────────────────────────────────────────\nclass LeafDataset(Dataset):\n    def __init__(self, df, img_dir=CFG.IMG_DIR, transform=None):\n        self.df        = df\n        self.img_dir   = img_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row      = self.df.iloc[idx]\n        img_path = os.path.join(self.img_dir, row['image'])\n        image    = cv2.imread(img_path)\n        image    = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        if self.transform:\n            image = self.transform(image=image)['image']\n        label = torch.tensor(row['label'], dtype=torch.long)\n        return image, label\n\nprint('Dataset class OK')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T12:25:00.146245Z","iopub.execute_input":"2026-04-11T12:25:00.146543Z","iopub.status.idle":"2026-04-11T12:25:00.152709Z","shell.execute_reply.started":"2026-04-11T12:25:00.146518Z","shell.execute_reply":"2026-04-11T12:25:00.152137Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Cell 7: DataLoaders ───────────────────────────────────────────────────────\ntrain_dataset = LeafDataset(train_df, transform=train_transform)\nval_dataset   = LeafDataset(val_df,   transform=val_transform)\n\n# Weighted sampler to handle class imbalance\nclass_counts      = train_df['label'].value_counts().sort_index().values\nclass_weights_s   = 1.0 / class_counts.astype(float)\nclass_weights_s[3] *= CFG.MULTI_DISEASE_BOOST\nsample_weights    = class_weights_s[train_df['label'].values]\n\nprint('Sampling weights per class:')\nfor i, cls in enumerate(CFG.CLASSES):\n    print(f'  {cls:<20} {class_weights_s[i]:.6f}')\n\nsampler = WeightedRandomSampler(\n    weights     = sample_weights,\n    num_samples = len(train_dataset),\n    replacement = True,\n)\ntrain_loader = DataLoader(\n    train_dataset, batch_size=CFG.BATCH_SIZE,\n    sampler=sampler, num_workers=CFG.NUM_WORKERS, pin_memory=True,\n)\nval_loader = DataLoader(\n    val_dataset, batch_size=CFG.BATCH_SIZE,\n    shuffle=False, num_workers=CFG.NUM_WORKERS, pin_memory=True,\n)\nprint(f'Train batches: {len(train_loader)}  |  Val batches: {len(val_loader)}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T12:25:04.861691Z","iopub.execute_input":"2026-04-11T12:25:04.862386Z","iopub.status.idle":"2026-04-11T12:25:04.870572Z","shell.execute_reply.started":"2026-04-11T12:25:04.862356Z","shell.execute_reply":"2026-04-11T12:25:04.869852Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Cell 8: MixUp & CutMix ───────────────────────────────────────────────────\ndef mixup_batch(images, labels, alpha=0.4):\n    lam   = np.random.beta(alpha, alpha)\n    idx   = torch.randperm(images.size(0), device=images.device)\n    mixed = lam * images + (1 - lam) * images[idx]\n    return mixed, labels, labels[idx], lam\n\ndef cutmix_batch(images, labels, alpha=1.0):\n    lam   = np.random.beta(alpha, alpha)\n    idx   = torch.randperm(images.size(0), device=images.device)\n    W, H  = images.shape[3], images.shape[2]\n    cut_w = int(W * np.sqrt(1 - lam))\n    cut_h = int(H * np.sqrt(1 - lam))\n    cx, cy = np.random.randint(W), np.random.randint(H)\n    x1 = np.clip(cx - cut_w // 2, 0, W)\n    x2 = np.clip(cx + cut_w // 2, 0, W)\n    y1 = np.clip(cy - cut_h // 2, 0, H)\n    y2 = np.clip(cy + cut_h // 2, 0, H)\n    mixed = images.clone()\n    mixed[:, :, y1:y2, x1:x2] = images[idx, :, y1:y2, x1:x2]\n    lam_actual = 1 - (x2 - x1) * (y2 - y1) / (W * H)\n    return mixed, labels, labels[idx], lam_actual\n\ndef augment_batch(images, labels):\n    if random.random() < 0.5:\n        return mixup_batch(images, labels, alpha=CFG.MIXUP_ALPHA)\n    else:\n        return cutmix_batch(images, labels, alpha=1.0)\n\nprint('Augmentation functions OK')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T12:25:07.697395Z","iopub.execute_input":"2026-04-11T12:25:07.697971Z","iopub.status.idle":"2026-04-11T12:25:07.705724Z","shell.execute_reply.started":"2026-04-11T12:25:07.697942Z","shell.execute_reply":"2026-04-11T12:25:07.705023Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Cell 9: Model ────────────────────────────────────────────────────────────\nclass LeafCNN(nn.Module):\n    def __init__(self, num_classes=4, dropout=0.5):\n        super(LeafCNN, self).__init__()\n        self.block1 = nn.Sequential(\n            nn.Conv2d(3,  32, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(32), nn.ReLU(inplace=True),\n            nn.Conv2d(32, 32, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(32), nn.ReLU(inplace=True),\n            nn.MaxPool2d(2, 2),\n        )\n        self.block2 = nn.Sequential(\n            nn.Conv2d(32, 64, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(64), nn.ReLU(inplace=True),\n            nn.Conv2d(64, 64, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(64), nn.ReLU(inplace=True),\n            nn.MaxPool2d(2, 2),\n        )\n        self.block3 = nn.Sequential(\n            nn.Conv2d(64,  128, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(128), nn.ReLU(inplace=True),\n            nn.Conv2d(128, 128, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(128), nn.ReLU(inplace=True),\n            nn.MaxPool2d(2, 2),\n        )\n        self.block4 = nn.Sequential(\n            nn.Conv2d(128, 256, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(256), nn.ReLU(inplace=True),\n            nn.Conv2d(256, 256, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(256), nn.ReLU(inplace=True),\n            nn.MaxPool2d(2, 2),\n        )\n        self.block5 = nn.Sequential(\n            nn.Conv2d(256, 512, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(512), nn.ReLU(inplace=True),\n            nn.Conv2d(512, 512, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(512), nn.ReLU(inplace=True),\n            nn.MaxPool2d(2, 2),\n        )\n        self.gap = nn.AdaptiveAvgPool2d(1)\n        self.classifier = nn.Sequential(\n            nn.Dropout(p=dropout),\n            nn.Linear(512, 256), nn.ReLU(inplace=True), nn.BatchNorm1d(256),\n            nn.Dropout(p=dropout * 0.6),\n            nn.Linear(256, 128), nn.ReLU(inplace=True), nn.BatchNorm1d(128),\n            nn.Dropout(p=dropout * 0.4),\n            nn.Linear(128, num_classes),\n        )\n\n    def forward(self, x):\n        x = self.block1(x)\n        x = self.block2(x)\n        x = self.block3(x)\n        x = self.block4(x)\n        x = self.block5(x)\n        x = self.gap(x)\n        x = x.view(x.size(0), -1)\n        x = self.classifier(x)\n        return x\n\nmodel     = LeafCNN(num_classes=CFG.NUM_CLASSES, dropout=0.5).to(DEVICE)\ntotal     = sum(p.numel() for p in model.parameters())\ntrainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f'Total parameters     : {total:,}')\nprint(f'Trainable parameters : {trainable:,}')\n\ndummy = torch.zeros(2, 3, CFG.IMG_SIZE, CFG.IMG_SIZE).to(DEVICE)\nwith torch.no_grad():\n    out = model(dummy)\nprint(f'Output shape: {out.shape}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T12:25:10.485331Z","iopub.execute_input":"2026-04-11T12:25:10.485990Z","iopub.status.idle":"2026-04-11T12:25:10.553256Z","shell.execute_reply.started":"2026-04-11T12:25:10.485963Z","shell.execute_reply":"2026-04-11T12:25:10.552498Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Cell 10: Loss, Optimizer, Scheduler ──────────────────────────────────────\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma=2.0, label_smoothing=0.0):\n        super().__init__()\n        self.gamma           = gamma\n        self.label_smoothing = label_smoothing\n\n    def forward(self, logits, targets):\n        ce   = F.cross_entropy(logits, targets,\n                               reduction='none',\n                               label_smoothing=self.label_smoothing)\n        pt   = torch.exp(-ce)\n        loss = ((1 - pt) ** self.gamma) * ce\n        return loss.mean()\n\ncriterion = FocalLoss(gamma=2.0, label_smoothing=CFG.LABEL_SMOOTHING)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=CFG.LR,\n    weight_decay=CFG.WEIGHT_DECAY,\n)\nscheduler = CosineAnnealingLR(\n    optimizer,\n    T_max=CFG.EPOCHS,\n    eta_min=1e-6,\n)\nscaler = torch.amp.GradScaler('cuda')\n\nprint('Loss / optimizer / scheduler OK')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T12:25:14.574422Z","iopub.execute_input":"2026-04-11T12:25:14.574985Z","iopub.status.idle":"2026-04-11T12:25:14.581907Z","shell.execute_reply.started":"2026-04-11T12:25:14.574957Z","shell.execute_reply":"2026-04-11T12:25:14.581247Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Cell 11: Train & Validate Functions ──────────────────────────────────────\ndef train_one_epoch(model, loader, optimizer, criterion, scaler, device):\n    model.train()\n    total_loss, all_preds, all_labels = 0.0, [], []\n    for images, labels in tqdm(loader, desc='  Train', leave=False):\n        images, labels = images.to(device), labels.to(device)\n        mixed, lab_a, lab_b, lam = augment_batch(images, labels)\n        optimizer.zero_grad()\n        with torch.amp.autocast('cuda'):\n            logits = model(mixed)\n            loss   = lam * criterion(logits, lab_a) + \\\n                     (1 - lam) * criterion(logits, lab_b)\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        scaler.step(optimizer)\n        scaler.update()\n        total_loss += loss.item()\n        all_preds.extend(logits.argmax(1).cpu().numpy())\n        all_labels.extend(lab_a.cpu().numpy())\n    return total_loss / len(loader), f1_score(all_labels, all_preds, average='macro')\n\n\ndef validate(model, loader, criterion, device):\n    model.eval()\n    total_loss, all_preds, all_labels, all_probs = 0.0, [], [], []\n    with torch.no_grad():\n        for images, labels in tqdm(loader, desc='  Val  ', leave=False):\n            images, labels = images.to(device), labels.to(device)\n            with torch.amp.autocast('cuda'):\n                logits = model(images)\n                loss   = criterion(logits, labels)\n            total_loss += loss.item()\n            probs = F.softmax(logits.float(), dim=1).cpu().numpy()\n            preds = probs.argmax(axis=1).copy()\n            preds[probs[:, 3] > CFG.MULTI_DISEASE_THRESH] = 3\n            all_probs.extend(probs)\n            all_preds.extend(preds)\n            all_labels.extend(labels.cpu().numpy())\n    all_probs = np.stack(all_probs)\n    f1  = f1_score(all_labels, all_preds, average='macro')\n    try:\n        auc = roc_auc_score(all_labels, all_probs,\n                            multi_class='ovr', average='macro')\n    except ValueError as e:\n        print(f'  [AUC WARNING] {e}')\n        auc = 0.0\n    return total_loss / len(loader), f1, auc, all_labels, all_preds\n\nprint('Training functions OK')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T12:25:21.200811Z","iopub.execute_input":"2026-04-11T12:25:21.201076Z","iopub.status.idle":"2026-04-11T12:25:21.211441Z","shell.execute_reply.started":"2026-04-11T12:25:21.201055Z","shell.execute_reply":"2026-04-11T12:25:21.210849Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Cell 12: Training Loop ───────────────────────────────────────────────────\nprint('=' * 60)\nprint('  Starting Training  —  LeafCNN  (2021 dataset)')\nprint('=' * 60)\n\nbest_val_f1 = 0.0\nhistory     = []\n\nfor epoch in range(1, CFG.EPOCHS + 1):\n    lr = optimizer.param_groups[0]['lr']\n    print(f'\\nEpoch {epoch:02d}/{CFG.EPOCHS}  |  LR: {lr:.2e}')\n    print('-' * 40)\n\n    train_loss, train_f1 = train_one_epoch(\n        model, train_loader, optimizer, criterion, scaler, DEVICE)\n    val_loss, val_f1, val_auc, val_labels, val_preds = validate(\n        model, val_loader, criterion, DEVICE)\n    scheduler.step()\n\n    history.append(dict(epoch=epoch, train_loss=train_loss, train_f1=train_f1,\n                        val_loss=val_loss, val_f1=val_f1, val_auc=val_auc))\n\n    print(f'  Train  loss: {train_loss:.4f}  F1: {train_f1:.4f}')\n    print(f'  Val    loss: {val_loss:.4f}  F1: {val_f1:.4f}  AUC: {val_auc:.4f}')\n\n    if val_f1 > best_val_f1:\n        best_val_f1 = val_f1\n        torch.save({\n            'epoch': epoch,\n            'model_state': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'scheduler_state_dict': scheduler.state_dict(),\n            'val_f1': val_f1,\n            'val_auc': val_auc,\n        }, CFG.SAVE_PATH)\n        print(f'  [SAVED] best model  (F1: {best_val_f1:.4f})')\n\n    # Safety checkpoint every 5 epochs\n    if epoch % 5 == 0:\n        torch.save({\n            'epoch': epoch,\n            'model_state': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'scheduler_state_dict': scheduler.state_dict(),\n            'val_f1': val_f1,\n            'val_auc': val_auc,\n        }, f'checkpoint_epoch{epoch}.pth')\n\n        print('\\n  Per-class report:')\n        report = classification_report(val_labels, val_preds,\n                                       target_names=CFG.CLASSES, digits=3)\n        for line in report.split('\\n'):\n            print('  ' + line)\n\nprint('\\n' + '=' * 60)\nprint(f'  Done.  Best Val F1: {best_val_f1:.4f}')\nprint('=' * 60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T12:25:24.761902Z","iopub.execute_input":"2026-04-11T12:25:24.762436Z","iopub.status.idle":"2026-04-11T21:27:12.188657Z","shell.execute_reply.started":"2026-04-11T12:25:24.762407Z","shell.execute_reply":"2026-04-11T21:27:12.187944Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Bonus","metadata":{}},{"cell_type":"code","source":"# ── TTA Validation ───────────────────────────────────────────────────────────\ndef validate_tta(model, loader, device, n_augments=5):\n    model.eval()\n    all_preds, all_labels, all_probs = [], [], []\n\n    # TTA transforms — lighter than train, no destructive augmentations\n    tta_transform = A.Compose([\n        A.LongestMaxSize(max_size=CFG.IMG_SIZE),\n        A.PadIfNeeded(\n            min_height=CFG.IMG_SIZE, min_width=CFG.IMG_SIZE,\n            border_mode=0, fill=0, position='random', p=1.0,\n        ),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.RandomBrightnessContrast(p=0.3),\n        A.Normalize(mean=MEAN, std=STD),\n        ToTensorV2(),\n    ])\n\n    with torch.no_grad():\n        for images, labels in tqdm(loader, desc='  TTA  ', leave=False):\n            images, labels = images.to(device), labels.to(device)\n            batch_probs = []\n\n            # Pass 1: original image (no augmentation)\n            with torch.amp.autocast('cuda'):\n                logits = model(images)\n            batch_probs.append(F.softmax(logits.float(), dim=1).cpu().numpy())\n\n            # Passes 2-N: augmented versions\n            for _ in range(n_augments - 1):\n                # Re-augment each image in the batch\n                aug_images = []\n                for img in images:\n                    # Convert tensor back to numpy for albumentations\n                    img_np = img.cpu().numpy().transpose(1, 2, 0)\n                    # Denormalize\n                    img_np = (img_np * np.array(STD) + np.array(MEAN))\n                    img_np = np.clip(img_np * 255, 0, 255).astype(np.uint8)\n                    aug = tta_transform(image=img_np)['image']\n                    aug_images.append(aug)\n                aug_tensor = torch.stack(aug_images).to(device)\n                with torch.amp.autocast('cuda'):\n                    logits = model(aug_tensor)\n                batch_probs.append(F.softmax(logits.float(), dim=1).cpu().numpy())\n\n            # Average probabilities across all passes\n            avg_probs = np.mean(batch_probs, axis=0)\n            preds = avg_probs.argmax(axis=1).copy()\n            preds[avg_probs[:, 3] > CFG.MULTI_DISEASE_THRESH] = 3\n\n            all_probs.extend(avg_probs)\n            all_preds.extend(preds)\n            all_labels.extend(labels.cpu().numpy())\n\n    all_probs = np.stack(all_probs)\n    f1  = f1_score(all_labels, all_preds, average='macro')\n    try:\n        auc = roc_auc_score(all_labels, all_probs,\n                            multi_class='ovr', average='macro')\n    except ValueError as e:\n        print(f'  [AUC WARNING] {e}')\n        auc = 0.0\n\n    print(f'\\n  TTA Val  F1: {f1:.4f}  AUC: {auc:.4f}')\n    print('\\n  Per-class report:')\n    report = classification_report(all_labels, all_preds,\n                                   target_names=CFG.CLASSES, digits=3)\n    for line in report.split('\\n'):\n        print('  ' + line)\n    return f1, auc\n\n\n# Load best model then run TTA\ncheckpoint = torch.load(CFG.SAVE_PATH, map_location=DEVICE, weights_only=False)\nmodel.load_state_dict(checkpoint['model_state'])\nprint(f\"Loaded best model from epoch {checkpoint['epoch']} (F1: {checkpoint['val_f1']:.4f})\")\n\ntta_f1, tta_auc = validate_tta(model, val_loader, DEVICE, n_augments=6)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T22:02:10.318181Z","iopub.execute_input":"2026-04-11T22:02:10.318957Z","iopub.status.idle":"2026-04-11T22:07:07.526585Z","shell.execute_reply.started":"2026-04-11T22:02:10.318915Z","shell.execute_reply":"2026-04-11T22:07:07.525608Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Cell 13: Plot Training Curves ────────────────────────────────────────────\nhist_df = pd.DataFrame(history)\n\nfig, axes = plt.subplots(1, 3, figsize=(15, 4))\n\naxes[0].plot(hist_df['epoch'], hist_df['train_loss'], label='Train')\naxes[0].plot(hist_df['epoch'], hist_df['val_loss'],   label='Val')\naxes[0].set_title('Loss'); axes[0].legend()\n\naxes[1].plot(hist_df['epoch'], hist_df['train_f1'], label='Train')\naxes[1].plot(hist_df['epoch'], hist_df['val_f1'],   label='Val')\naxes[1].set_title('F1 Score'); axes[1].legend()\n\naxes[2].plot(hist_df['epoch'], hist_df['val_auc'], label='Val AUC', color='green')\naxes[2].set_title('Val AUC'); axes[2].legend()\n\nplt.tight_layout()\nplt.savefig('training_curves.png', dpi=150)\nplt.show()\nprint('Saved training_curves.png')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\ndef plot_class_distribution(df, classes, save_path='class_distribution.png'):\n    counts = [len(df[df['label_name'] == cls]) for cls in classes]\n    labels = [cls.replace('_', '\\n').title() for cls in classes]\n\n    colors = ['#4CAF50', '#2196F3', '#FF9800', '#F44336']\n\n    fig, ax = plt.subplots(figsize=(9, 6))\n    fig.patch.set_facecolor('white')\n    ax.set_facecolor('white')\n\n    bars = ax.bar(labels, counts, color=colors, width=0.55,\n                  edgecolor='white', linewidth=1.5)\n\n    # Count labels on top of each bar\n    for bar, count in zip(bars, counts):\n        ax.text(bar.get_x() + bar.get_width() / 2,\n                bar.get_height() + 40,\n                f'{count:,}', ha='center', va='bottom',\n                fontsize=13, fontweight='bold', color='#222222')\n\n    ax.set_ylabel('Number of Images', fontsize=13, labelpad=10)\n    ax.set_ylim(0, max(counts) * 1.18)\n    ax.set_xticks(range(len(labels)))\n    ax.set_xticklabels(labels, fontsize=12)\n    ax.yaxis.grid(True, linestyle='--', alpha=0.5, color='#cccccc')\n    ax.set_axisbelow(True)\n    ax.spines['top'].set_visible(False)\n    ax.spines['right'].set_visible(False)\n    ax.spines['left'].set_color('#cccccc')\n    ax.spines['bottom'].set_color('#cccccc')\n\n    fig.text(\n        0.5, -0.04,\n        \"Figure 3.3: Class distribution in the Plant Pathology 2021 training set.\\n\"\n        \"The multiple_diseases class is significantly underrepresented relative to the\\n\"\n        \"other three classes, motivating the use of Focal Loss and class-specific\\n\"\n        \"weighting in the training strategy.\",\n        ha='center', va='top', fontsize=10, color='#444444',\n        style='italic', wrap=True\n    )\n\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=200, bbox_inches='tight', facecolor='white')\n    plt.show()\n    print(f'Saved to {save_path}')\n\nplot_class_distribution(df, CFG.CLASSES, save_path='class_distribution.png')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T22:40:29.814288Z","iopub.execute_input":"2026-04-19T22:40:29.814626Z","iopub.status.idle":"2026-04-19T22:40:30.576152Z","shell.execute_reply.started":"2026-04-19T22:40:29.814600Z","shell.execute_reply":"2026-04-19T22:40:30.575553Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport matplotlib.patches as FancyBboxPatch\nfrom matplotlib.patches import FancyArrowPatch\nimport numpy as np\n\ndef plot_dl_pipeline(save_path='figure_3_2_dl_pipeline.png'):\n    fig, ax = plt.subplots(figsize=(14, 4))\n    fig.patch.set_facecolor('white')\n    ax.set_facecolor('white')\n    ax.axis('off')\n\n    steps = [\n        ('Raw Image\\n(Pixels)', '#4CAF50'),\n        ('Conv Block 1\\n32 filters', '#2196F3'),\n        ('Conv Block 2\\n64 filters', '#2196F3'),\n        ('Conv Block 3\\n128 filters', '#2196F3'),\n        ('Conv Block 4\\n256 filters', '#2196F3'),\n        ('Conv Block 5\\n512 filters', '#2196F3'),\n        ('GAP +\\nClassifier', '#FF9800'),\n        ('Predicted\\nClass', '#F44336'),\n    ]\n\n    box_w, box_h = 0.10, 0.45\n    gap = 0.035\n    start_x = 0.01\n    y = 0.5\n\n    positions = []\n    for i, (label, color) in enumerate(steps):\n        x = start_x + i * (box_w + gap)\n        fancy = plt.Rectangle((x, y - box_h / 2), box_w, box_h,\n                               linewidth=1.5, edgecolor='white',\n                               facecolor=color, zorder=3,\n                               transform=ax.transAxes, clip_on=False)\n        ax.add_patch(fancy)\n        ax.text(x + box_w / 2, y, label, transform=ax.transAxes,\n                ha='center', va='center', fontsize=8.5,\n                fontweight='bold', color='white', zorder=4,\n                multialignment='center')\n        positions.append((x, x + box_w))\n\n    # Arrows between boxes\n    for i in range(len(positions) - 1):\n        x_start = positions[i][1]\n        x_end   = positions[i + 1][0]\n        ax.annotate('', xy=(x_end, y), xytext=(x_start, y),\n                    xycoords='axes fraction', textcoords='axes fraction',\n                    arrowprops=dict(arrowstyle='->', color='#555555',\n                                   lw=1.8), zorder=5)\n\n    # Labels below\n    category_labels = [\n        (positions[0][0] + box_w / 2, 'Input'),\n        (positions[1][0] + (4 * (box_w + gap) + box_w) / 2, 'Automatic Feature Extraction (CNN)'),\n        (positions[6][0] + box_w / 2, 'Head'),\n        (positions[7][0] + box_w / 2, 'Output'),\n    ]\n    for xpos, txt in category_labels:\n        ax.text(xpos, y - box_h / 2 - 0.08, txt,\n                transform=ax.transAxes, ha='center', va='top',\n                fontsize=9, color='#444444', style='italic')\n\n    fig.text(0.5, 0.02,\n             \"Figure 3.2: Deep learning pipeline for plant disease classification. \"\n             \"Unlike the classical approach, no manual feature\\n\"\n             \"engineering is required. The CNN automatically learns a hierarchical set of features \"\n             \"directly from raw pixel values during training.\",\n             ha='center', fontsize=9, color='#444444', style='italic')\n\n    plt.savefig(save_path, dpi=200, bbox_inches='tight', facecolor='white')\n    plt.show()\n    print(f'Saved to {save_path}')\n\nplot_dl_pipeline()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-20T15:25:46.275062Z","iopub.execute_input":"2026-04-20T15:25:46.275282Z","iopub.status.idle":"2026-04-20T15:25:46.675596Z","shell.execute_reply.started":"2026-04-20T15:25:46.275257Z","shell.execute_reply":"2026-04-20T15:25:46.674809Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_preprocessing_pipeline(df, img_dir, save_path='figure_4_3_preprocessing.png'):\n    # Pick one image\n    row      = df[df['label_name'] == 'scab'].iloc[2]\n    img_path = os.path.join(img_dir, row['image'])\n    orig     = cv2.cvtColor(cv2.imread(img_path), cv2.COLOR_BGR2RGB)\n\n    SIZE = CFG.IMG_SIZE\n\n    # Step 1: LongestMaxSize\n    h, w  = orig.shape[:2]\n    scale = SIZE / max(h, w)\n    resized = cv2.resize(orig, (int(w * scale), int(h * scale)))\n\n    # Step 2: PadIfNeeded (center)\n    rh, rw = resized.shape[:2]\n    pad_top    = (SIZE - rh) // 2\n    pad_bottom = SIZE - rh - pad_top\n    pad_left   = (SIZE - rw) // 2\n    pad_right  = SIZE - rw - pad_left\n    padded = cv2.copyMakeBorder(resized, pad_top, pad_bottom,\n                                pad_left, pad_right,\n                                cv2.BORDER_CONSTANT, value=0)\n\n    # Step 3: CenterCrop (already 512x512 here, just show final)\n    cropped = padded[:SIZE, :SIZE]\n\n    stages = [\n        (orig,    f'Original\\n({w}×{h} px)'),\n        (resized, f'LongestMaxSize\\n({resized.shape[1]}×{resized.shape[0]} px)'),\n        (padded,  f'PadIfNeeded\\n({SIZE}×{SIZE} px)'),\n        (cropped, f'Final\\n({SIZE}×{SIZE} px)'),\n    ]\n\n    fig, axes = plt.subplots(1, 4, figsize=(16, 5))\n    fig.patch.set_facecolor('white')\n\n    for ax, (img, title) in zip(axes, stages):\n        ax.imshow(img)\n        ax.set_title(title, fontsize=11, fontweight='bold', pad=8)\n        ax.axis('off')\n\n    # Arrows between subplots\n    for i in range(3):\n        fig.text(0.245 + i * 0.188, 0.52, '→',\n                 fontsize=22, ha='center', color='#555555')\n\n    fig.text(0.5, 0.01,\n             \"Figure 4.3: Image preprocessing pipeline. The original high-resolution photograph is first resized so its longest\\n\"\n             \"dimension equals 512 pixels while preserving the aspect ratio, then padded with black pixels to reach 512×512,\\n\"\n             \"and finally cropped to the target dimensions. Black padding regions are visible as dark borders around the leaf content.\",\n             ha='center', fontsize=9, color='#444444', style='italic')\n\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=200, bbox_inches='tight', facecolor='white')\n    plt.show()\n    print(f'Saved to {save_path}')\n\nplot_preprocessing_pipeline(df, CFG.IMG_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-20T15:26:55.090379Z","iopub.execute_input":"2026-04-20T15:26:55.090710Z","iopub.status.idle":"2026-04-20T15:26:58.218832Z","shell.execute_reply.started":"2026-04-20T15:26:55.090678Z","shell.execute_reply":"2026-04-20T15:26:58.218107Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_augmentation_grid(df, img_dir, save_path='figure_4_4_augmentations.png'):\n    row      = df[df['label_name'] == 'rust'].iloc[0]\n    img_path = os.path.join(img_dir, row['image'])\n    orig     = cv2.cvtColor(cv2.imread(img_path), cv2.COLOR_BGR2RGB)\n\n    # Resize to 512 first\n    base = A.Compose([\n        A.LongestMaxSize(max_size=CFG.IMG_SIZE),\n        A.PadIfNeeded(min_height=CFG.IMG_SIZE, min_width=CFG.IMG_SIZE,\n                      border_mode=0, fill=0, position='center', p=1.0),\n    ])(image=orig)['image']\n\n    augmentations = [\n        ('Original',           A.Compose([])),\n        ('HorizontalFlip',     A.Compose([A.HorizontalFlip(p=1.0)])),\n        ('VerticalFlip',       A.Compose([A.VerticalFlip(p=1.0)])),\n        ('RandomRotate90',     A.Compose([A.RandomRotate90(p=1.0)])),\n        ('GridDistortion',     A.Compose([A.GridDistortion(num_steps=5, distort_limit=0.3, p=1.0)])),\n        ('ElasticTransform',   A.Compose([A.ElasticTransform(alpha=120, sigma=6, p=1.0)])),\n        ('ColorJitter',        A.Compose([A.ColorJitter(brightness=0.4, contrast=0.4, saturation=0.4, hue=0.1, p=1.0)])),\n        ('GaussNoise',         A.Compose([A.GaussNoise(p=1.0)])),\n        ('MotionBlur',         A.Compose([A.MotionBlur(blur_limit=9, p=1.0)])),\n        ('CoarseDropout',      A.Compose([A.CoarseDropout(num_holes_range=(4, 8),\n                                          hole_height_range=(64, 128),\n                                          hole_width_range=(64, 128),\n                                          fill=0, p=1.0)])),\n        ('Sharpen',            A.Compose([A.Sharpen(alpha=(0.5, 0.8), p=1.0)])),\n        ('Affine',             A.Compose([A.Affine(translate_percent=0.1,\n                                          scale=(0.85, 1.15), rotate=(-45, 45), p=1.0)])),\n    ]\n\n    cols = 4\n    rows = int(np.ceil(len(augmentations) / cols))\n    fig, axes = plt.subplots(rows, cols, figsize=(cols * 4, rows * 4))\n    fig.patch.set_facecolor('white')\n    axes = axes.flatten()\n\n    for i, (name, aug) in enumerate(augmentations):\n        result = aug(image=base.copy())['image']\n        axes[i].imshow(result)\n        axes[i].set_title(name, fontsize=10, fontweight='bold', pad=6)\n        axes[i].axis('off')\n\n    # Hide unused axes\n    for j in range(len(augmentations), len(axes)):\n        axes[j].axis('off')\n\n    fig.text(0.5, 0.01,\n             \"Figure 4.4: Examples of augmentation transforms applied to a single training image. \"\n             \"Each cell shows the result of one augmentation\\napplied in isolation. In practice, multiple transforms are composed \"\n             \"and applied simultaneously during training, producing a\\ndiverse set of perturbed views of each image.\",\n             ha='center', fontsize=9, color='#444444', style='italic')\n\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=200, bbox_inches='tight', facecolor='white')\n    plt.show()\n    print(f'Saved to {save_path}')\n\nplot_augmentation_grid(df, CFG.IMG_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-20T15:28:02.889329Z","iopub.execute_input":"2026-04-20T15:28:02.890108Z","iopub.status.idle":"2026-04-20T15:28:07.369185Z","shell.execute_reply.started":"2026-04-20T15:28:02.890075Z","shell.execute_reply":"2026-04-20T15:28:07.368321Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Figure 4.5: MixUp & CutMix","metadata":{}},{"cell_type":"code","source":"def plot_mixup_cutmix(df, img_dir, save_path='figure_4_5_mixup_cutmix.png'):\n    base_tf = A.Compose([\n        A.LongestMaxSize(max_size=256),\n        A.PadIfNeeded(min_height=256, min_width=256,\n                      border_mode=0, fill=0, position='center', p=1.0),\n    ])\n\n    row1 = df[df['label_name'] == 'scab'].iloc[0]\n    row2 = df[df['label_name'] == 'rust'].iloc[0]\n\n    img1 = cv2.cvtColor(cv2.imread(os.path.join(img_dir, row1['image'])), cv2.COLOR_BGR2RGB)\n    img2 = cv2.cvtColor(cv2.imread(os.path.join(img_dir, row2['image'])), cv2.COLOR_BGR2RGB)\n    img1 = base_tf(image=img1)['image'].astype(np.float32) / 255.0\n    img2 = base_tf(image=img2)['image'].astype(np.float32) / 255.0\n\n    # MixUp\n    lam_mix  = 0.6\n    mixed    = lam_mix * img1 + (1 - lam_mix) * img2\n\n    # CutMix\n    cutmixed = img1.copy()\n    cx, cy   = 128, 128\n    cut_w, cut_h = 100, 100\n    x1, x2   = cx - cut_w // 2, cx + cut_w // 2\n    y1, y2   = cy - cut_h // 2, cy + cut_h // 2\n    cutmixed[y1:y2, x1:x2] = img2[y1:y2, x1:x2]\n\n    panels = [\n        (img1,     f'Image A\\n(scab)'),\n        (img2,     f'Image B\\n(rust)'),\n        (mixed,    f'MixUp\\n(λ={lam_mix})'),\n        (cutmixed, 'CutMix\\n(patch from B)'),\n    ]\n\n    fig, axes = plt.subplots(1, 4, figsize=(16, 5))\n    fig.patch.set_facecolor('white')\n\n    for ax, (img, title) in zip(axes, panels):\n        ax.imshow(np.clip(img, 0, 1))\n        ax.set_title(title, fontsize=12, fontweight='bold', pad=8)\n        ax.axis('off')\n\n    # Arrow and plus signs\n    for i, sym in enumerate(['+', '→\\nMixUp', '→\\nCutMix']):\n        fig.text(0.245 + i * 0.188, 0.52, sym,\n                 fontsize=13, ha='center', va='center', color='#555555')\n\n    fig.text(0.5, 0.01,\n             \"Figure 4.5: Illustration of MixUp and CutMix augmentation. MixUp (third column) produces a pixel-wise blend of two images\\n\"\n             \"with soft combined labels. CutMix (fourth column) replaces a rectangular region of one image with the corresponding region\\n\"\n             \"from another, preserving local realism while mixing label information proportionally to the area of the pasted region.\",\n             ha='center', fontsize=9, color='#444444', style='italic')\n\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=200, bbox_inches='tight', facecolor='white')\n    plt.show()\n    print(f'Saved to {save_path}')\n\nplot_mixup_cutmix(df, CFG.IMG_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-20T15:30:42.645325Z","iopub.execute_input":"2026-04-20T15:30:42.645675Z","iopub.status.idle":"2026-04-20T15:30:44.296560Z","shell.execute_reply.started":"2026-04-20T15:30:42.645646Z","shell.execute_reply":"2026-04-20T15:30:44.295804Z"}},"outputs":[],"execution_count":null}]}