{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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},{"sourceType":"datasetVersion","sourceId":1105512,"datasetId":619070,"databundleVersionId":1135665}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport random\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.cuda.amp import GradScaler, autocast\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import f1_score, roc_auc_score, classification_report","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-08T11:31:23.607321Z","iopub.execute_input":"2026-04-08T11:31:23.608212Z","iopub.status.idle":"2026-04-08T11:31:23.613218Z","shell.execute_reply.started":"2026-04-08T11:31:23.608177Z","shell.execute_reply":"2026-04-08T11:31:23.612379Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"class CFG:\n    CSV_PATH  = \"/kaggle/input/datasets/piantic/plantpathology-apple-dataset/train.csv\"\n    IMG_DIR   = \"/kaggle/input/datasets/piantic/plantpathology-apple-dataset/images\"\n    SAVE_PATH = \"best_model_v2.pth\"      # NEW filename, don't overwrite your old model\n\n    IMG_SIZE    = 512\n    VAL_SPLIT   = 0.2\n    NUM_CLASSES = 4\n\n    EPOCHS      = 45          # middle of your 40-50 range\n    BATCH_SIZE  = 16\n    NUM_WORKERS = 2\n    SEED        = 42\n\n    LR           = 1e-4       # full LR since training from scratch\n    WEIGHT_DECAY = 1e-2\n\n    LABEL_SMOOTHING = 0.1\n\n    MULTI_DISEASE_BOOST  = 3.0\n    MULTI_DISEASE_THRESH = 0.50   # best threshold from your test\n\n    MIXUP_ALPHA = 0.4\n\n    CLASSES = [\"healthy\", \"scab\", \"rust\", \"multiple_diseases\"]\n\n    PATH_2021_CSV = \"/kaggle/input/competitions/plant-pathology-2021-fgvc8/train.csv\"\n    PATH_2021_IMG = \"/kaggle/input/competitions/plant-pathology-2021-fgvc8/train_images\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T11:31:27.917895Z","iopub.execute_input":"2026-04-08T11:31:27.918265Z","iopub.status.idle":"2026-04-08T11:31:27.923114Z","shell.execute_reply.started":"2026-04-08T11:31:27.918238Z","shell.execute_reply":"2026-04-08T11:31:27.922514Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Reproducibility","metadata":{}},{"cell_type":"code","source":"def 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    torch.backends.cudnn.benchmark     = False\n\nseed_everything(CFG.SEED)\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {DEVICE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T11:31:31.264736Z","iopub.execute_input":"2026-04-08T11:31:31.265058Z","iopub.status.idle":"2026-04-08T11:31:31.272869Z","shell.execute_reply.started":"2026-04-08T11:31:31.265028Z","shell.execute_reply":"2026-04-08T11:31:31.272131Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load & Merge Datasets","metadata":{}},{"cell_type":"code","source":"# ── 2020 dataset ──────────────────────────────────────────────────────────────\ndf = pd.read_csv(CFG.CSV_PATH)\nlabel_cols  = [\"healthy\", \"scab\", \"rust\", \"multiple_diseases\"]\ndf[\"label\"] = df[label_cols].values.argmax(axis=1)\ndf[\"_img_dir\"] = CFG.IMG_DIR\n\nprint(\"Class distribution (2020):\")\nfor i, cls in enumerate(CFG.CLASSES):\n    n = (df[\"label\"] == i).sum()\n    print(f\"  {cls:<20} {n}  ({100*n/len(df):.1f}%)\")\n\n# ── 2021 dataset ──────────────────────────────────────────────────────────────\nLABEL_MAP_2021 = {\"healthy\": 0, \"scab\": 1, \"rust\": 2, \"complex\": 3}\n\ndef load_2021(csv_path, img_dir):\n    raw = pd.read_csv(csv_path)\n    rows = []\n    skipped_multi, skipped_unknown = 0, 0\n\n    for _, row in raw.iterrows():\n        parts = row[\"labels\"].strip().split()\n        if len(parts) > 1:\n            skipped_multi += 1\n            continue\n        label_str = parts[0]\n        if label_str not in LABEL_MAP_2021:\n            skipped_unknown += 1\n            continue\n        label_int = LABEL_MAP_2021[label_str]\n        rows.append({\n            \"image_id\"          : os.path.splitext(row[\"image\"])[0],\n            \"healthy\"           : int(label_int == 0),\n            \"scab\"              : int(label_int == 1),\n            \"rust\"              : int(label_int == 2),\n            \"multiple_diseases\" : int(label_int == 3),\n            \"label\"             : label_int,\n            \"_img_dir\"          : img_dir,\n        })\n\n    df_2021 = pd.DataFrame(rows)\n    print(f\"\\n2021 kept: {len(df_2021)} rows \"\n          f\"(skipped multi={skipped_multi}, unknown={skipped_unknown})\")\n    return df_2021\n\ndf_2021   = load_2021(CFG.PATH_2021_CSV, CFG.PATH_2021_IMG)\ndf_merged = pd.concat([df, df_2021], ignore_index=True)\n\nprint(f\"\\nMerged total: {len(df_merged)} rows\")\nfor i, cls in enumerate(CFG.CLASSES):\n    n = (df_merged[\"label\"] == i).sum()\n    print(f\"  {cls:<20} {n}  ({100*n/len(df_merged):.1f}%)\")\n\n# ── Train / Val split ─────────────────────────────────────────────────────────\ntrain_df, val_df = train_test_split(\n    df_merged,\n    test_size    = CFG.VAL_SPLIT,\n    stratify     = df_merged[\"label\"],\n    random_state = CFG.SEED,\n)\ntrain_df = train_df.reset_index(drop=True)\nval_df   = val_df.reset_index(drop=True)\nprint(f\"\\nTrain: {len(train_df)}  |  Val: {len(val_df)}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset Class","metadata":{}},{"cell_type":"code","source":"class LeafDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.df        = dataframe\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(row[\"_img_dir\"], row[\"image_id\"] + \".jpg\")\n        image    = np.array(Image.open(img_path).convert(\"RGB\"))\n        label    = int(row[\"label\"])\n        if self.transform:\n            image = self.transform(image=image)[\"image\"]\n        return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T11:31:43.473033Z","iopub.execute_input":"2026-04-08T11:31:43.473810Z","iopub.status.idle":"2026-04-08T11:31:43.479057Z","shell.execute_reply.started":"2026-04-08T11:31:43.473776Z","shell.execute_reply":"2026-04-08T11:31:43.478212Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Augmentation Pipelines","metadata":{}},{"cell_type":"code","source":"MEAN = [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\n    A.PadIfNeeded(\n        min_height=CFG.IMG_SIZE,\n        min_width=CFG.IMG_SIZE,\n        border_mode=0,\n        fill=0,  # ✅ FIX: 'value' → 'fill'\n        position=\"random\",\n        p=1.0,\n    ),\n\n    A.RandomCrop(height=CFG.IMG_SIZE, width=CFG.IMG_SIZE),\n\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n\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\n    # ✅ FIX: ShiftScaleRotate → Affine\n    A.Affine(\n        translate_percent=0.1,\n        scale=(0.9, 1.1),\n        rotate=(-30, 30),\n        p=0.5\n    ),\n\n    A.ColorJitter(\n        brightness=0.2, contrast=0.2,\n        saturation=0.2, hue=0.1, p=0.5\n    ),\n    A.RandomBrightnessContrast(p=0.3),\n\n    A.GaussNoise(p=0.2),\n    A.MotionBlur(blur_limit=5, p=0.2),\n\n    # ✅ FIX: CoarseDropout new API\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,\n        p=0.5,\n    ),\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\n    A.PadIfNeeded(\n        min_height=CFG.IMG_SIZE,\n        min_width=CFG.IMG_SIZE,\n        border_mode=0,\n        fill=0,  # ✅ FIX\n        position=\"center\",\n        p=1.0,\n    ),\n\n    A.CenterCrop(height=CFG.IMG_SIZE, width=CFG.IMG_SIZE),\n\n    A.Normalize(mean=MEAN, std=STD),\n    ToTensorV2(),\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T11:31:50.752948Z","iopub.execute_input":"2026-04-08T11:31:50.753282Z","iopub.status.idle":"2026-04-08T11:31:50.774751Z","shell.execute_reply.started":"2026-04-08T11:31:50.753251Z","shell.execute_reply":"2026-04-08T11:31:50.773898Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# DataLoaders","metadata":{}},{"cell_type":"code","source":"train_dataset = LeafDataset(train_df, transform=train_transform)\nval_dataset   = LeafDataset(val_df,   transform=val_transform)\n\n# Weighted sampler — still needed even with merged data\nclass_counts      = train_df[\"label\"].value_counts().sort_index().values\nclass_weights     = 1.0 / class_counts.astype(float)\nclass_weights[3] *= CFG.MULTI_DISEASE_BOOST      # boost multiple_diseases\nsample_weights    = class_weights[train_df[\"label\"].values]\n\nprint(\"Sampling weights per class:\")\nfor i, cls in enumerate(CFG.CLASSES):\n    print(f\"  {cls:<20} {class_weights[i]:.6f}\")\n\nsampler = WeightedRandomSampler(\n    weights     = sample_weights,\n    num_samples = len(train_dataset),\n    replacement = True,\n)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size  = CFG.BATCH_SIZE,\n    sampler     = sampler,\n    num_workers = CFG.NUM_WORKERS,\n    pin_memory  = True,\n)\nval_loader = DataLoader(\n    val_dataset,\n    batch_size  = CFG.BATCH_SIZE,\n    shuffle     = False,\n    num_workers = CFG.NUM_WORKERS,\n    pin_memory  = True,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T11:32:02.286402Z","iopub.execute_input":"2026-04-08T11:32:02.286727Z","iopub.status.idle":"2026-04-08T11:32:02.295209Z","shell.execute_reply.started":"2026-04-08T11:32:02.286692Z","shell.execute_reply":"2026-04-08T11:32:02.294458Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# The CNN Architecture ","metadata":{}},{"cell_type":"code","source":"class LeafCNN(nn.Module):\n    def __init__(self, num_classes=4, dropout=0.5):\n        super(LeafCNN, self).__init__()\n\n        # Block 1: (B, 3, 512, 512) → (B, 32, 256, 256)\n        self.block1 = nn.Sequential(\n            nn.Conv2d(3,  32, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(32),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(32, 32, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(32),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(2, 2),\n        )\n\n        # Block 2: (B, 32, 256, 256) → (B, 64, 128, 128)\n        self.block2 = nn.Sequential(\n            nn.Conv2d(32, 64, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(64, 64, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(2, 2),\n        )\n\n        # Block 3: (B, 64, 128, 128) → (B, 128, 64, 64)\n        self.block3 = nn.Sequential(\n            nn.Conv2d(64,  128, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(128, 128, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(2, 2),\n        )\n\n        # Block 4: (B, 128, 64, 64) → (B, 256, 32, 32)\n        self.block4 = nn.Sequential(\n            nn.Conv2d(128, 256, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(256, 256, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(2, 2),\n        )\n\n        # Block 5: (B, 256, 32, 32) → (B, 512, 16, 16)\n        self.block5 = nn.Sequential(\n            nn.Conv2d(256, 512, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(512, 512, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(2, 2),\n        )\n\n        # GAP: (B, 512, 16, 16) → (B, 512, 1, 1)\n        self.gap = nn.AdaptiveAvgPool2d(1)\n\n        # Deeper classifier head (NEW)\n        # 512 → 256 → 128 → 4  with BN between each FC layer\n        self.classifier = nn.Sequential(\n            nn.Dropout(p=dropout),\n            nn.Linear(512, 256),\n            nn.ReLU(inplace=True),\n            nn.BatchNorm1d(256),\n            nn.Dropout(p=dropout * 0.6),\n            nn.Linear(256, 128),\n            nn.ReLU(inplace=True),\n            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)          # (B, 32,  256, 256)\n        x = self.block2(x)          # (B, 64,  128, 128)\n        x = self.block3(x)          # (B, 128,  64,  64)\n        x = self.block4(x)          # (B, 256,  32,  32)\n        x = self.block5(x)          # (B, 512,  16,  16)\n        x = self.gap(x)             # (B, 512,   1,   1)\n        x = x.view(x.size(0), -1)  # (B, 512)\n        x = self.classifier(x)     # (B, 4)\n        return x\n\n\n# Instantiate fresh model with upgraded architecture\nmodel = LeafCNN(num_classes=CFG.NUM_CLASSES, dropout=0.5).to(DEVICE)\n\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\n# Sanity check\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}\")    # should be torch.Size([2, 4])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T11:32:11.203308Z","iopub.execute_input":"2026-04-08T11:32:11.203962Z","iopub.status.idle":"2026-04-08T11:32:11.264919Z","shell.execute_reply.started":"2026-04-08T11:32:11.203931Z","shell.execute_reply":"2026-04-08T11:32:11.264180Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loss, Optimizer, Scheduler","metadata":{}},{"cell_type":"code","source":"class 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)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max   = CFG.EPOCHS,\n    eta_min = 1e-6,\n)\n\nscaler = torch.amp.GradScaler('cuda')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T11:32:17.498696Z","iopub.execute_input":"2026-04-08T11:32:17.499019Z","iopub.status.idle":"2026-04-08T11:32:17.505899Z","shell.execute_reply.started":"2026-04-08T11:32:17.498989Z","shell.execute_reply":"2026-04-08T11:32:17.505226Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# MixUp & CutMix","metadata":{}},{"cell_type":"code","source":"def 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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T11:32:23.002948Z","iopub.execute_input":"2026-04-08T11:32:23.003261Z","iopub.status.idle":"2026-04-08T11:32:23.011372Z","shell.execute_reply.started":"2026-04-08T11:32:23.003233Z","shell.execute_reply":"2026-04-08T11:32:23.010738Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train & Validation Functions","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, criterion, scaler, device):\n    model.train()\n    total_loss, all_preds, all_labels = 0.0, [], []\n\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\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\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\n        total_loss += loss.item()\n        all_preds.extend(logits.argmax(1).cpu().numpy())\n        all_labels.extend(lab_a.cpu().numpy())\n\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\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\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\n            all_probs.extend(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    return total_loss / len(loader), f1, auc, all_labels, all_preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T11:32:27.728986Z","iopub.execute_input":"2026-04-08T11:32:27.729292Z","iopub.status.idle":"2026-04-08T11:32:27.739468Z","shell.execute_reply.started":"2026-04-08T11:32:27.729260Z","shell.execute_reply":"2026-04-08T11:32:27.738748Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Main Training Loop","metadata":{}},{"cell_type":"code","source":"print(\"=\" * 60)\nprint(\"  Starting Training  —  Custom LeafCNN\")\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\n    scheduler.step()\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({\"epoch\": epoch, \"model_state\": model.state_dict(),\n                    \"val_f1\": val_f1, \"val_auc\": val_auc}, CFG.SAVE_PATH)\n        print(f\"  [SAVED] best model  (F1: {best_val_f1:.4f})\")\n\n    if epoch % 5 == 0:\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},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Plot Training Curves","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nhist_df = pd.DataFrame(history)\nfig, axes = plt.subplots(1, 3, figsize=(15, 4))\nfig.suptitle(\"LeafCNN — Training History\", fontsize=14, fontweight=\"bold\")\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].set_xlabel(\"Epoch\"); 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(\"Macro F1\"); axes[1].set_xlabel(\"Epoch\"); axes[1].legend()\n\naxes[2].plot(hist_df[\"epoch\"], hist_df[\"val_auc\"], color=\"purple\")\naxes[2].set_title(\"Val ROC-AUC\"); axes[2].set_xlabel(\"Epoch\")\n\nplt.tight_layout()\nplt.savefig(\"training_curves.png\", dpi=150)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T10:24:58.979443Z","iopub.execute_input":"2026-04-08T10:24:58.980078Z","iopub.status.idle":"2026-04-08T10:24:59.725615Z","shell.execute_reply.started":"2026-04-08T10:24:58.980045Z","shell.execute_reply":"2026-04-08T10:24:59.725039Z"}},"outputs":[],"execution_count":null}]}