{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":25563,"databundleVersionId":2094376,"sourceType":"competition"},{"sourceId":121906917,"sourceType":"kernelVersion"},{"sourceId":670827,"sourceType":"modelInstanceVersion","modelInstanceId":508097,"modelId":522766}],"dockerImageVersionId":30408,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Installation and Imports\nimport os\nimport math\nimport random\nimport json\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom pathlib import Path\nfrom tqdm import tqdm\nfrom typing import Dict, List, Tuple\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (classification_report, f1_score,\n                             multilabel_confusion_matrix, confusion_matrix)\n\nfrom PIL import Image\n\n# Set random seeds for reproducibility\nRANDOM_SEED = 42\ntorch.manual_seed(RANDOM_SEED)\ntorch.cuda.manual_seed_all(RANDOM_SEED)\nnp.random.seed(RANDOM_SEED)\nrandom.seed(RANDOM_SEED)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False\n\nprint(\"Libraries imported successfully!\")\nprint(f\"PyTorch version: {torch.__version__}\")\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"CUDA device: {torch.cuda.get_device_name(0)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T09:19:52.788926Z","iopub.execute_input":"2026-01-07T09:19:52.789811Z","iopub.status.idle":"2026-01-07T09:19:54.562645Z","shell.execute_reply.started":"2026-01-07T09:19:52.789764Z","shell.execute_reply":"2026-01-07T09:19:54.561611Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class config:\n    # specify the paths to datasets\n    DATA_DIR = Path('../input/plant-pathology-2021-fgvc8/train_images')\n    ROOT_DIR = Path('./data')\n    TRAIN_DIR = ROOT_DIR.joinpath('train')\n    TEST_DIR = ROOT_DIR.joinpath('test')\n    VAL_DIR = ROOT_DIR.joinpath('val')\n\n    # set the input height and width\n    INPUT_HEIGHT = 224\n    INPUT_WIDTH = 224\n\n    # set the input heig/ht and width\n    IMAGENET_MEAN = [0.485, 0.456, 0.406]\n    IMAGENET_STD = [0.229, 0.224, 0.225]\n    \n    IMAGE_TYPE = '.jpg'\n    BATCH_SIZE = 32\n    # will use the vision transformer\n    MODEL_NAME = 'vit_base'\n    \n    DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n    TRAINING_PARAMS = 'training_hyperparams/default_train_params'\n    LABELS = ['complex', 'frog_eye_leaf_spot', 'healthy', 'powdery_mildew', 'rust', 'scab']\n    NUM_CLASSES = len(LABELS)\n    CHECKPOINT_DIR = 'checkpoints'\n","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:19:54.564957Z","iopub.execute_input":"2026-01-07T09:19:54.565615Z","iopub.status.idle":"2026-01-07T09:19:54.573132Z","shell.execute_reply.started":"2026-01-07T09:19:54.565576Z","shell.execute_reply":"2026-01-07T09:19:54.572240Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Configuration\nclass Config:\n    # Data paths - Using relative paths\n    DATA_DIR = Path('../input/plant-pathology-2021-fgvc8')\n    TRAIN_CSV = '/kaggle/input/plant-pathology-2021-fgvc8/train.csv'\n    IMAGES_DIR = '/kaggle/input/multi-label-classification-plant-pathology/data/'\n    \n    # Model parameters\n    INPUT_HEIGHT = 224\n    INPUT_WIDTH = 224\n    BATCH_SIZE = 32\n    NUM_WORKERS = 2\n    \n    # Normalization (ImageNet statistics)\n    IMAGENET_MEAN = [0.485, 0.456, 0.406]\n    IMAGENET_STD = [0.229, 0.224, 0.225]\n    \n    # Classes\n    LABELS = ['complex', 'frog_eye_leaf_spot', 'healthy', 'powdery_mildew', 'rust', 'scab']\n    NUM_CLASSES = len(LABELS)\n    \n    # Training parameters\n    NUM_EPOCHS = 30\n    LEARNING_RATE = 1e-3\n    WEIGHT_DECAY = 1e-4\n    PATIENCE = 3\n    \n    # Device\n    DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n    \n    # Paths for saving outputs\n    OUTPUT_DIR = Path('baseline_outputs')\n    CHECKPOINT_DIR = OUTPUT_DIR / 'checkpoints'\n    PLOT_DIR = OUTPUT_DIR / 'plots'\n    METRICS_DIR = OUTPUT_DIR / 'metrics'\n    \n    @classmethod\n    def create_dirs(cls):\n        \"\"\"Create output directories if they don't exist\"\"\"\n        cls.OUTPUT_DIR.mkdir(exist_ok=True)\n        cls.CHECKPOINT_DIR.mkdir(exist_ok=True)\n        cls.PLOT_DIR.mkdir(exist_ok=True)\n        cls.METRICS_DIR.mkdir(exist_ok=True)\n\n# Create directories\nConfig.create_dirs()\n\nconfig = Config()\nprint(f\"Configuration loaded successfully!\")\nprint(f\"Device: {config.DEVICE}\")\nprint(f\"Output directory: {config.OUTPUT_DIR.absolute()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T09:19:54.574123Z","iopub.execute_input":"2026-01-07T09:19:54.574379Z","iopub.status.idle":"2026-01-07T09:19:54.584745Z","shell.execute_reply.started":"2026-01-07T09:19:54.574355Z","shell.execute_reply":"2026-01-07T09:19:54.583757Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Data Loading and Splitting\ndef load_and_split_data(csv_path: Path, train_size: float = 0.7, \n                        val_size: float = 0.15, test_size: float = 0.15,\n                        random_state: int = 42):\n    \"\"\"\n    Load CSV and split into train/validation/test sets.\n    \n    Args:\n        csv_path: Path to CSV file\n        train_size: Proportion of training data\n        val_size: Proportion of validation data\n        test_size: Proportion of test data\n        random_state: Random seed for reproducibility\n    \n    Returns:\n        train_df, val_df, test_df: DataFrames with image paths and labels\n    \"\"\"\n    # Load CSV\n    df = pd.read_csv(csv_path)\n    \n    # Update image paths to point to local data directory\n    df['image'] = df['image'].apply(lambda x: config.IMAGES_DIR +x)\n    \n    # First split: train vs (val + test)\n    train_df, temp_df = train_test_split(\n        df, train_size=train_size, shuffle=True, random_state=random_state\n    )\n    \n    # Second split: val vs test\n    val_ratio = val_size / (val_size + test_size)\n    val_df, test_df = train_test_split(\n        temp_df, train_size=val_ratio, shuffle=True, random_state=random_state\n    )\n    \n    print(f\"Dataset split complete:\")\n    print(f\"  Train: {len(train_df)} samples\")\n    print(f\"  Validation: {len(val_df)} samples\")\n    print(f\"  Test: {len(test_df)} samples\")\n    \n    return train_df, val_df, test_df\n\n# Load and split data\ntrain_df, val_df, test_df = load_and_split_data(config.TRAIN_CSV)","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:19:54.585914Z","iopub.execute_input":"2026-01-07T09:19:54.586259Z","iopub.status.idle":"2026-01-07T09:19:54.635888Z","shell.execute_reply.started":"2026-01-07T09:19:54.586208Z","shell.execute_reply":"2026-01-07T09:19:54.635107Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Analyze Class Distribution\ndef analyze_class_distribution(train_df: pd.DataFrame):\n    \"\"\"\n    Analyze and visualize class distribution in the training set.\n    \n    Args:\n        train_df: Training DataFrame with 'labels' column\n    \"\"\"\n    fig, axes = plt.subplots(1, 2, figsize=(16, 5))\n    \n    # Plot 1: Label combination distribution\n    train_df['labels'].value_counts().plot(kind='bar', ax=axes[0])\n    axes[0].set_title('Label Combination Distribution')\n    axes[0].set_xlabel('Label Combination')\n    axes[0].set_ylabel('Count')\n    axes[0].tick_params(axis='x', rotation=45)\n    \n    # Plot 2: Individual class distribution\n    all_labels = train_df['labels'].str.split(expand=True).stack().reset_index(drop=True)\n    class_counts = all_labels.value_counts()[config.LABELS]\n    \n    # Calculate class weights (inverse frequency)\n    class_weights = torch.reciprocal(torch.tensor(class_counts.values, dtype=torch.float))\n    class_weights = class_weights / torch.max(class_weights)\n    \n    # Plot bar chart\n    bars = axes[1].bar(class_counts.index, class_counts.values)\n    axes[1].set_title('Individual Class Distribution (with Weights)')\n    axes[1].set_xlabel('Class')\n    axes[1].set_ylabel('Count')\n    axes[1].tick_params(axis='x', rotation=45)\n    \n    # Add weight annotations on bars\n    for bar, weight in zip(bars, class_weights):\n        height = bar.get_height()\n        axes[1].text(bar.get_x() + bar.get_width()/2., height,\n                    f'w={weight:.2f}',\n                    ha='center', va='bottom', fontsize=10)\n    \n    plt.tight_layout()\n    plt.savefig(config.PLOT_DIR / 'class_distribution.png', dpi=300, bbox_inches='tight')\n    plt.show()\n    \n    # Print statistics\n    print(\"\\nClass Statistics:\")\n    print(\"-\" * 50)\n    for label, count, weight in zip(class_counts.index, class_counts.values, class_weights):\n        print(f\"{label:25s}: Count={count:4d} | Weight={weight:.4f}\")\n    \n    return class_counts, class_weights\n\n# Analyze distribution\nclass_counts, class_weights = analyze_class_distribution(train_df)","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:19:54.637842Z","iopub.execute_input":"2026-01-07T09:19:54.638103Z","iopub.status.idle":"2026-01-07T09:19:55.993639Z","shell.execute_reply.started":"2026-01-07T09:19:54.638078Z","shell.execute_reply":"2026-01-07T09:19:55.992634Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Data Augmentation and Preprocessing\n\nprint(\"Setting up data augmentation pipelines...\\n\")\n\n# Training transforms with augmentation (matching old notebook parameters)\ntrain_transforms = transforms.Compose([\n    # Geometric augmentations\n    transforms.RandomResizedCrop(\n        size=(config.INPUT_HEIGHT, config.INPUT_WIDTH),\n        scale=(0.6, 1.0),  # Crop 60%-100% of original area\n        ratio=(0.75, 1.33)  # Maintain reasonable aspect ratios\n    ),\n    transforms.RandomHorizontalFlip(p=0.5),  # Can be increased to 0.75 for stronger augmentation\n    transforms.RandomVerticalFlip(p=0.5),   # Can be increased to 0.75 for stronger augmentation\n    transforms.RandomRotation(degrees=30),  # Can be increased to 90 for stronger augmentation\n    \n    # Photometric augmentations\n    transforms.ColorJitter(\n        brightness=0.2,\n        contrast=0.2,\n        saturation=0.2,\n        hue=0.0  # Avoid hue shift for disease color consistency\n    ),\n    \n    # Convert to tensor and normalize\n    transforms.ToTensor(),\n    transforms.Normalize(mean=config.IMAGENET_MEAN, std=config.IMAGENET_STD)\n])\n\n# Validation/test transforms (no augmentation)\nval_transforms = transforms.Compose([\n    transforms.Resize((config.INPUT_HEIGHT, config.INPUT_WIDTH)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=config.IMAGENET_MEAN, std=config.IMAGENET_STD)\n])\n\nprint(\"Training transforms:\")\nprint(\"  - RandomResizedCrop (scale: 0.6-1.0)\")\nprint(\"  - RandomHorizontalFlip (p=0.5)\")\nprint(\"  - RandomVerticalFlip (p=0.5)\")\nprint(\"  - RandomRotation (±30°)\")\nprint(\"  - ColorJitter (brightness, contrast, saturation: ±0.2)\")\nprint(\"\\nValidation transforms:\")\nprint(\"  - Resize to (224, 224)\")\nprint(\"  - Normalize with ImageNet statistics\")\nprint(\"\\nNOTE: For stronger augmentation (like old notebook), you can change:\")\nprint(\"  - RandomHorizontalFlip(p=0.75)\")\nprint(\"  - RandomVerticalFlip(p=0.75)\")\nprint(\"  - RandomRotation(degrees=90)\")","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:19:55.994905Z","iopub.execute_input":"2026-01-07T09:19:55.995188Z","iopub.status.idle":"2026-01-07T09:19:56.004663Z","shell.execute_reply.started":"2026-01-07T09:19:55.995161Z","shell.execute_reply":"2026-01-07T09:19:56.003781Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Dataset and DataLoader\ndef encode_label(labels: List[str], class_list: List[str]) -> torch.Tensor:\n    \"\"\"\n    Encode a list of labels into multi-hot encoded tensor.\n\n    Args:\n        labels: List of label strings\n        class_list: Ordered list of all possible classes\n\n    Returns:\n        Multi-hot encoded tensor of shape (num_classes,)\n    \"\"\"\n    target = torch.zeros(len(class_list))\n    for label in labels:\n        idx = class_list.index(label)\n        target[idx] = 1\n    return target\n\nclass PlantDataset(Dataset):\n    \"\"\"\n    Custom Dataset for Plant Pathology images.\n\n    Args:\n        dataframe: DataFrame with 'image' and 'labels' columns\n        transform: torchvision.transforms to apply to images\n    \"\"\"\n    def __init__(self, dataframe: pd.DataFrame, transform=None):\n        self.dataframe = dataframe\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        # Load image\n        image_path = self.dataframe['image'].iloc[idx]\n        image = Image.open(image_path).convert('RGB')\n\n        # Parse labels\n        labels = self.dataframe.iloc[idx]['labels'].split(' ')\n        encoded_labels = encode_label(labels, config.LABELS)\n\n        # Apply transforms\n        if self.transform:\n            image = self.transform(image)\n\n        return image, encoded_labels\n\n# Create datasets\ntrain_dataset = PlantDataset(train_df, transform=train_transforms)\nval_dataset = PlantDataset(val_df, transform=val_transforms)\ntest_dataset = PlantDataset(test_df, transform=val_transforms)\n\n# Create dataloaders\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=config.BATCH_SIZE,\n    shuffle=True,\n    num_workers=config.NUM_WORKERS,\n    pin_memory=True if config.DEVICE == 'cuda' else False\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=config.BATCH_SIZE,\n    shuffle=False,\n    num_workers=config.NUM_WORKERS,\n    pin_memory=True if config.DEVICE == 'cuda' else False\n)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=config.BATCH_SIZE,\n    shuffle=False,\n    num_workers=config.NUM_WORKERS,\n    pin_memory=True if config.DEVICE == 'cuda' else False\n)\n\nprint(f\"Datasets and DataLoaders created!\")\nprint(f\"Train batches: {len(train_loader)}\")\nprint(f\"Validation batches: {len(val_loader)}\")\nprint(f\"Test batches: {len(test_loader)}\")","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:19:56.006014Z","iopub.execute_input":"2026-01-07T09:19:56.006612Z","iopub.status.idle":"2026-01-07T09:19:56.027457Z","shell.execute_reply.started":"2026-01-07T09:19:56.006576Z","shell.execute_reply":"2026-01-07T09:19:56.026543Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Model Architecture: CustomPlantCNN\nclass SimpleResBlock(nn.Module):\n    \"\"\"\n    A simple residual block with two 3x3 convolutions and a skip connection.\n\n    Structure:\n        Input -> Conv3x3 -> BN -> ReLU -> Conv3x3 -> BN -> Add Skip -> ReLU -> Output\n                                          |                              ^\n                                          |______________________________|\n\n    Args:\n        in_channels: Number of input feature maps\n        out_channels: Number of output feature maps\n        stride: Stride for first convolution (1 or 2 for downsampling)\n    \"\"\"\n    def __init__(self, in_channels: int, out_channels: int, stride: int = 1):\n        super(SimpleResBlock, self).__init__()\n\n        # First convolution\n        self.conv1 = nn.Conv2d(\n            in_channels, out_channels, kernel_size=3,\n            stride=stride, padding=1, bias=False\n        )\n        self.bn1 = nn.BatchNorm2d(out_channels)\n\n        # Second convolution\n        self.conv2 = nn.Conv2d(\n            out_channels, out_channels, kernel_size=3,\n            stride=1, padding=1, bias=False\n        )\n        self.bn2 = nn.BatchNorm2d(out_channels)\n\n        # Skip connection (residual pathway)\n        # If dimensions change, use 1x1 conv to match them\n        self.shortcut = nn.Sequential()\n        if stride != 1 or in_channels != out_channels:\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, kernel_size=1,\n                         stride=stride, bias=False),\n                nn.BatchNorm2d(out_channels)\n            )\n\n    def forward(self, x):\n        # Main pathway\n        out = F.relu(self.bn1(self.conv1(x)))\n        out = self.bn2(self.conv2(out))\n\n        # Add skip connection\n        out += self.shortcut(x)\n\n        # Final activation after addition\n        out = F.relu(out)\n        return out\n\n\nclass CustomPlantCNN(nn.Module):\n    \"\"\"\n    Custom CNN for Plant Pathology multi-label classification.\n\n    Architecture:\n        1. Initial Conv (7x7, stride=2) + BN + ReLU + MaxPool\n        2. Four SimpleResBlocks with progressive channel expansion:\n           - Layer1: 64 -> 64 channels\n           - Layer2: 64 -> 128 channels (stride=2, downsample)\n           - Layer3: 128 -> 256 channels (stride=2, downsample)\n           - Layer4: 256 -> 512 channels (stride=2, downsample)\n        3. Global Average Pooling\n        4. Dropout (p=0.5)\n        5. Fully Connected layer (512 -> num_classes)\n\n    Parameters: ~3.5M trainable parameters\n\n    Args:\n        num_classes: Number of output classes (6 for Plant Pathology)\n    \"\"\"\n    def __init__(self, num_classes: int = 6):\n        super(CustomPlantCNN, self).__init__()\n\n        # Initial convolution (stem)\n        # Reduces spatial dimensions: 224x224 -> 112x112 -> 56x56\n        self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        self.bn1 = nn.BatchNorm2d(64)\n        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n\n        # Residual blocks (feature extraction)\n        self.layer1 = SimpleResBlock(64, 64, stride=1)    # 56x56\n        self.layer2 = SimpleResBlock(64, 128, stride=2)   # 28x28\n        self.layer3 = SimpleResBlock(128, 256, stride=2)  # 14x14\n        self.layer4 = SimpleResBlock(256, 512, stride=2)  # 7x7\n\n        # Classification head\n        self.avgpool = nn.AdaptiveAvgPool2d((1, 1))  # Global pooling: 7x7 -> 1x1\n        self.dropout = nn.Dropout(p=0.5)\n        self.fc = nn.Linear(512, num_classes)\n\n    def forward(self, x):\n        # Stem\n        x = F.relu(self.bn1(self.conv1(x)))\n        x = self.maxpool(x)\n\n        # Feature extraction\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n\n        # Classification\n        x = self.avgpool(x)           # Shape: (B, 512, 1, 1)\n        x = torch.flatten(x, 1)       # Shape: (B, 512)\n        x = self.dropout(x)\n        x = self.fc(x)                # Shape: (B, num_classes) - LOGITS\n\n        return x\n\n# Initialize model\nmodel = CustomPlantCNN(num_classes=config.NUM_CLASSES).to(config.DEVICE)\n\n# Count parameters\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n\nprint(f\"CustomPlantCNN initialized!\")\nprint(f\"Total parameters: {total_params:,}\")\nprint(f\"Trainable parameters: {trainable_params:,}\")\nprint(f\"\\nModel architecture:\")\nprint(model)","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:19:56.028715Z","iopub.execute_input":"2026-01-07T09:19:56.029281Z","iopub.status.idle":"2026-01-07T09:19:58.608255Z","shell.execute_reply.started":"2026-01-07T09:19:56.029243Z","shell.execute_reply":"2026-01-07T09:19:58.607296Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Loss Function with Class Weighting\ncriterion = nn.BCEWithLogitsLoss(pos_weight=class_weights.to(config.DEVICE))\n\n# Optimizer: AdamW (Adam with decoupled weight decay)\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=config.LEARNING_RATE,\n    weight_decay=config.WEIGHT_DECAY\n)\n\n# Learning Rate Scheduler: ReduceLROnPlateau\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode='min',        # Minimize validation loss\n    factor=0.1,        # Reduce LR by factor of 10\n    patience=config.PATIENCE,  # Wait 3 epochs before reducing\n    verbose=True,\n    threshold=1e-4\n)\n\nprint(\"Training components configured:\")\nprint(f\"  Loss: BCEWithLogitsLoss with pos_weight\")\nprint(f\"  Optimizer: AdamW (lr={config.LEARNING_RATE}, weight_decay={config.WEIGHT_DECAY})\")\nprint(f\"  Scheduler: ReduceLROnPlateau (patience={config.PATIENCE}, factor=0.1)\")","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:19:58.609645Z","iopub.execute_input":"2026-01-07T09:19:58.610110Z","iopub.status.idle":"2026-01-07T09:19:58.617118Z","shell.execute_reply.started":"2026-01-07T09:19:58.610070Z","shell.execute_reply":"2026-01-07T09:19:58.616153Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Training and Evaluation Functions\ndef train_one_epoch(model, dataloader, criterion, optimizer, device):\n    \"\"\"\n    Train model for one epoch.\n    \n    Args:\n        model: Neural network model\n        dataloader: Training data loader\n        criterion: Loss function\n        optimizer: Optimizer\n        device: Device to train on\n    \n    Returns:\n        avg_loss: Average training loss\n        avg_acc: Average training accuracy (element-wise for multi-label)\n        avg_f1: Average training F1 score (macro, better metric for multi-label)\n    \"\"\"\n    model.train()\n    running_loss = 0.0\n    running_correct = 0\n    total_preds = 0\n    total_samples = 0\n    \n    all_preds_list = []\n    all_labels_list = []\n    \n    progress_bar = tqdm(dataloader, desc='Training', leave=False)\n    \n    for images, labels in progress_bar:\n        images, labels = images.to(device), labels.to(device)\n        \n        # Forward pass\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        \n        # Backward pass\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        # Statistics\n        batch_size = images.size(0)\n        running_loss += loss.item() * batch_size\n        total_samples += batch_size\n        \n        # Calculate element-wise accuracy (each label prediction counts separately)\n        preds = (torch.sigmoid(outputs) > 0.5).float()\n        running_correct += (preds == labels).float().sum()\n        total_preds += labels.numel()  # Total number of individual label predictions\n        \n        # Store for F1 calculation\n        all_preds_list.append(preds.cpu().numpy())\n        all_labels_list.append(labels.cpu().numpy())\n        \n        # Update progress bar\n        progress_bar.set_postfix({\n            'loss': f'{loss.item():.4f}',\n            'acc': f'{(preds == labels).float().mean().item():.4f}'\n        })\n    \n    avg_loss = running_loss / total_samples\n    avg_acc = running_correct / total_preds\n    \n    # Calculate macro F1 score (better metric for multi-label)\n    all_preds = np.vstack(all_preds_list)\n    all_labels = np.vstack(all_labels_list)\n    avg_f1 = f1_score(all_labels, all_preds, average='macro', zero_division=0)\n    \n    return avg_loss, avg_acc, avg_f1\n\n\ndef validate(model, dataloader, criterion, device, return_probs=False):\n    \"\"\"\n    Validate model on validation set.\n    \n    Args:\n        model: Neural network model\n        dataloader: Validation data loader\n        criterion: Loss function\n        device: Device to validate on\n        return_probs: If True, return probabilities along with binary predictions\n    \n    Returns:\n        avg_loss: Average validation loss\n        avg_acc: Average validation accuracy (element-wise for multi-label)\n        avg_f1: Average validation F1 score (macro)\n        all_preds: All predictions (numpy array)\n        all_labels: All true labels (numpy array)\n        all_probs: All probabilities (numpy array) - only if return_probs=True\n    \"\"\"\n    model.eval()\n    running_loss = 0.0\n    running_correct = 0\n    total_preds = 0\n    total_samples = 0\n    \n    all_preds = []\n    all_labels = []\n    all_probs = []  # Store probabilities if requested\n    \n    with torch.no_grad():\n        for images, labels in tqdm(dataloader, desc='Validation', leave=False):\n            images, labels = images.to(device), labels.to(device)\n            \n            # Forward pass\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            \n            # Statistics\n            batch_size = images.size(0)\n            running_loss += loss.item() * batch_size\n            total_samples += batch_size\n            \n            # Predictions - get probabilities FIRST\n            probs = torch.sigmoid(outputs)\n            preds = (probs > 0.5).float()\n            \n            # Store for later analysis\n            all_preds.append(preds.cpu().numpy())\n            all_labels.append(labels.cpu().numpy())\n            if return_probs:\n                all_probs.append(probs.cpu().numpy())  # Save probabilities\n            \n            # Element-wise accuracy (each label prediction counts separately)\n            running_correct += (preds == labels).float().sum()\n            total_preds += labels.numel()\n    \n    avg_loss = running_loss / total_samples\n    avg_acc = running_correct / total_preds\n    \n    # Calculate macro F1 score (better metric for multi-label)\n    all_preds_np = np.vstack(all_preds)\n    all_labels_np = np.vstack(all_labels)\n    avg_f1 = f1_score(all_labels_np, all_preds_np, average='macro', zero_division=0)\n    \n    if return_probs:\n        all_probs = np.vstack(all_probs)\n        return avg_loss, avg_acc, avg_f1, all_preds_np, all_labels_np, all_probs\n    else:\n        return avg_loss, avg_acc, avg_f1, all_preds_np, all_labels_np\n\nprint(\"Training and evaluation functions defined!\")\nprint(\"\\n⚠️  NOTE: Accuracy in multi-label classification can be misleading!\")\nprint(\"   - Element-wise accuracy: Includes correct negative predictions (easy)\")\nprint(\"   - F1 score: Balances precision & recall (better metric for imbalanced data)\")\nprint(\"   - We track BOTH to monitor training properly\")","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:19:58.618418Z","iopub.execute_input":"2026-01-07T09:19:58.618712Z","iopub.status.idle":"2026-01-07T09:19:58.638032Z","shell.execute_reply.started":"2026-01-07T09:19:58.618687Z","shell.execute_reply":"2026-01-07T09:19:58.637281Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Loop\n\nThe model will be trained for **30 epochs** with the following protocol:\n\n1. **Forward pass**: Compute logits for each batch\n2. **Loss calculation**: BCEWithLogitsLoss with class-weighted positives\n3. **Backward pass**: Compute gradients via backpropagation\n4. **Optimizer step**: Update parameters via AdamW\n5. **Validation**: Evaluate on validation set after each epoch\n6. **Learning rate scheduling**: Reduce LR if validation loss plateaus\n7. **Checkpointing**: Save best model based on validation loss\n\n### Key Metrics Tracked:\n- **Training Loss**: BCE loss on training batches\n- **Training Accuracy**: Exact match accuracy (all labels correct)\n- **Validation Loss**: BCE loss on validation set\n- **Validation Accuracy**: Exact match accuracy on validation set\n- **Learning Rate**: Current learning rate after scheduling","metadata":{}},{"cell_type":"code","source":"# Training Loop\nhistory = {\n    'train_loss': [],\n    'train_acc': [],\n    'train_f1': [],\n    'val_loss': [],\n    'val_acc': [],\n    'val_f1': [],\n    'lr': []\n}\n\nbest_val_f1 = float('-inf')  # FIXED: Changed from float('inf') to float('-inf') for maximizing F1\nbest_epoch = 0\n\nprint(\"Starting training...\")\nprint(f\"Device: {config.DEVICE}\")\nprint(f\"Epochs: {config.NUM_EPOCHS}\")\nprint(f\"Batch size: {config.BATCH_SIZE}\")\nprint(f\"Initial learning rate: {config.LEARNING_RATE}\")\nprint(\"-\" * 80)\nprint(\"⚠️  Monitoring both Accuracy and F1:\")\nprint(\"   - Accuracy: Will be high (70-90%) due to easy negative predictions\")\nprint(\"   - F1 Score: Better indicator of actual performance\")\nprint(\"   - Model selected based on validation F1 (not accuracy)\")\nprint(\"-\" * 80)\n\nfor epoch in range(config.NUM_EPOCHS):\n    print(f\"\\nEpoch {epoch+1}/{config.NUM_EPOCHS}\")\n    \n    # Train\n    train_loss, train_acc, train_f1 = train_one_epoch(\n        model, train_loader, criterion, optimizer, config.DEVICE\n    )\n    \n    # Validate\n    val_loss, val_acc, val_f1, _, _ = validate(\n        model, val_loader, criterion, config.DEVICE\n    )\n    \n    # Get current learning rate\n    current_lr = optimizer.param_groups[0]['lr']\n    \n    # Update scheduler\n    scheduler.step(val_loss)\n    \n    # FIXED: Convert tensors to Python native types for JSON serialization\n    history['train_loss'].append(float(train_loss))\n    history['train_acc'].append(float(train_acc) if torch.is_tensor(train_acc) else train_acc)\n    history['train_f1'].append(float(train_f1))\n    history['val_loss'].append(float(val_loss))\n    history['val_acc'].append(float(val_acc) if torch.is_tensor(val_acc) else val_acc)\n    history['val_f1'].append(float(val_f1))\n    history['lr'].append(float(current_lr))\n    \n    # Print epoch summary\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | Train F1: {train_f1:.4f}\")\n    print(f\"Val Loss:   {val_loss:.4f} | Val Acc:   {val_acc:.4f} | Val F1:   {val_f1:.4f}\")\n    print(f\"LR: {current_lr:.2e}\")\n    \n    # Save best model based on F1 (better metric for multi-label)\n    if val_f1 > best_val_f1:\n        best_val_f1 = val_f1\n        best_epoch = epoch + 1\n        \n        # Save checkpoint\n        checkpoint = {\n            'epoch': epoch + 1,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'scheduler_state_dict': scheduler.state_dict(),\n            'val_loss': float(val_loss),\n            'val_acc': float(val_acc) if torch.is_tensor(val_acc) else val_acc,\n            'val_f1': float(val_f1),\n            'history': history\n        }\n        torch.save(checkpoint, config.CHECKPOINT_DIR / 'best_model.pth')\n        print(f\"✓ Saved best model (validation F1: {val_f1:.4f})\")\n\nprint(\"\\n\" + \"-\" * 80)\nprint(f\"Training complete!\")\nprint(f\"Best epoch: {best_epoch}\")\nprint(f\"Best validation F1: {best_val_f1:.4f}\")\n\n# FIXED: Save training history with proper type conversion\n# All values are now Python native types (float), not tensors\nwith open(config.METRICS_DIR / 'training_history.json', 'w') as f:\n    json.dump(history, f, indent=2)\nprint(f\"Training history saved to {config.METRICS_DIR / 'training_history.json'}\")","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:19:58.639265Z","iopub.execute_input":"2026-01-07T09:19:58.639745Z","iopub.status.idle":"2026-01-07T09:48:09.595359Z","shell.execute_reply.started":"2026-01-07T09:19:58.639719Z","shell.execute_reply":"2026-01-07T09:48:09.594212Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Curves\n\nThe following plots visualize the training dynamics:\n\n### 1. Loss Curves\n- **Training Loss**: Should decrease monotonically as model learns\n- **Validation Loss**: Should decrease but may plateau or increase if overfitting\n- **Gap between curves**: Indicates overfitting if large\n\n### 2. Accuracy Curves  \n- **Training Accuracy**: Proportion of training samples with all labels correct\n- **Validation Accuracy**: Proportion of validation samples with all labels correct\n- **Convergence**: When both stabilize, model has learned the data distribution\n\n### 3. Learning Rate Schedule\n- Shows when ReduceLROnPlateau reduced the learning rate\n- Steps indicate validation loss plateaued for 3 consecutive epochs","metadata":{}},{"cell_type":"code","source":"# Plot Training Curves\ndef plot_training_curves(history, save_path):\n    \"\"\"\n    Plot training and validation loss/accuracy/F1 curves.\n\n    Args:\n        history: Dictionary with training history\n        save_path: Path to save the plot\n    \"\"\"\n    epochs = range(1, len(history['train_loss']) + 1)\n\n    # UPDATED: Changed to 3 subplots including F1 score\n    fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\n    # Plot 1: Loss\n    axes[0].plot(epochs, history['train_loss'], 'b-', label='Train Loss', linewidth=2)\n    axes[0].plot(epochs, history['val_loss'], 'r-', label='Val Loss', linewidth=2)\n    axes[0].set_xlabel('Epoch', fontsize=12)\n    axes[0].set_ylabel('Loss', fontsize=12)\n    axes[0].set_title('Training and Validation Loss', fontsize=14, fontweight='bold')\n    axes[0].legend(fontsize=11)\n    axes[0].grid(True, alpha=0.3)\n\n    # Plot 2: Accuracy\n    axes[1].plot(epochs, history['train_acc'], 'b-', label='Train Acc', linewidth=2)\n    axes[1].plot(epochs, history['val_acc'], 'r-', label='Val Acc', linewidth=2)\n    axes[1].set_xlabel('Epoch', fontsize=12)\n    axes[1].set_ylabel('Accuracy', fontsize=12)\n    axes[1].set_title('Training and Validation Accuracy', fontsize=14, fontweight='bold')\n    axes[1].legend(fontsize=11)\n    axes[1].grid(True, alpha=0.3)\n\n    # UPDATED: Plot 3: F1 Score (more important than LR for multi-label)\n    axes[2].plot(epochs, history['train_f1'], 'b-', label='Train F1', linewidth=2)\n    axes[2].plot(epochs, history['val_f1'], 'r-', label='Val F1', linewidth=2)\n    axes[2].set_xlabel('Epoch', fontsize=12)\n    axes[2].set_ylabel('F1 Score (Macro)', fontsize=12)\n    axes[2].set_title('Training and Validation F1 Score', fontsize=14, fontweight='bold')\n    axes[2].legend(fontsize=11)\n    axes[2].grid(True, alpha=0.3)\n\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=300, bbox_inches='tight')\n    plt.show()\n    print(f\"Training curves saved to {save_path}\")\n\n# Plot curves\nplot_training_curves(history, config.PLOT_DIR / 'baseline_training_curves.png')","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:48:09.597131Z","iopub.execute_input":"2026-01-07T09:48:09.598016Z","iopub.status.idle":"2026-01-07T09:48:11.414662Z","shell.execute_reply.started":"2026-01-07T09:48:09.597971Z","shell.execute_reply":"2026-01-07T09:48:11.413742Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Test Set Evaluation (Default Threshold = 0.5)\n\nBefore any threshold optimization, we evaluate the baseline model on the test set using the default 0.5 threshold for all classes. This gives us a **baseline performance metric** to compare against after threshold optimization.\n\n### Evaluation Metrics:\n- **Exact Match Accuracy**: Percentage of samples where ALL labels are predicted correctly\n- **Macro F1 Score**: Average F1 score across all classes (treats all classes equally)\n- **Per-Class Precision/Recall/F1**: Detailed metrics for each disease class","metadata":{}},{"cell_type":"code","source":"# Load Best Model for Testing\ncheckpoint = torch.load(config.CHECKPOINT_DIR / 'best_model.pth')\nmodel.load_state_dict(checkpoint['model_state_dict'])\nprint(f\"Loaded best model from epoch {checkpoint['epoch']} (val_loss={checkpoint['val_loss']:.4f})\")\n\n# FIXED: validate() now returns 5 values (loss, acc, f1, preds, labels)\ntest_loss, test_acc, test_f1, test_preds, test_labels = validate(\n    model, test_loader, criterion, config.DEVICE\n)\n\n# Note: test_f1 is already calculated by validate(), but we'll recalculate for clarity\nfrom sklearn.metrics import f1_score, precision_score, recall_score\n\nmacro_f1 = f1_score(test_labels, test_preds, average='macro', zero_division=0)\nmicro_f1 = f1_score(test_labels, test_preds, average='micro', zero_division=0)\nmacro_precision = precision_score(test_labels, test_preds, average='macro', zero_division=0)\nmacro_recall = recall_score(test_labels, test_preds, average='macro', zero_division=0)\n\nprint(\"\\n\" + \"=\" * 60)\nprint(\"TEST SET PERFORMANCE SUMMARY (Default Threshold = 0.5)\")\nprint(\"=\" * 60)\nprint(f\"Test Loss:        {test_loss:.4f}\")\nprint(f\"Exact Match Acc:  {test_acc:.4f}\")\nprint(f\"Macro F1:          {macro_f1:.4f}\")\nprint(f\"Micro F1:          {micro_f1:.4f}\")\nprint(f\"Macro Precision:   {macro_precision:.4f}\")\nprint(f\"Macro Recall:      {macro_recall:.4f}\")\nprint(\"=\" * 60)\n\n# Detailed per-class report\nprint(\"\\nPer-Class Classification Report:\")\nprint(\"-\" * 60)\nreport = classification_report(\n    test_labels, test_preds,\n    target_names=config.LABELS,\n    zero_division=0,\n    digits=4\n)\nprint(report)\n\n# Save metrics\nbaseline_metrics = {\n    'test_loss': float(test_loss),\n    'exact_match_accuracy': float(test_acc),\n    'macro_f1': float(macro_f1),\n    'micro_f1': float(micro_f1),\n    'macro_precision': float(macro_precision),\n    'macro_recall': float(macro_recall),\n    'threshold': 0.5,\n    'per_class_metrics': classification_report(\n        test_labels, test_preds,\n        target_names=config.LABELS,\n        zero_division=0,\n        output_dict=True\n    )\n}\n\nwith open(config.METRICS_DIR / 'baseline_test_metrics.json', 'w') as f:\n    json.dump(baseline_metrics, f, indent=2)\n\nprint(f\"\\nMetrics saved to {config.METRICS_DIR / 'baseline_test_metrics.json'}\")","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:57:25.415406Z","iopub.execute_input":"2026-01-07T09:57:25.415790Z","iopub.status.idle":"2026-01-07T09:57:31.207447Z","shell.execute_reply.started":"2026-01-07T09:57:25.415758Z","shell.execute_reply":"2026-01-07T09:57:31.206311Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Collect Baseline Predictions with Probabilities\ndef collect_predictions_and_probabilities(model, test_loader, device):\n    \"\"\"\n    Collect predictions AND probabilities for threshold optimization.\n\n    Returns:\n        all_labels: True labels (multi-hot)\n        all_preds: Binary predictions (threshold=0.5)\n        all_probs: Raw probabilities (before thresholding)\n    \"\"\"\n    model.eval()\n    all_labels = []\n    all_preds = []\n    all_probs = []\n\n    with torch.no_grad():\n        for images, labels in tqdm(test_loader, desc='Collecting predictions', leave=False):\n            images = images.to(device)\n            labels_np = labels.numpy()\n\n            # Forward pass\n            outputs = model(images)\n            probs = torch.sigmoid(outputs).cpu().numpy()  # PROBABILITIES\n            preds = (probs > 0.5).astype(float)            # BINARY PREDICTIONS\n\n            all_labels.append(labels_np)\n            all_preds.append(preds)\n            all_probs.append(probs)\n\n    all_labels = np.vstack(all_labels)\n    all_preds = np.vstack(all_preds)\n    all_probs = np.vstack(all_probs)\n\n    return all_labels, all_preds, all_probs\n\n# Collect baseline predictions with probabilities\ntest_labels, test_preds_baseline, test_probs = collect_predictions_and_probabilities(\n    model, test_loader, config.DEVICE\n)\n\nprint(f\"Collected {len(test_labels)} test samples\")\nprint(f\"Labels shape: {test_labels.shape}\")\nprint(f\"Predictions shape: {test_preds_baseline.shape}\")\nprint(f\"Probabilities shape: {test_probs.shape}\")\nprint(\"\\nThese will be used for threshold optimization and comparison.\")","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:57:51.710812Z","iopub.execute_input":"2026-01-07T09:57:51.711594Z","iopub.status.idle":"2026-01-07T09:57:57.369979Z","shell.execute_reply.started":"2026-01-07T09:57:51.711556Z","shell.execute_reply":"2026-01-07T09:57:57.368829Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Threshold Optimization with Comparison Table","metadata":{}},{"cell_type":"code","source":"def find_optimal_thresholds_with_comparison(model, val_loader, device, class_names):\n    \"\"\"\n    Find optimal threshold for each class using validation set.\n    Shows comparison with default 0.5 threshold.\n\n    Returns:\n        optimal_thresholds: Dictionary mapping class names to optimal thresholds\n        threshold_scores: Dictionary mapping class names to best F1 scores\n        comparison_df: DataFrame with detailed comparison\n    \"\"\"\n    print(\"Finding optimal thresholds...\")\n    print(\"-\" * 85)\n    print(f\"{'Class':<25} {'Opt Thresh':<10} {'Best F1':<10} {'Default F1':<10} {'Improvement':<12}\")\n    print(\"-\" * 85)\n\n    # Collect all predictions\n    model.eval()\n    all_probs = []\n    all_labels = []\n\n    with torch.no_grad():\n        for images, labels in tqdm(val_loader, desc='Collecting val predictions', leave=False):\n            images = images.to(device)\n            outputs = model(images)\n            probs = torch.sigmoid(outputs).cpu().numpy()\n\n            all_probs.append(probs)\n            all_labels.append(labels.numpy())\n\n    all_probs = np.vstack(all_probs)\n    all_labels = np.vstack(all_labels)\n\n    # Find optimal threshold for each class\n    optimal_thresholds = {}\n    threshold_scores = {}\n    comparison_data = []\n    threshold_range = np.arange(0.15, 0.86, 0.02)  # Finer granularity: 0.15 to 0.85\n\n    for i, class_name in enumerate(class_names):\n        best_f1 = 0\n        best_thresh = 0.5\n\n        # Test each threshold\n        for thresh in threshold_range:\n            preds = (all_probs[:, i] > thresh).astype(float)\n            f1 = f1_score(all_labels[:, i], preds, zero_division=0)\n\n            if f1 > best_f1:\n                best_f1 = f1\n                best_thresh = thresh\n\n        # Calculate F1 with default 0.5 threshold\n        default_preds = (all_probs[:, i] > 0.5).astype(float)\n        default_f1 = f1_score(all_labels[:, i], default_preds, zero_division=0)\n\n        # Calculate improvement\n        improvement = 0\n        if default_f1 > 0:\n            improvement = ((best_f1 - default_f1) / default_f1) * 100\n\n        optimal_thresholds[class_name] = float(best_thresh)\n        threshold_scores[class_name] = float(best_f1)\n\n        # Add marker for significant improvements\n        if improvement > 5:\n            marker = \" ⭐\"\n        elif improvement > 0:\n            marker = \" ✓\"\n        else:\n            marker = \"\"\n\n        print(f\"{class_name:<25} {best_thresh:<10.3f} {best_f1:<10.4f} {default_f1:<10.4f} {improvement:>+10.2f}%{marker}\")\n\n        # Store for comparison DataFrame\n        comparison_data.append({\n            'Class': class_name,\n            'Default_Thresh': 0.5,\n            'Optimal_Thresh': best_thresh,\n            'Default_F1': default_f1,\n            'Optimal_F1': best_f1,\n            'Improvement_%': improvement\n        })\n\n    print(\"-\" * 85)\n\n    # Create comparison DataFrame\n    comparison_df = pd.DataFrame(comparison_data)\n\n    return optimal_thresholds, threshold_scores, comparison_df\n\n# Find optimal thresholds with comparison\noptimal_thresholds, threshold_scores, comparison_df = find_optimal_thresholds_with_comparison(\n    model, val_loader, config.DEVICE, config.LABELS\n)\n\n# Save comparison to CSV\ncomparison_df.to_csv(config.METRICS_DIR / 'threshold_optimization_comparison.csv', index=False)\nprint(f\"\\nComparison saved to {config.METRICS_DIR / 'threshold_optimization_comparison.csv'}\")\n\n# Save thresholds\nwith open(config.METRICS_DIR / 'optimal_thresholds.json', 'w') as f:\n    json.dump({\n        'thresholds': optimal_thresholds,\n        'scores': threshold_scores\n    }, f, indent=2)\nprint(f\"Optimal thresholds saved to {config.METRICS_DIR / 'optimal_thresholds.json'}\")","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:58:03.868427Z","iopub.execute_input":"2026-01-07T09:58:03.868805Z","iopub.status.idle":"2026-01-07T09:58:09.873591Z","shell.execute_reply.started":"2026-01-07T09:58:03.868771Z","shell.execute_reply":"2026-01-07T09:58:09.872431Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Threshold Comparison Visualization\ndef plot_threshold_comparison(comparison_df, save_path):\n    \"\"\"\n    Create side-by-side comparison plots for threshold optimization.\n\n    Plot 1: F1 Score comparison (default vs optimal)\n    Plot 2: Optimal thresholds per class (with color coding)\n    \"\"\"\n    fig, axes = plt.subplots(1, 2, figsize=(16, 6))\n\n    # Plot 1: F1 Score Comparison\n    x = np.arange(len(comparison_df))\n    width = 0.35\n\n    axes[0].bar(\n        x - width/2,\n        comparison_df['Default_F1'] * 100,\n        width,\n        label='Default (0.5)',\n        alpha=0.8,\n        color='coral'\n    )\n    axes[0].bar(\n        x + width/2,\n        comparison_df['Optimal_F1'] * 100,\n        width,\n        label='Optimal Thresholds',\n        alpha=0.8,\n        color='steelblue'\n    )\n\n    axes[0].set_xlabel('Class', fontsize=12)\n    axes[0].set_ylabel('F1 Score (%)', fontsize=12)\n    axes[0].set_title('F1 Score: Default vs Optimal Thresholds', fontsize=14, fontweight='bold')\n    axes[0].set_xticks(x)\n    axes[0].set_xticklabels(comparison_df['Class'], rotation=45, ha='right')\n    axes[0].legend(fontsize=11)\n    axes[0].grid(axis='y', alpha=0.3)\n    axes[0].set_ylim(0, 105)\n\n    # Add improvement annotations\n    for i, (default_f1, optimal_f1) in enumerate(zip(comparison_df['Default_F1'], comparison_df['Optimal_F1'])):\n        delta = (optimal_f1 - default_f1) * 100\n        if abs(delta) > 3:\n            axes[0].annotate(\n                f'{delta:+.1f}%',\n                xy=(i + width/2, optimal_f1 * 100),\n                xytext=(i + width/2, optimal_f1 * 100 + 5),\n                ha='center',\n                fontsize=9,\n                fontweight='bold',\n                color='green' if delta > 0 else 'red'\n            )\n\n    # Plot 2: Optimal Thresholds per Class\n    colors = [\n        'green' if t < 0.5\n        else 'orange' if t < 0.6\n        else 'red'\n        for t in comparison_df['Optimal_Thresh']\n    ]\n\n    bars = axes[1].bar(\n        comparison_df['Class'],\n        comparison_df['Optimal_Thresh'],\n        color=colors,\n        alpha=0.7\n    )\n\n    axes[1].axhline(\n        y=0.5,\n        color='black',\n        linestyle='--',\n        linewidth=2,\n        label='Default (0.5)'\n    )\n\n    axes[1].set_xlabel('Class', fontsize=12)\n    axes[1].set_ylabel('Optimal Threshold', fontsize=12)\n    axes[1].set_title('Optimal Threshold per Class', fontsize=14, fontweight='bold')\n    axes[1].set_xticks(x)\n    axes[1].set_xticklabels(comparison_df['Class'], rotation=45, ha='right')\n    axes[1].legend(fontsize=11)\n    axes[1].grid(axis='y', alpha=0.3)\n    axes[1].set_ylim(0, 1.0)\n\n    # Add threshold values on bars\n    for bar, thresh in zip(bars, comparison_df['Optimal_Thresh']):\n        height = bar.get_height()\n        axes[1].text(\n            bar.get_x() + bar.get_width()/2.,\n            height,\n            f'{thresh:.2f}',\n            ha='center',\n            va='bottom',\n            fontsize=9\n        )\n\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=300, bbox_inches='tight')\n    plt.show()\n    print(f\"Threshold comparison plot saved to {save_path}\")\n\n# Create comparison plot\nplot_threshold_comparison(\n    comparison_df,\n    config.PLOT_DIR / 'threshold_optimization_comparison.png'\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T09:58:30.392296Z","iopub.execute_input":"2026-01-07T09:58:30.393186Z","iopub.status.idle":"2026-01-07T09:58:31.808966Z","shell.execute_reply.started":"2026-01-07T09:58:30.393148Z","shell.execute_reply":"2026-01-07T09:58:31.808015Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Evaluation After Threshold Optimization\n\nNow we evaluate the model on the **test set** using the optimized thresholds. This demonstrates the **improvement gained by threshold calibration**.\n\n### Expected Improvements:\n- **Higher F1 scores**: Each class operates at its optimal precision-recall tradeoff\n- **Better rare-class performance**: Lower thresholds for rare classes (e.g., powdery_mildew) increase recall\n- **Balanced predictions**: Majority classes won't dominate predictions\n\n### Comparison Strategy:\nWe'll compare:\n1. **Baseline**: All thresholds = 0.5\n2. **Optimized**: Class-specific thresholds from validation set","metadata":{}},{"cell_type":"code","source":"# Evaluate with Optimized Thresholds\ndef evaluate_with_thresholds(model, test_loader, device, thresholds_dict, class_names):\n    \"\"\"\n    Evaluate model with custom thresholds for each class.\n\n    Args:\n        model: Trained model\n        test_loader: Test data loader\n        device: Device to run on\n        thresholds_dict: Dictionary mapping class names to thresholds\n        class_names: Ordered list of class names\n\n    Returns:\n        all_preds: Predictions using custom thresholds\n        all_labels: True labels\n    \"\"\"\n    model.eval()\n    all_preds = []\n    all_labels = []\n\n    # Convert thresholds dict to tensor in correct order\n    thresholds = torch.tensor([thresholds_dict[name] for name in class_names]).to(device)\n\n    with torch.no_grad():\n        for images, labels in tqdm(test_loader, desc='Testing with optimized thresholds', leave=False):\n            images = images.to(device)\n            outputs = model(images)\n            probs = torch.sigmoid(outputs)\n\n            # Apply class-specific thresholds\n            preds = (probs > thresholds.unsqueeze(0)).float()\n\n            all_preds.append(preds.cpu().numpy())\n            all_labels.append(labels.numpy())\n\n    return np.vstack(all_preds), np.vstack(all_labels)\n\n# Evaluate with optimized thresholds\nopt_preds, opt_labels = evaluate_with_thresholds(\n    model, test_loader, config.DEVICE, optimal_thresholds, config.LABELS\n)\n\n# Calculate metrics\nopt_macro_f1 = f1_score(opt_labels, opt_preds, average='macro', zero_division=0)\nopt_micro_f1 = f1_score(opt_labels, opt_preds, average='micro', zero_division=0)\nopt_acc = (opt_labels == opt_preds).all(axis=1).mean()\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TEST SET PERFORMANCE AFTER THRESHOLD OPTIMIZATION\")\nprint(\"=\" * 70)\nprint(f\"\\nOptimized Thresholds:\")\nfor name, thresh in optimal_thresholds.items():\n    print(f\"  {name:25s}: {thresh:.2f}\")\n\nprint(f\"\\nPerformance Metrics:\")\nprint(f\"  Exact Match Accuracy:  {opt_acc:.4f} (was {test_acc:.4f})\")\nprint(f\"  Change:               {opt_acc - test_acc:+.4f} ({(opt_acc/test_acc - 1)*100:+.2f}%)\")\nprint(f\"\\n  Macro F1:            {opt_macro_f1:.4f} (was {macro_f1:.4f})\")\nprint(f\"  Change:               {opt_macro_f1 - macro_f1:+.4f} ({(opt_macro_f1/macro_f1 - 1)*100:+.2f}%)\")\nprint(f\"\\n  Micro F1:            {opt_micro_f1:.4f} (was {micro_f1:.4f})\")\nprint(f\"  Change:               {opt_micro_f1 - micro_f1:+.4f} ({(opt_micro_f1/micro_f1 - 1)*100:+.2f}%)\")\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"THRESHOLD OPTIMIZATION IMPROVEMENT ANALYSIS\")\nprint(\"=\" * 70)\nprint(\"\\nHow threshold optimization helped:\")\nprint(\"1. Rare classes (e.g., powdery_mildew) now use LOWER thresholds → higher recall\")\nprint(\"2. Common classes (e.g., scab) may use HIGHER thresholds → higher precision\")\nprint(\"3. Each class operates at its optimal precision-recall tradeoff\")\nprint(\"4. Overall macro-F1 increases due to balanced per-class performance\")\n\n# Detailed comparison\nprint(\"\\nPer-Class Metrics Comparison:\")\nprint(\"-\" * 70)\n\nfor i, class_name in enumerate(config.LABELS):\n    default_pred = test_preds[:, i]\n    opt_pred = opt_preds[:, i]\n    true = opt_labels[:, i]\n\n    default_f1 = f1_score(true, default_pred, zero_division=0)\n    opt_f1 = f1_score(true, opt_pred, zero_division=0)\n\n    default_prec = precision_score(true, default_pred, zero_division=0)\n    opt_prec = precision_score(true, opt_pred, zero_division=0)\n\n    default_rec = recall_score(true, default_pred, zero_division=0)\n    opt_rec = recall_score(true, opt_pred, zero_division=0)\n\n    print(f\"\\n{class_name}:\")\n    print(f\"  Threshold: 0.50 → {optimal_thresholds[class_name]:.2f}\")\n    print(f\"  F1:        {default_f1:.4f} → {opt_f1:.4f} ({opt_f1-default_f1:+.4f})\")\n    print(f\"  Precision: {default_prec:.4f} → {opt_prec:.4f} ({opt_prec-default_prec:+.4f})\")\n    print(f\"  Recall:    {default_rec:.4f} → {opt_rec:.4f} ({opt_rec-default_rec:+.4f})\")\n\n# Save optimized metrics\noptimized_metrics = {\n    'test_loss': float(test_loss),  # Same as baseline\n    'exact_match_accuracy': float(opt_acc),\n    'macro_f1': float(opt_macro_f1),\n    'micro_f1': float(opt_micro_f1),\n    'thresholds': optimal_thresholds,\n    'per_class_metrics': classification_report(\n        opt_labels, opt_preds,\n        target_names=config.LABELS,\n        zero_division=0,\n        output_dict=True\n    ),\n    'improvement_over_baseline': {\n        'accuracy': float(opt_acc - test_acc),\n        'macro_f1': float(opt_macro_f1 - macro_f1),\n        'micro_f1': float(opt_micro_f1 - micro_f1)\n    }\n}\n\nwith open(config.METRICS_DIR / 'optimized_test_metrics.json', 'w') as f:\n    json.dump(optimized_metrics, f, indent=2)\n\nprint(f\"\\nOptimized metrics saved to {config.METRICS_DIR / 'optimized_test_metrics.json'}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T09:58:47.576350Z","iopub.execute_input":"2026-01-07T09:58:47.576672Z","iopub.status.idle":"2026-01-07T09:58:53.199455Z","shell.execute_reply.started":"2026-01-07T09:58:47.576646Z","shell.execute_reply":"2026-01-07T09:58:53.198417Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Combined Confusion Matrix (Dominant Class Style)\n\nFor multi-label problems, a single confusion matrix can be confusing because each sample has multiple labels. To simplify visualization, we use the **dominant class approach**:\n\n- **Dominant True Label**: The class with the highest probability in the ground truth (or first if multiple)\n- **Dominant Predicted Label**: The class with the highest predicted probability\n\nThis gives us a **single-label confusion matrix** that shows which diseases are most commonly confused with each other.\n\n### Interpreting the Matrix:\n- **Diagonal elements**: Correct predictions (true positives)\n- **Off-diagonal elements**: Misclassifications (confusion between classes)\n- **Rows**: True labels\n- **Columns**: Predicted labels\n- **Darker colors**: Higher counts","metadata":{}},{"cell_type":"code","source":"# Save Classification Reports to Text Files\ndef save_classification_report(labels, preds, class_names, filepath, title):\n    \"\"\"\n    Generate and save classification report to text file.\n\n    Args:\n        labels: True labels\n        preds: Predicted labels\n        class_names: List of class names\n        filepath: Path to save the report\n        title: Title for the report\n    \"\"\"\n    from sklearn.metrics import classification_report\n\n    report = classification_report(\n        labels, preds,\n        target_names=class_names,\n        zero_division=0,\n        digits=4\n    )\n\n    with open(filepath, 'w') as f:\n        f.write(f\"{title}\\n\")\n        f.write(\"=\" * 80 + \"\\n\\n\")\n        f.write(report)\n        f.write(\"\\n\" + \"=\" * 80 + \"\\n\")\n\n    print(f\"Report saved to {filepath}\")\n    return report\n\n# Save baseline report (default 0.5 threshold)\nprint(\"Saving baseline classification report...\")\nbaseline_report = save_classification_report(\n    test_labels, test_preds_baseline, config.LABELS,\n    config.METRICS_DIR / 'baseline_classification_report.txt',\n    'BASELINE CLASSIFICATION REPORT (Default Threshold = 0.5)'\n)\nprint(\"\\nBaseline Report:\")\nprint(\"-\" * 80)\nprint(baseline_report)\nprint(\"-\" * 80)\n\n# Save optimized report (optimal thresholds)\nprint(\"\\nSaving optimized classification report...\")\noptimized_report = save_classification_report(\n    opt_labels, opt_preds, config.LABELS,\n    config.METRICS_DIR / 'optimized_classification_report.txt',\n    'OPTIMIZED CLASSIFICATION REPORT (Optimal Thresholds)'\n)\nprint(\"\\nOptimized Report:\")\nprint(\"-\" * 80)\nprint(optimized_report)\nprint(\"-\" * 80)","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:59:17.296469Z","iopub.execute_input":"2026-01-07T09:59:17.296845Z","iopub.status.idle":"2026-01-07T09:59:17.337185Z","shell.execute_reply.started":"2026-01-07T09:59:17.296812Z","shell.execute_reply":"2026-01-07T09:59:17.336135Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Combined Confusion Matrix (Dominant Class)\ndef plot_dominant_confusion_matrix(labels, preds, class_names, save_path):\n    \"\"\"\n    Plot confusion matrix for multi-label predictions using dominant class.\n\n    For multi-label, we select the dominant class (highest probability)\n    to create a single-label confusion matrix.\n\n    Args:\n        labels: True labels (multi-hot encoded)\n        preds: Predicted labels (multi-hot encoded)\n        class_names: List of class names\n        save_path: Path to save the plot\n    \"\"\"\n    # Get dominant class (argmax) for each sample\n    # For true labels: use first positive label (arbitrary but consistent)\n    # For predictions: use class with highest probability\n\n    # For simplicity, use argmax (works well if labels are mostly single)\n    y_true_indices = np.argmax(labels, axis=1)\n    y_pred_indices = np.argmax(preds, axis=1)\n\n    # Compute confusion matrix\n    cm = confusion_matrix(y_true_indices, y_pred_indices)\n\n    # Plot\n    plt.figure(figsize=(10, 8))\n    sns.heatmap(\n        cm, annot=True, fmt='d', cmap='Blues',\n        xticklabels=class_names, yticklabels=class_names,\n        cbar_kws={'label': 'Count'}\n    )\n    plt.xlabel('Predicted Label', fontsize=12, fontweight='bold')\n    plt.ylabel('True Label', fontsize=12, fontweight='bold')\n    plt.title('Combined Confusion Matrix (Dominant Class)', fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=300, bbox_inches='tight')\n    plt.show()\n\n    print(f\"Confusion matrix saved to {save_path}\")\n\n    # Print analysis\n    print(\"\\nConfusion Matrix Analysis:\")\n    print(\"-\" * 50)\n\n    for i, class_name in enumerate(class_names):\n        true_count = cm[i].sum()\n        correct = cm[i, i]\n        accuracy = correct / true_count if true_count > 0 else 0\n        print(f\"{class_name:25s}: {correct:3d}/{true_count:3d} correct ({accuracy:.2%})\")\n\n    return cm\n\n# Plot confusion matrix with optimized predictions\ncm = plot_dominant_confusion_matrix(\n    opt_labels, opt_preds, config.LABELS,\n    config.PLOT_DIR / 'baseline_confusion_matrix.png'\n)","metadata":{"execution":{"iopub.status.busy":"2026-01-07T09:59:22.940458Z","iopub.execute_input":"2026-01-07T09:59:22.941184Z","iopub.status.idle":"2026-01-07T09:59:24.239649Z","shell.execute_reply.started":"2026-01-07T09:59:22.941147Z","shell.execute_reply":"2026-01-07T09:59:24.238706Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Per-Class Multi-Label Confusion Matrices\n\nFor a **more detailed analysis**, we plot individual confusion matrices for each class. This treats each class as a binary classification problem:\n\n- **True Positive (TP)**: Sample has the class AND model predicts it\n- **False Positive (FP)**: Sample doesn't have the class BUT model predicts it\n- **True Negative (TN)**: Sample doesn't have the class AND model doesn't predict it\n- **False Negative (FN)**: Sample has the class BUT model doesn't predict it\n\n### Matrix Structure:\n```\n                Predicted\n              Negative  Positive\nActual Neg |    TN    |    FP    |\n       Pos |    FN    |    TP    |\n```\n\nThis reveals **which classes have high false positive rates** (over-prediction) vs **high false negative rates** (under-prediction).","metadata":{}},{"cell_type":"code","source":"# Per-Class Multi-Label Confusion Matrices\ndef plot_per_class_confusion_matrices(labels, preds, class_names, save_path):\n    \"\"\"\n    Plot individual binary confusion matrices for each class.\n\n    Args:\n        labels: True labels (multi-hot encoded)\n        preds: Predicted labels (multi-hot encoded)\n        class_names: List of class names\n        save_path: Path to save the combined plot\n    \"\"\"\n    # Compute multi-label confusion matrices\n    ml_cms = multilabel_confusion_matrix(labels, preds)\n\n    # Setup plot grid\n    n_classes = len(class_names)\n    n_cols = 3\n    n_rows = (n_classes + n_cols - 1) // n_cols\n\n    fig, axes = plt.subplots(n_rows, n_cols, figsize=(15, 5*n_rows))\n    axes = axes.ravel() if n_classes > 1 else [axes]\n\n    for i, (class_name, cm) in enumerate(zip(class_names, ml_cms)):\n        ax = axes[i]\n\n        # Plot heatmap\n        sns.heatmap(\n            cm, annot=True, fmt='d', cmap='Blues',\n            ax=ax, cbar=False,\n            xticklabels=['Pred Neg', 'Pred Pos'],\n            yticklabels=['True Neg', 'True Pos'],\n            annot_kws={'size': 12}\n        )\n\n        # Calculate metrics\n        tn, fp, fn, tp = cm.ravel()\n        precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n        recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n        f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0\n\n        ax.set_title(\n            f'{class_name}\\n'\n            f'P={precision:.3f}, R={recall:.3f}, F1={f1:.3f}',\n            fontsize=11, fontweight='bold'\n        )\n        ax.set_xlabel('Predicted', fontsize=10)\n        ax.set_ylabel('Actual', fontsize=10)\n\n    # Hide extra subplots\n    for i in range(n_classes, len(axes)):\n        axes[i].axis('off')\n\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=300, bbox_inches='tight')\n    plt.show()\n\n    print(f\"Per-class confusion matrices saved to {save_path}\")\n\n    # Print summary statistics\n    print(\"\\nPer-Class Confusion Matrix Summary:\")\n    print(\"-\" * 70)\n    print(f\"{'Class':<25} {'TP':>5} {'FP':>5} {'FN':>5} {'TN':>5} {'Prec':>5} {'Rec':>5} {'F1':>5}\")\n    print(\"-\" * 70)\n\n    for i, (class_name, cm) in enumerate(zip(class_names, ml_cms)):\n        tn, fp, fn, tp = cm.ravel()\n        precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n        recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n        f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0\n\n        print(f\"{class_name:<25} {tp:>5} {fp:>5} {fn:>5} {tn:>5} {precision:>5.3f} {recall:>5.3f} {f1:>5.3f}\")\n\n# Plot per-class confusion matrices\nplot_per_class_confusion_matrices(\n    opt_labels, opt_preds, config.LABELS,\n    config.PLOT_DIR / 'baseline_per_class_confusion_matrices.png'\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T09:59:41.644757Z","iopub.execute_input":"2026-01-07T09:59:41.645793Z","iopub.status.idle":"2026-01-07T09:59:43.773361Z","shell.execute_reply.started":"2026-01-07T09:59:41.645739Z","shell.execute_reply":"2026-01-07T09:59:43.772428Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Per-Class Performance Summary CSV\nfrom sklearn.metrics import accuracy_score\n\ndef create_per_class_performance_csv(labels_baseline, preds_baseline,\n                                     labels_opt, preds_opt,\n                                     class_names, thresholds_dict, save_path):\n    \"\"\"\n    Create a comprehensive CSV with per-class metrics for baseline and optimized.\n\n    Args:\n        labels_baseline: True labels (same for both)\n        preds_baseline: Predictions with default 0.5 threshold\n        labels_opt: True labels (same as baseline)\n        preds_opt: Predictions with optimal thresholds\n        class_names: List of class names\n        thresholds_dict: Optimal thresholds for each class\n        save_path: Path to save CSV\n    \"\"\"\n    results = []\n\n    for i, class_name in enumerate(class_names):\n        # Baseline metrics\n        true = labels_baseline[:, i]\n        pred_baseline = preds_baseline[:, i]\n\n        acc_baseline = accuracy_score(true, pred_baseline)\n        f1_baseline = f1_score(true, pred_baseline, zero_division=0)\n        prec_baseline = precision_score(true, pred_baseline, zero_division=0)\n        rec_baseline = recall_score(true, pred_baseline, zero_division=0)\n\n        # Optimized metrics\n        pred_opt = preds_opt[:, i]\n\n        acc_opt = accuracy_score(true, pred_opt)\n        f1_opt = f1_score(true, pred_opt, zero_division=0)\n        prec_opt = precision_score(true, pred_opt, zero_division=0)\n        rec_opt = recall_score(true, pred_opt, zero_division=0)\n\n        results.append({\n            'Class': class_name,\n            'Optimal_Threshold': thresholds_dict[class_name],\n            # Baseline\n            'Baseline_Acc': acc_baseline,\n            'Baseline_Precision': prec_baseline,\n            'Baseline_Recall': rec_baseline,\n            'Baseline_F1': f1_baseline,\n            # Optimized\n            'Opt_Acc': acc_opt,\n            'Opt_Precision': prec_opt,\n            'Opt_Recall': rec_opt,\n            'Opt_F1': f1_opt,\n            # Improvements\n            'Acc_Improvement': acc_opt - acc_baseline,\n            'Prec_Improvement': prec_opt - prec_baseline,\n            'Rec_Improvement': rec_opt - rec_baseline,\n            'F1_Improvement': f1_opt - f1_baseline\n        })\n\n    df = pd.DataFrame(results)\n    df.to_csv(save_path, index=False)\n\n    print(f\"Per-class performance CSV saved to {save_path}\")\n    return df\n\n# Create per-class performance summary\nper_class_df = create_per_class_performance_csv(\n    test_labels, test_preds_baseline,\n    opt_labels, opt_preds,\n    config.LABELS, optimal_thresholds,\n    config.METRICS_DIR / 'per_class_performance_summary.csv'\n)\n\n# Display summary\nprint(\"\\n\" + \"=\" * 100)\nprint(\"PER-CLASS PERFORMANCE SUMMARY\")\nprint(\"=\" * 100)\nprint(per_class_df.to_string(index=False))\nprint(\"=\" * 100)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T10:00:09.170008Z","iopub.execute_input":"2026-01-07T10:00:09.170494Z","iopub.status.idle":"2026-01-07T10:00:09.266396Z","shell.execute_reply.started":"2026-01-07T10:00:09.170460Z","shell.execute_reply":"2026-01-07T10:00:09.265408Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Visualizing Predictions on Test Samples\n\nLet's visualize the model's predictions on random test samples to **qualitatively assess performance**. This helps us understand:\n\n1. **Which diseases are easily confused?** (e.g., rust vs. complex)\n2. **Does the model miss obvious symptoms?** (false negatives)\n3. **Does the model over-predict certain diseases?** (false positives)\n4. **Are predictions consistent with visual symptoms?**","metadata":{}},{"cell_type":"code","source":"# Visualize Test Predictions\ndef visualize_predictions(test_df, labels, preds, class_names,\n                          thresholds_dict, num_samples=15):\n    \"\"\"\n    Visualize model predictions on random test samples.\n\n    Args:\n        test_df: Test DataFrame with image paths and labels\n        labels: True labels (multi-hot encoded)\n        preds: Predicted labels (multi-hot encoded)\n        class_names: List of class names\n        thresholds_dict: Optimal thresholds for each class\n        num_samples: Number of samples to visualize\n    \"\"\"\n    # Sample random indices\n    indices = np.random.choice(len(test_df), size=min(num_samples, len(test_df)), replace=False)\n\n    # Setup plot\n    n_cols = 5\n    n_rows = (num_samples + n_cols - 1) // n_cols\n    fig, axes = plt.subplots(n_rows, n_cols, figsize=(20, 4*n_rows))\n    axes = axes.ravel() if num_samples > 1 else [axes]\n\n    for plot_idx, sample_idx in enumerate(indices):\n        ax = axes[plot_idx]\n\n        # Load image\n        img_path = test_df.iloc[sample_idx]['image']\n        img = Image.open(img_path)\n\n        # Get labels\n        true_labels = [class_names[i] for i, val in enumerate(labels[sample_idx]) if val == 1]\n        pred_labels = [class_names[i] for i, val in enumerate(preds[sample_idx]) if val == 1]\n\n        if not pred_labels:\n            pred_labels = ['None']\n\n        # Check if prediction is correct\n        is_correct = set(true_labels) == set(pred_labels)\n        color = 'green' if is_correct else 'red'\n\n        # Plot\n        ax.imshow(img)\n        ax.axis('off')\n        ax.set_title(\n            f\"True: {', '.join(true_labels)}\\n\"\n            f\"Pred: {', '.join(pred_labels)}\\n\"\n            f\"{'✓ Correct' if is_correct else '✗ Wrong'}\",\n            fontsize=9, color=color, fontweight='bold'\n        )\n\n    # Hide extra subplots\n    for i in range(num_samples, len(axes)):\n        axes[i].axis('off')\n\n    plt.tight_layout()\n    save_path = config.PLOT_DIR / 'baseline_test_predictions.png'\n    plt.savefig(save_path, dpi=300, bbox_inches='tight')\n    plt.show()\n    print(f\"Predictions visualization saved to {save_path}\")\n\n# Visualize predictions\nvisualize_predictions(\n    test_df, opt_labels, opt_preds, config.LABELS,\n    optimal_thresholds, num_samples=15\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T10:01:21.216966Z","iopub.execute_input":"2026-01-07T10:01:21.218107Z","iopub.status.idle":"2026-01-07T10:01:25.402361Z","shell.execute_reply.started":"2026-01-07T10:01:21.218068Z","shell.execute_reply":"2026-01-07T10:01:25.401289Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate Final Summary Report\ndef generate_summary_report(config, history, baseline_metrics, optimized_metrics,\n                            optimal_thresholds, checkpoint, save_path):\n    \"\"\"\n    Generate a comprehensive summary report in markdown format.\n    \n    FIXED: Added checkpoint parameter to avoid NameError\n    \"\"\"\n    report = f\"\"\"\n# Custom CNN Baseline Model - Training Report\n\n## Training Configuration\n\n- **Model Architecture**: CustomPlantCNN (4 ResBlocks, 3.5M parameters)\n- **Dataset**: Plant Pathology 2021\n- **Training Samples**: {len(train_df)}\n- **Validation Samples**: {len(val_df)}\n- **Test Samples**: {len(test_df)}\n- **Batch Size**: {config.BATCH_SIZE}\n- **Epochs**: {config.NUM_EPOCHS}\n- **Initial Learning Rate**: {config.LEARNING_RATE}\n- **Optimizer**: AdamW (weight_decay={config.WEIGHT_DECAY})\n- **Scheduler**: ReduceLROnPlateau (patience={config.PATIENCE}, factor=0.1)\n- **Device**: {config.DEVICE}\n\n## Training Results\n\n### Best Model\n- **Best Epoch**: {checkpoint['epoch']}\n- **Validation Loss**: {checkpoint['val_loss']:.4f}\n- **Validation Accuracy**: {checkpoint['val_acc']:.4f}\n- **Validation F1**: {checkpoint.get('val_f1', 'N/A')}\n\n### Final Training Metrics\n- **Training Loss**: {history['train_loss'][-1]:.4f}\n- **Training Accuracy**: {history['train_acc'][-1]:.4f}\n- **Training F1**: {history['train_f1'][-1]:.4f}\n- **Validation Loss**: {history['val_loss'][-1]:.4f}\n- **Validation Accuracy**: {history['val_acc'][-1]:.4f}\n- **Validation F1**: {history['val_f1'][-1]:.4f}\n\n## Test Set Performance\n\n### Baseline (Default Threshold = 0.5)\n- **Exact Match Accuracy**: {baseline_metrics['exact_match_accuracy']:.4f}\n- **Macro F1**: {baseline_metrics['macro_f1']:.4f}\n- **Micro F1**: {baseline_metrics['micro_f1']:.4f}\n- **Macro Precision**: {baseline_metrics['macro_precision']:.4f}\n- **Macro Recall**: {baseline_metrics['macro_recall']:.4f}\n\n### After Threshold Optimization\nOptimal Thresholds:\n\"\"\"\n\n    for name, thresh in optimal_thresholds.items():\n        report += f\"- **{name}**: {thresh:.2f}\\n\"\n\n    report += f\"\"\"\n\nPerformance:\n- **Exact Match Accuracy**: {optimized_metrics['exact_match_accuracy']:.4f} (Δ {optimized_metrics['improvement_over_baseline']['accuracy']:+.4f})\n- **Macro F1**: {optimized_metrics['macro_f1']:.4f} (Δ {optimized_metrics['improvement_over_baseline']['macro_f1']:+.4f})\n- **Micro F1**: {optimized_metrics['micro_f1']:.4f} (Δ {optimized_metrics['improvement_over_baseline']['micro_f1']:+.4f})\n\n## Key Insights\n\n1. **Threshold Optimization Impact**: Class-specific thresholds improve macro-F1 by {(optimized_metrics['improvement_over_baseline']['macro_f1']/baseline_metrics['macro_f1'])*100:.2f}%\n\n2. **Model Strengths**: Lightweight (3.5M params), fast training, competitive performance\n\n3. **Class Imbalance Handling**: Inverse frequency weighting helps with rare classes\n\n4. **Confusion Patterns**: Refer to confusion matrix plots for detailed analysis\n\n## Output Files\n\nAll outputs have been saved to: `{config.OUTPUT_DIR.absolute()}`\n\n### Checkpoints\n- `best_model.pth`: Best model checkpoint\n\n### Plots\n- `class_distribution.png`: Class distribution with weights\n- `baseline_training_curves.png`: Loss/accuracy/F1 curves\n- `baseline_confusion_matrix.png`: Combined confusion matrix\n- `baseline_per_class_confusion_matrices.png`: Per-class confusion matrices\n- `baseline_test_predictions.png`: Visual predictions on test samples\n\n### Metrics\n- `training_history.json`: Training metrics per epoch\n- `baseline_test_metrics.json`: Test metrics (default threshold)\n- `optimized_test_metrics.json`: Test metrics (optimized thresholds)\n- `optimal_thresholds.json`: Optimal thresholds for each class\n\n---\n*Report generated by Custom CNN Baseline Model*\n\"\"\"\n\n    # Save report\n    with open(save_path, 'w') as f:\n        f.write(report)\n\n    print(report)\n    print(f\"\\nReport saved to {save_path}\")\n\n# Generate report - FIXED: Pass checkpoint as parameter\ngenerate_summary_report(\n    config, history, baseline_metrics, optimized_metrics,\n    optimal_thresholds, checkpoint,  # ← Added checkpoint parameter\n    config.OUTPUT_DIR / 'TRAINING_REPORT.md'\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T10:00:45.552572Z","iopub.execute_input":"2026-01-07T10:00:45.553125Z","iopub.status.idle":"2026-01-07T10:00:45.564112Z","shell.execute_reply.started":"2026-01-07T10:00:45.553071Z","shell.execute_reply":"2026-01-07T10:00:45.563113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Final Summary\nprint(\"\\n\" + \"=\"*80)\nprint(\"BASELINE MODEL TRAINING COMPLETE!\")\nprint(\"=\"*80)\nprint(f\"\\nAll outputs saved to: {config.OUTPUT_DIR.absolute()}\")\nprint(f\"\\nSummary of saved files:\")\nprint(f\"  - Checkpoints: {list(config.CHECKPOINT_DIR.glob('*.pth'))}\")\nprint(f\"  - Plots: {len(list(config.PLOT_DIR.glob('*.png')))} PNG files\")\nprint(f\"  - Metrics: {len(list(config.METRICS_DIR.glob('*.json')))} JSON files\")\nprint(f\"  - Report: {config.OUTPUT_DIR / 'TRAINING_REPORT.md'}\")\nprint(\"\\nNext steps:\")\nprint(\"  1. Review the training curves and confusion matrices\")\nprint(\"  2. Analyze per-class metrics to identify weaknesses\")\nprint(\"  3. Consider improvements: transfer learning, advanced augmentations, etc.\")\nprint(\"=\"*80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-07T10:00:48.417707Z","iopub.execute_input":"2026-01-07T10:00:48.418389Z","iopub.status.idle":"2026-01-07T10:00:48.426410Z","shell.execute_reply.started":"2026-01-07T10:00:48.418352Z","shell.execute_reply":"2026-01-07T10:00:48.425388Z"}},"outputs":[],"execution_count":null}]}