{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":25563,"databundleVersionId":2094376,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":121906917,"sourceType":"kernelVersion"}],"dockerImageVersionId":30097,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# --- 1. IMPORTS AND SETUP ---\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom pathlib import Path\nfrom PIL import Image\nfrom tqdm import tqdm\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.metrics import (\n    classification_report, \n    confusion_matrix, \n    accuracy_score, \n    multilabel_confusion_matrix,\n    f1_score,\n    precision_score,\n    recall_score\n)\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport json\nimport warnings\n\nwarnings.filterwarnings('ignore')\n\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-04T13:53:02.960872Z","iopub.execute_input":"2026-01-04T13:53:02.961170Z","iopub.status.idle":"2026-01-04T13:53:02.968960Z","shell.execute_reply.started":"2026-01-04T13:53:02.961145Z","shell.execute_reply":"2026-01-04T13:53:02.968172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 2. CONFIGURATION ---\nclass config:\n    # FIXED: Use local data paths instead of Kaggle paths\n    DATA_DIR = Path('/kaggle/input/multi-label-classification-plant-pathology/data')\n    CSV_PATH = Path('../input/plant-pathology-2021-fgvc8/train.csv')\n    ROOT_DIR = Path('./data')  # For resized images\n    \n    # Model parameters\n    INPUT_HEIGHT = 224\n    INPUT_WIDTH = 224\n    BATCH_SIZE = 32\n    \n    # DenseNet architecture\n    GROWTH_RATE = 32\n    BLOCK_CONFIG = (6, 12, 16)  # Number of layers in each dense block\n    \n    # Training parameters\n    EPOCHS = 20\n    LEARNING_RATE = 1e-3\n    WEIGHT_DECAY = 1e-4\n    LABEL_SMOOTHING = 0.1\n    \n    # Device\n    DEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n    # ImageNet normalization\n    IMAGENET_MEAN = [0.485, 0.456, 0.406]\n    IMAGENET_STD = [0.229, 0.224, 0.225]\n    \n    # Labels (will be updated from data)\n    NUM_CLASSES = 0\n    LABELS = []\n    \n    # Multi-label mode\n    MULTILABEL = True  # Set to False for single-label mode\n\nprint(\"Configuration loaded successfully\")\nprint(f\"Data directory: {config.DATA_DIR}\")\nprint(f\"Device: {config.DEVICE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:02.970346Z","iopub.execute_input":"2026-01-04T13:53:02.970696Z","iopub.status.idle":"2026-01-04T13:53:02.988816Z","shell.execute_reply.started":"2026-01-04T13:53:02.970660Z","shell.execute_reply":"2026-01-04T13:53:02.988000Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 3. DATA SPLITTING ---\ndef split_df(csv_dir):\n    \"\"\"\n    Load CSV and split into train/validation/test sets (70/15/15)\n    \"\"\"\n    df = pd.read_csv(csv_dir)\n    \n    # Update image paths to use local data directory\n    df['image'] = df['image'].apply(lambda x: str(config.DATA_DIR / x))\n    \n    # Split data\n    train_df, dummy_df = train_test_split(\n        df, \n        train_size=0.7, \n        shuffle=True, \n        random_state=42\n    )\n    \n    valid_df, test_df = train_test_split(\n        dummy_df, \n        train_size=0.5, \n        shuffle=True, \n        random_state=42\n    )\n    \n    return train_df, valid_df, test_df\n\nprint(\"Loading and splitting data...\")\ntrain_df, valid_df, test_df = split_df(config.CSV_PATH)\n\nprint(f\"\\nDataset sizes:\")\nprint(f\"  Train: {len(train_df)} images\")\nprint(f\"  Valid: {len(valid_df)} images\")\nprint(f\"  Test:  {len(test_df)} images\")\n\n# Analyze labels\nif config.MULTILABEL:\n    # Multi-label: split label strings and get unique classes\n    all_labels = train_df['labels'].str.split(expand=True).stack().reset_index(drop=True)\n    label_counts = all_labels.value_counts()\n    config.LABELS = ['complex', 'frog_eye_leaf_spot', 'healthy', 'powdery_mildew', 'rust', 'scab']\n    config.NUM_CLASSES = len(config.LABELS)\n    \n    print(f\"\\nClass distribution (train set):\")\n    for label in config.LABELS:\n        count = label_counts.get(label, 0)\n        print(f\"  {label}: {count}\")\nelse:\n    # Single-label: use LabelEncoder\n    le = LabelEncoder()\n    train_df['labels_n'] = le.fit_transform(train_df['labels'].values)\n    config.LABELS = list(le.classes_)\n    config.NUM_CLASSES = len(le.classes_)\n    print(f\"\\nClasses: {config.LABELS}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:02.990125Z","iopub.execute_input":"2026-01-04T13:53:02.990362Z","iopub.status.idle":"2026-01-04T13:53:03.128278Z","shell.execute_reply.started":"2026-01-04T13:53:02.990327Z","shell.execute_reply":"2026-01-04T13:53:03.127556Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Class Imbalance Analysis\n\nThe dataset has significant class imbalance:\n- **scab**: ~4000 samples (majority)\n- **powdery_mildew**: ~900 samples (minority)\n\nThis will be handled using **class weighting** in the loss function.","metadata":{}},{"cell_type":"code","source":"# --- 4. CLASS WEIGHTING FOR IMBALANCE ---\nif config.MULTILABEL:\n    # Get class counts from training data\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    print(\"Class counts:\")\n    print(class_counts)\n    \n    # Compute inverse class frequency weights\n    class_weights = torch.reciprocal(torch.tensor(class_counts.values).float())\n    class_weights /= torch.max(class_weights)  # Normalize\n    \n    print(\"\\nClass weights (inverse frequency):\")\n    for label, weight in zip(config.LABELS, class_weights):\n        print(f\"  {label}: {weight:.4f}\")\nelse:\n    class_weights = None\n\n# Plot class distribution\nif config.MULTILABEL:\n    plt.figure(figsize=(10, 5))\n    class_counts.plot(kind='bar')\n    plt.title('Class Distribution (Train Set)')\n    plt.xlabel('Class')\n    plt.ylabel('Count')\n    plt.xticks(rotation=45, ha='right')\n    plt.grid(axis='y', alpha=0.3)\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:03.129286Z","iopub.execute_input":"2026-01-04T13:53:03.129503Z","iopub.status.idle":"2026-01-04T13:53:03.358429Z","shell.execute_reply.started":"2026-01-04T13:53:03.129480Z","shell.execute_reply":"2026-01-04T13:53:03.357135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 5. LABEL ENCODING/DECODING ---\ndef encode_label(labels_str, class_list):\n    \"\"\"\n    Convert space-separated label string to one-hot encoded tensor.\n    Example: 'scab frog_eye_leaf_spot' -> [0, 1, 0, 0, 0, 1]\n    \"\"\"\n    if not config.MULTILABEL:\n        # Single-label: use integer encoding\n        return torch.tensor(class_list.index(labels_str), dtype=torch.long)\n    \n    # Multi-label: one-hot encoding\n    labels = labels_str.split(' ')\n    target = torch.zeros(len(class_list))\n    for label in labels:\n        if label in class_list:\n            idx = class_list.index(label)\n            target[idx] = 1\n    return target\n\ndef decode_label(encoded_label, class_list):\n    \"\"\"\n    Convert encoded tensor back to list of label strings.\n    \"\"\"\n    if not config.MULTILABEL:\n        # Single-label\n        return class_list[encoded_label]\n    \n    # Multi-label\n    if isinstance(encoded_label, torch.Tensor):\n        encoded_label = encoded_label.cpu().numpy()\n    return [class_list[i] for i, val in enumerate(encoded_label) if val == 1]\n\n# Test encoding/decoding\nif config.MULTILABEL:\n    test_labels = 'scab healthy'\n    encoded = encode_label(test_labels, config.LABELS)\n    decoded = decode_label(encoded, config.LABELS)\n    print(f\"Test encoding: '{test_labels}' -> {encoded}\")\n    print(f\"Test decoding: {encoded} -> {decoded}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:03.360136Z","iopub.execute_input":"2026-01-04T13:53:03.360604Z","iopub.status.idle":"2026-01-04T13:53:03.369480Z","shell.execute_reply.started":"2026-01-04T13:53:03.360578Z","shell.execute_reply":"2026-01-04T13:53:03.368650Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 6. DATASET CLASS ---\nclass PlantDataset(Dataset):\n    \"\"\"\n    Custom PyTorch Dataset for plant pathology images.\n    \"\"\"\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        image_path = row['image']\n        labels_str = row['labels']\n        \n        # Load image\n        try:\n            image = Image.open(image_path).convert(\"RGB\")\n        except Exception as e:\n            print(f\"Error loading image {image_path}: {e}\")\n            # Return a blank image if loading fails\n            image = Image.new('RGB', (config.INPUT_WIDTH, config.INPUT_HEIGHT))\n        \n        # Encode labels\n        label = encode_label(labels_str, config.LABELS)\n        \n        # Apply transforms\n        if self.transform:\n            image = self.transform(image)\n        \n        return image, label\n\nprint(\"Dataset class defined successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:03.370951Z","iopub.execute_input":"2026-01-04T13:53:03.371179Z","iopub.status.idle":"2026-01-04T13:53:03.383004Z","shell.execute_reply.started":"2026-01-04T13:53:03.371156Z","shell.execute_reply":"2026-01-04T13:53:03.382325Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 7. DATA AUGMENTATION ---\n# Training transforms with strong augmentation\ntrain_transforms = transforms.Compose([\n    transforms.RandomResizedCrop(\n        size=(config.INPUT_HEIGHT, config.INPUT_WIDTH), \n        scale=(0.8, 1.0),\n        ratio=(0.9, 1.1)\n    ),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(degrees=30),\n    transforms.ColorJitter(\n        brightness=0.2, \n        contrast=0.2, \n        saturation=0.2, \n        hue=0.1\n    ),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=config.IMAGENET_MEAN, \n        std=config.IMAGENET_STD\n    )\n])\n\n# Validation/test transforms (no augmentation)\ntest_transforms = transforms.Compose([\n    transforms.Resize(\n        size=(config.INPUT_HEIGHT, config.INPUT_WIDTH)\n    ),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=config.IMAGENET_MEAN, \n        std=config.IMAGENET_STD\n    )\n])\n\nprint(\"Data transforms defined:\")\nprint(f\"  Train transforms: {len(train_transforms.transforms)} augmentations\")\nprint(f\"  Test transforms: Resize + Normalize only\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:03.384097Z","iopub.execute_input":"2026-01-04T13:53:03.384318Z","iopub.status.idle":"2026-01-04T13:53:03.399812Z","shell.execute_reply.started":"2026-01-04T13:53:03.384291Z","shell.execute_reply":"2026-01-04T13:53:03.399119Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Augmentation Explained\n\n**Why augment training data?**\n\n1. **Prevents overfitting**: Model sees different variations of same image\n2. **Improves generalization**: Learns invariant features (disease recognizable from any angle)\n3. **Effective for small datasets**: Artificially increases training data size\n\n**Augmentations used:**\n- **RandomResizedCrop**: Simulates different distances/zoom levels\n- **RandomHorizontalFlip/VerticalFlip**: Leaves can be photographed from any orientation\n- **RandomRotation**: Adds rotational invariance\n- **ColorJitter**: Simulates different lighting conditions\n\n**Important**: Validation/test data only gets resized and normalized - no augmentation!","metadata":{}},{"cell_type":"code","source":"# --- 8. CREATE DATALOADERS ---\ntrain_dataset = PlantDataset(train_df, transform=train_transforms)\nvalid_dataset = PlantDataset(valid_df, transform=test_transforms)\ntest_dataset = PlantDataset(test_df, transform=test_transforms)\n\ntrain_loader = DataLoader(\n    train_dataset, \n    batch_size=config.BATCH_SIZE, \n    shuffle=True, \n    num_workers=0,  # Set to 0 for Windows compatibility\n    pin_memory=True if torch.cuda.is_available() else False\n)\n\nvalid_loader = DataLoader(\n    valid_dataset, \n    batch_size=config.BATCH_SIZE, \n    shuffle=False, \n    num_workers=0,\n    pin_memory=True if torch.cuda.is_available() else False\n)\n\ntest_loader = DataLoader(\n    test_dataset, \n    batch_size=config.BATCH_SIZE, \n    shuffle=False, \n    num_workers=0,\n    pin_memory=True if torch.cuda.is_available() else False\n)\n\nprint(f\"Dataloaders created successfully\")\nprint(f\"Train batches: {len(train_loader)}\")\nprint(f\"Valid batches: {len(valid_loader)}\")\nprint(f\"Test batches: {len(test_loader)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:03.400810Z","iopub.execute_input":"2026-01-04T13:53:03.401118Z","iopub.status.idle":"2026-01-04T13:53:03.417152Z","shell.execute_reply.started":"2026-01-04T13:53:03.401087Z","shell.execute_reply":"2026-01-04T13:53:03.416437Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 9. DENSENET MODEL DEFINITION ---\nclass DenseLayer(nn.Module):\n    \"\"\"\n    Single layer in a DenseBlock.\n    Structure: BN -> ReLU -> Conv(1x1) -> BN -> ReLU -> Conv(3x3)\n    \"\"\"\n    def __init__(self, in_channels, growth_rate):\n        super(DenseLayer, self).__init__()\n        self.bn1 = nn.BatchNorm2d(in_channels)\n        self.relu = nn.ReLU(inplace=True)\n        # 1x1 convolution reduces dimensions (bottleneck)\n        self.conv1 = nn.Conv2d(in_channels, 4 * growth_rate, kernel_size=1, bias=False)\n        self.bn2 = nn.BatchNorm2d(4 * growth_rate)\n        # 3x3 convolution produces growth_rate feature maps\n        self.conv2 = nn.Conv2d(4 * growth_rate, growth_rate, kernel_size=3, padding=1, bias=False)\n\n    def forward(self, x):\n        out = self.conv1(self.relu(self.bn1(x)))\n        out = self.conv2(self.relu(self.bn2(out)))\n        # Concatenate with input (dense connection)\n        return torch.cat([x, out], 1)\n\n\nclass DenseBlock(nn.Module):\n    \"\"\"\n    DenseBlock: Stack of DenseLayers where each layer receives all previous layers' outputs.\n    \"\"\"\n    def __init__(self, num_layers, in_channels, growth_rate):\n        super(DenseBlock, self).__init__()\n        layers = []\n        for i in range(num_layers):\n            # Each layer receives input from all previous layers\n            layers.append(DenseLayer(in_channels + i * growth_rate, growth_rate))\n        self.layer = nn.Sequential(*layers)\n\n    def forward(self, x):\n        return self.layer(x)\n\n\nclass TransitionLayer(nn.Module):\n    \"\"\"\n    TransitionLayer: Reduces spatial dimensions and number of channels between DenseBlocks.\n    Structure: BN -> ReLU -> Conv(1x1) -> AvgPool(2x2)\n    \"\"\"\n    def __init__(self, in_channels, out_channels):\n        super(TransitionLayer, self).__init__()\n        self.bn = nn.BatchNorm2d(in_channels)\n        self.relu = nn.ReLU(inplace=True)\n        # 1x1 convolution to reduce channels\n        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)\n        # Average pooling to reduce spatial dimensions\n        self.avg_pool = nn.AvgPool2d(kernel_size=2, stride=2)\n\n    def forward(self, x):\n        out = self.conv(self.relu(self.bn(x)))\n        out = self.avg_pool(out)\n        return out\n\n\nclass ResearcherDenseNet(nn.Module):\n    \"\"\"\n    DenseNet architecture for plant disease classification.\n    \"\"\"\n    def __init__(self, num_classes, growth_rate=32, block_config=(6, 12, 16)):\n        super(ResearcherDenseNet, self).__init__()\n        \n        # Initial convolution (before dense blocks)\n        self.features = nn.Sequential(\n            nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n        )\n        \n        num_features = 64\n        \n        # DenseBlocks and TransitionLayers\n        for i, num_layers in enumerate(block_config):\n            # Add DenseBlock\n            block = DenseBlock(num_layers, num_features, growth_rate)\n            self.features.add_module(f'denseblock{i+1}', block)\n            num_features = num_features + num_layers * growth_rate\n            \n            # Add TransitionLayer (except after last block)\n            if i != len(block_config) - 1:\n                out_features = num_features // 2\n                trans = TransitionLayer(num_features, out_features)\n                self.features.add_module(f'transition{i+1}', trans)\n                num_features = out_features\n\n        # Final batch norm and global pooling\n        self.final_bn = nn.BatchNorm2d(num_features)\n        \n        # Classifier\n        self.classifier = nn.Linear(num_features, num_classes)\n\n    def forward(self, x):\n        features = self.features(x)\n        out = torch.relu(self.final_bn(features))\n        out = F.adaptive_avg_pool2d(out, (1, 1))\n        out = torch.flatten(out, 1)\n        out = self.classifier(out)\n        return out\n\n\n# Initialize model\nmodel = ResearcherDenseNet(\n    num_classes=config.NUM_CLASSES,\n    growth_rate=config.GROWTH_RATE,\n    block_config=config.BLOCK_CONFIG\n).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\"\\nModel initialized successfully\")\nprint(f\"Total parameters: {total_params:,}\")\nprint(f\"Trainable parameters: {trainable_params:,}\")\nprint(f\"\\nModel architecture:\")\nprint(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:03.418314Z","iopub.execute_input":"2026-01-04T13:53:03.418639Z","iopub.status.idle":"2026-01-04T13:53:03.495575Z","shell.execute_reply.started":"2026-01-04T13:53:03.418607Z","shell.execute_reply":"2026-01-04T13:53:03.494525Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loss Function and Optimizer\n\n### BCEWithLogitsLoss (Multi-label)\n\nFor multi-label classification, we use **Binary Cross-Entropy with Logits**:\n\n- Combines sigmoid activation + BCE loss in one function (numerically stable)\n- Each class treated as independent binary classification\n- Outputs probabilities for each disease\n\n### CrossEntropyLoss (Single-label)\n\nFor single-label classification:\n\n- Standard softmax cross-entropy\n- Mutually exclusive classes\n\n### Class Weighting\n\nMinority classes get higher weights in the loss:\n- `pos_weight` parameter in BCEWithLogitsLoss\n- Balances contribution of each class to total loss\n\n### AdamW Optimizer\n\n**AdamW = Adam + Decoupled Weight Decay**\n\n- Adaptive learning rates (per-parameter)\n- Momentum for faster convergence\n- Proper L2 regularization (weight decay)\n- State-of-the-art for computer vision","metadata":{}},{"cell_type":"code","source":"# --- 10. LOSS FUNCTION AND OPTIMIZER ---\nif config.MULTILABEL:\n    # Multi-label: BCEWithLogitsLoss with class weighting\n    criterion = nn.BCEWithLogitsLoss(\n        pos_weight=class_weights.to(config.DEVICE) if class_weights is not None else None\n    )\nelse:\n    # Single-label: CrossEntropyLoss with label smoothing\n    criterion = nn.CrossEntropyLoss(\n        label_smoothing=config.LABEL_SMOOTHING\n    )\n\n# AdamW optimizer (Adam + 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',  # Reduce LR when validation loss plateaus\n    factor=0.1,  # Multiply LR by 0.1\n    patience=3,  # Wait 3 epochs before reducing LR\n    verbose=True\n)\n\nprint(f\"Loss function: {criterion.__class__.__name__}\")\nprint(f\"Optimizer: AdamW (lr={config.LEARNING_RATE}, weight_decay={config.WEIGHT_DECAY})\")\nprint(f\"Scheduler: ReduceLROnPlateau (patience=3, factor=0.1)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:03.496577Z","iopub.execute_input":"2026-01-04T13:53:03.496795Z","iopub.status.idle":"2026-01-04T13:53:03.506950Z","shell.execute_reply.started":"2026-01-04T13:53:03.496774Z","shell.execute_reply":"2026-01-04T13:53:03.505472Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 11. TRAINING FUNCTIONS ---\ndef train_one_epoch(model, loader, criterion, optimizer, device):\n    \"\"\"\n    Train for one epoch.\n    Returns: (avg_loss, avg_accuracy, avg_f1_macro)\n    \"\"\"\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    # Collect all predictions and labels for F1 calculation\n    all_preds = []\n    all_labels = []\n\n    loop = tqdm(loader, desc=\"Training\", leave=False)\n    for images, labels in loop:\n        images, labels = images.to(device), labels.to(device)\n\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        # Statistics\n        running_loss += loss.item() * images.size(0)\n\n        if config.MULTILABEL:\n            # Multi-label accuracy (threshold 0.5)\n            preds = (torch.sigmoid(outputs) > 0.5).int()\n            correct += (preds == labels.int()).sum().item()\n            total += labels.numel()\n            \n            # Store for F1 calculation\n            all_preds.append(preds.cpu())\n            all_labels.append(labels.cpu())\n        else:\n            # Single-label accuracy\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n\n        loop.set_postfix({'loss': f'{loss.item():.4f}'})\n\n    avg_loss = running_loss / len(loader.dataset)\n    avg_acc = correct / total\n    \n    # Calculate macro F1 score for multi-label\n    if config.MULTILABEL and len(all_preds) > 0:\n        all_preds = torch.cat(all_preds).numpy()\n        all_labels = torch.cat(all_labels).numpy()\n        avg_f1 = f1_score(all_labels, all_preds, average='macro', zero_division=0)\n    else:\n        avg_f1 = 0.0\n    \n    return avg_loss, avg_acc, avg_f1\n\n\ndef validate(model, loader, criterion, device):\n    \"\"\"\n    Validate the model.\n    Returns: (avg_loss, avg_accuracy, avg_f1_macro)\n    \"\"\"\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    # Collect all predictions and labels for F1 calculation\n    all_preds = []\n    all_labels = []\n\n    with torch.no_grad():\n        loop = tqdm(loader, desc=\"Validating\", leave=False)\n        for images, labels in loop:\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item() * images.size(0)\n\n            if config.MULTILABEL:\n                preds = (torch.sigmoid(outputs) > 0.5).int()\n                correct += (preds == labels.int()).sum().item()\n                total += labels.numel()\n                \n                # Store for F1 calculation\n                all_preds.append(preds.cpu())\n                all_labels.append(labels.cpu())\n            else:\n                _, predicted = outputs.max(1)\n                total += labels.size(0)\n                correct += predicted.eq(labels).sum().item()\n\n    avg_loss = running_loss / len(loader.dataset)\n    avg_acc = correct / total\n    \n    # Calculate macro F1 score for multi-label\n    if config.MULTILABEL and len(all_preds) > 0:\n        all_preds = torch.cat(all_preds).numpy()\n        all_labels = torch.cat(all_labels).numpy()\n        avg_f1 = f1_score(all_labels, all_preds, average='macro', zero_division=0)\n    else:\n        avg_f1 = 0.0\n    \n    return avg_loss, avg_acc, avg_f1\n\nprint(\"Training functions defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:03.507994Z","iopub.execute_input":"2026-01-04T13:53:03.508303Z","iopub.status.idle":"2026-01-04T13:53:03.525358Z","shell.execute_reply.started":"2026-01-04T13:53:03.508273Z","shell.execute_reply":"2026-01-04T13:53:03.524637Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 12. TRAINING LOOP ---\n# Training history tracking\nhistory = {\n    'train_loss': [],\n    'train_acc': [],\n    'train_f1': [], \n    'val_loss': [],\n    'val_acc': [],\n    'val_f1': []\n}\n\nbest_val_f1 = float('-inf')  # Changed from best_val_loss\npatience_counter = 0\nEARLY_STOPPING_PATIENCE = 5\n\nprint(f\"Starting training for {config.EPOCHS} epochs...\")\nprint(f\"=\"*60)\nprint(f\"This ensures we select the best model for multi-label classification\")\nprint(f\"=\"*60)\n\nfor epoch in range(config.EPOCHS):\n    print(f\"\\nEpoch {epoch+1}/{config.EPOCHS}\")\n    print(\"-\" * 60)\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, valid_loader, criterion, config.DEVICE\n    )\n\n    # Update scheduler\n    scheduler.step(val_loss)\n\n    # Track history\n    history['train_loss'].append(train_loss)\n    history['train_acc'].append(train_acc)\n    history['train_f1'].append(train_f1) \n    history['val_loss'].append(val_loss)\n    history['val_acc'].append(val_acc)\n    history['val_f1'].append(val_f1)\n\n    # Print epoch summary with F1 scores\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: {optimizer.param_groups[0]['lr']:.2e}\")\n\n    if val_f1 > best_val_f1:\n        best_val_f1 = val_f1\n        patience_counter = 0\n        torch.save({\n            'epoch': epoch,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'val_loss': val_loss,\n            'val_acc': val_acc,\n            'val_f1': val_f1,\n        }, 'best_model_densenet.pth')\n        print(f\"*** Saved best model (Val F1 improved: {val_f1:.4f}) ***\")\n    else:\n        patience_counter += 1\n        print(f\"Val F1 did not improve from {best_val_f1:.4f}\")\n\n    # Early stopping\n    if patience_counter >= EARLY_STOPPING_PATIENCE:\n        print(f\"\\nEarly stopping triggered after {epoch+1} epochs\")\n        break\n\nprint(f\"\\n{'='*60}\")\nprint(f\"Training completed!\")\nprint(f\"Best validation F1: {best_val_f1:.4f}\")  # UPDATED: Show F1\nprint(f\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:03.526442Z","iopub.execute_input":"2026-01-04T13:53:03.526714Z","iopub.status.idle":"2026-01-04T13:53:50.506146Z","shell.execute_reply.started":"2026-01-04T13:53:03.526692Z","shell.execute_reply":"2026-01-04T13:53:50.504246Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 13. PLOT TRAINING CURVES ---\nepochs = range(1, len(history['train_loss']) + 1)\n\nfig, axes = plt.subplots(1, 3, figsize=(20, 5))\n\n# Loss plot\naxes[0].plot(epochs, history['train_loss'], 'b-o', label='Train Loss', markersize=4)\naxes[0].plot(epochs, history['val_loss'], 'r-s', label='Val Loss', markersize=4)\naxes[0].set_title('Training and Validation Loss', fontsize=14, fontweight='bold')\naxes[0].set_xlabel('Epoch', fontsize=12)\naxes[0].set_ylabel('Loss', fontsize=12)\naxes[0].legend(fontsize=11)\naxes[0].grid(True, alpha=0.3)\n\n# Accuracy plot\naxes[1].plot(epochs, history['train_acc'], 'b-o', label='Train Acc', markersize=4)\naxes[1].plot(epochs, history['val_acc'], 'r-s', label='Val Acc', markersize=4)\naxes[1].set_title('Training and Validation Accuracy', fontsize=14, fontweight='bold')\naxes[1].set_xlabel('Epoch', fontsize=12)\naxes[1].set_ylabel('Accuracy', fontsize=12)\naxes[1].legend(fontsize=11)\naxes[1].grid(True, alpha=0.3)\n\naxes[2].plot(epochs, history['train_f1'], 'b-o', label='Train F1', markersize=4)\naxes[2].plot(epochs, history['val_f1'], 'r-s', label='Val F1', markersize=4)\naxes[2].set_title('Training and Validation F1 Score (Macro)', fontsize=14, fontweight='bold')\naxes[2].set_xlabel('Epoch', fontsize=12)\naxes[2].set_ylabel('F1 Score', fontsize=12)\naxes[2].legend(fontsize=11)\naxes[2].grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.savefig('training_curves.png', dpi=150, bbox_inches='tight')\nplt.show()\n\nprint(\"Training curves saved to training_curves.png\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ========================================\n# PHASE 2: MODEL EVALUATION\n# ========================================\n\n## Baseline Evaluation (Default Threshold)\n\nFirst, evaluate the model with **default 0.5 threshold** for all classes to establish baseline performance.","metadata":{}},{"cell_type":"code","source":"# --- 14. LOAD BEST MODEL ---\ncheckpoint = torch.load('best_model_densenet.pth', map_location=config.DEVICE)\nmodel.load_state_dict(checkpoint['model_state_dict'])\nprint(f\"Loaded best model from epoch {checkpoint['epoch']+1}\")\nprint(f\"Validation loss: {checkpoint['val_loss']:.4f}\")\nprint(f\"Validation accuracy: {checkpoint['val_acc']:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:50.508287Z","iopub.status.idle":"2026-01-04T13:53:50.508823Z","shell.execute_reply":"2026-01-04T13:53:50.508568Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 15. BASELINE TEST EVALUATION (DEFAULT 0.5 THRESHOLD) ---\ndef evaluate_baseline(model, test_loader, device, threshold=0.5):\n    \"\"\"\n    Evaluate model with default threshold.\n    Returns predictions, targets, and probabilities.\n    \"\"\"\n    model.eval()\n    all_targets = []\n    all_preds = []\n    all_probs = []\n\n    with torch.no_grad():\n        for images, labels in tqdm(test_loader, desc=\"Evaluating baseline\"):\n            images = images.to(device)\n\n            outputs = model(images)\n            probs = torch.sigmoid(outputs)  # Convert logits to probabilities\n            preds = (probs > threshold).int()  # Apply threshold\n\n            all_targets.append(labels.cpu().numpy())\n            all_preds.append(preds.cpu().numpy())\n            all_probs.append(probs.cpu().numpy())\n\n    all_targets = np.vstack(all_targets)\n    all_preds = np.vstack(all_preds)\n    all_probs = np.vstack(all_probs)\n\n    return all_targets, all_preds, all_probs\n\n# Evaluate with default threshold\ny_true, y_pred, y_probs = evaluate_baseline(model, test_loader, config.DEVICE, threshold=0.5)\n\n# Calculate metrics\nexact_match_acc = accuracy_score(y_true, y_pred)\nmacro_f1 = f1_score(y_true, y_pred, average='macro', zero_division=0)\nmicro_f1 = f1_score(y_true, y_pred, average='micro', zero_division=0)\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"BASELINE TEST PERFORMANCE (Default 0.5 Threshold)\")\nprint(\"=\"*60)\nprint(f\"Exact Match Accuracy:  {exact_match_acc*100:6.2f}%\")\nprint(f\"Macro F1 Score:         {macro_f1:6.4f}\")\nprint(f\"Micro F1 Score:         {micro_f1:6.4f}\")\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:50.509924Z","iopub.status.idle":"2026-01-04T13:53:50.510435Z","shell.execute_reply":"2026-01-04T13:53:50.510173Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 16. DETAILED CLASSIFICATION REPORT (BASELINE) ---\nprint(\"\\nDetailed Classification Report (Default 0.5 Threshold):\")\nprint(\"=\"*80)\nreport_baseline = classification_report(\n    y_true, y_pred,\n    target_names=config.LABELS,\n    zero_division=0\n)\nprint(report_baseline)\nprint(\"=\"*80)\n\n# Save to file\nwith open('baseline_report.txt', 'w') as f:\n    f.write(f\"Baseline Performance (Default 0.5 Threshold)\\n\")\n    f.write(f\"=\"*60 + \"\\n\")\n    f.write(f\"Exact Match Accuracy: {exact_match_acc:.4f}\\n\")\n    f.write(f\"Macro F1: {macro_f1:.4f}\\n\")\n    f.write(f\"Micro F1: {micro_f1:.4f}\\n\\n\")\n    f.write(f\"Detailed Classification Report:\\n\")\n    f.write(report_baseline)\n\nprint(\"\\nBaseline report saved to baseline_report.txt\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:50.511366Z","iopub.status.idle":"2026-01-04T13:53:50.511892Z","shell.execute_reply":"2026-01-04T13:53:50.511645Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Confusion Matrix Analysis (Baseline)\n\nFor multi-label classification, we create two types of confusion matrices:\n\n1. **Combined Confusion Matrix (Dominant Class)**: Converts multi-label to single-label by taking the class with highest probability\n2. **Per-Class Confusion Matrices**: Individual 2x2 confusion matrices for each class (TP, FP, FN, TN)","metadata":{}},{"cell_type":"code","source":"# --- 17. COMBINED CONFUSION MATRIX (DOMINANT CLASS) ---\n# Get dominant class (highest probability) for each sample\npred_dominant = np.argmax(y_probs, axis=1)\ntrue_dominant = np.argmax(y_true, axis=1)\n\n# Calculate confusion matrix\ncm_combined = confusion_matrix(true_dominant, pred_dominant)\n\n# Plot\nplt.figure(figsize=(12, 10))\nsns.heatmap(\n    cm_combined,\n    annot=True,\n    fmt='d',\n    cmap='Blues',\n    xticklabels=config.LABELS,\n    yticklabels=config.LABELS,\n    cbar_kws={'label': 'Count'}\n)\nplt.xlabel('Predicted Class', fontsize=12)\nplt.ylabel('True Class', fontsize=12)\nplt.title(f'Combined Confusion Matrix (Dominant Class)\\nBaseline - Accuracy: {exact_match_acc*100:.2f}%',\n          fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.savefig('confusion_matrix_baseline.png', dpi=150, bbox_inches='tight')\nplt.show()\n\nprint(\"Combined confusion matrix saved to confusion_matrix_baseline.png\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:50.512885Z","iopub.status.idle":"2026-01-04T13:53:50.513375Z","shell.execute_reply":"2026-01-04T13:53:50.513132Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 18. PER-CLASS MULTI-LABEL CONFUSION MATRICES ---\nmcm = multilabel_confusion_matrix(y_true, y_pred)\n\nfig, axes = plt.subplots(2, 3, figsize=(18, 12))\naxes = axes.ravel()\n\nfor idx, label in enumerate(config.LABELS):\n    tn, fp, fn, tp = mcm[idx].ravel()\n\n    # Calculate metrics\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    sns.heatmap(\n        [[tp, fp], [fn, tn]],\n        annot=True,\n        fmt='d',\n        cmap='Greens',\n        xticklabels=['Pred Pos', 'Pred Neg'],\n        yticklabels=['True Pos', 'True Neg'],\n        ax=axes[idx],\n        cbar_kws={'label': 'Count'}\n    )\n\n    axes[idx].set_title(\n        f'{label}\\nPrec: {precision:.3f} | Rec: {recall:.3f} | F1: {f1:.3f}',\n        fontsize=11, fontweight='bold'\n    )\n\nplt.suptitle('Per-Class Confusion Matrices (Baseline)', fontsize=14, fontweight='bold', y=1.00)\nplt.tight_layout()\nplt.savefig('per_class_confusion_matrices_baseline.png', dpi=150, bbox_inches='tight')\nplt.show()\n\nprint(\"Per-class confusion matrices saved to per_class_confusion_matrices_baseline.png\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:50.514584Z","iopub.status.idle":"2026-01-04T13:53:50.515082Z","shell.execute_reply":"2026-01-04T13:53:50.514833Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ========================================\n# PHASE 3: THRESHOLD OPTIMIZATION\n# ========================================\n\n### What is Threshold Optimization?\n\nIn multi-label classification, using a **uniform 0.5 threshold** for all classes is suboptimal because:\n\n1. **Class imbalance**: Rare classes need lower thresholds to improve recall\n2. **Varying confidence**: Some classes have systematically higher/lower predictions\n3. **Cost asymmetry**: Missing a disease (false negative) may be worse than false alarm\n\n### How It Works\n\nFor each class, we:\n1. Collect validation set predictions\n2. Test thresholds from 0.15 to 0.85 (step 0.02)\n3. Select threshold that maximizes **F1 score**\n4. Apply class-specific thresholds for test evaluation","metadata":{}},{"cell_type":"code","source":"# --- 19. THRESHOLD OPTIMIZATION FUNCTION ---\ndef find_optimal_thresholds(model, val_loader, device, class_names):\n    \"\"\"\n    Find optimal threshold for each class using validation set.\n\n    Args:\n        model: Trained model\n        val_loader: Validation data loader\n        device: torch device\n        class_names: List of class names\n\n    Returns:\n        numpy array of optimal thresholds for each class\n    \"\"\"\n    model.eval()\n    val_probs = []\n    val_targets = []\n\n    # Collect all validation predictions\n    print(\"Collecting validation predictions...\")\n    with torch.no_grad():\n        for images, labels in tqdm(val_loader, leave=False):\n            images = images.to(device)\n            outputs = model(images)\n            probs = torch.sigmoid(outputs)\n\n            val_probs.append(probs.cpu())\n            val_targets.append(labels)\n\n    val_probs = torch.cat(val_probs).numpy()\n    val_targets = torch.cat(val_targets).numpy()\n\n    print(f\"Collected {len(val_probs)} validation samples\")\n\n    # Find optimal threshold for each class\n    optimal_thresholds = []\n    threshold_range = np.arange(0.15, 0.86, 0.02)\n\n    print(\"\\nFinding optimal thresholds for each class...\")\n    print(\"-\" * 60)\n    print(f\"{'Class':<25} {'Opt Thresh':<10} {'Best F1':<10} {'Default F1':<10} {'Improvement':<12}\")\n    print(\"-\" * 60)\n\n    for i, class_name in enumerate(class_names):\n        best_f1 = 0\n        best_thresh = 0.5\n\n        # Try each threshold\n        for thresh in threshold_range:\n            preds = (val_probs[:, i] > thresh).astype(float)\n            score = f1_score(val_targets[:, i], preds, zero_division=0)\n            if score > best_f1:\n                best_f1 = score\n                best_thresh = thresh\n\n        # Calculate F1 with default 0.5 threshold\n        default_preds = (val_probs[:, i] > 0.5).astype(float)\n        default_f1 = f1_score(val_targets[:, i], default_preds, zero_division=0)\n\n        # Calculate improvement\n        improvement = ((best_f1 - default_f1) / default_f1 * 100) if default_f1 > 0 else 0\n\n        optimal_thresholds.append(best_thresh)\n\n        # Print results\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    print(\"-\" * 60)\n    return np.array(optimal_thresholds)\n\n# Find optimal thresholds\noptimal_thresholds = find_optimal_thresholds(\n    model, valid_loader, config.DEVICE, config.LABELS\n)\n\n# Save thresholds to JSON\nthreshold_dict = {\n    name: float(thresh)\n    for name, thresh in zip(config.LABELS, optimal_thresholds)\n}\nwith open('optimal_thresholds.json', 'w') as f:\n    json.dump(threshold_dict, f, indent=2)\n\nprint(f\"\\nOptimal thresholds saved to optimal_thresholds.json\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:50.516174Z","iopub.status.idle":"2026-01-04T13:53:50.516708Z","shell.execute_reply":"2026-01-04T13:53:50.516432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 20. RE-EVALUATE WITH OPTIMAL THRESHOLDS ---\ndef evaluate_with_thresholds(model, test_loader, device, thresholds):\n    \"\"\"\n    Evaluate model with class-specific thresholds.\n    \"\"\"\n    model.eval()\n    all_targets = []\n    all_preds = []\n    all_probs = []\n\n    with torch.no_grad():\n        for images, labels in tqdm(test_loader, desc=\"Evaluating with optimal thresholds\"):\n            images = images.to(device)\n\n            outputs = model(images)\n            probs = torch.sigmoid(outputs)\n\n            # Apply class-specific thresholds\n            preds = torch.zeros_like(probs)\n            for i in range(len(config.LABELS)):\n                preds[:, i] = (probs[:, i] > thresholds[i]).float()\n\n            all_targets.append(labels.numpy())\n            all_preds.append(preds.cpu().numpy())\n            all_probs.append(probs.cpu().numpy())\n\n    all_targets = np.vstack(all_targets)\n    all_preds = np.vstack(all_preds)\n    all_probs = np.vstack(all_probs)\n\n    return all_targets, all_preds, all_probs\n\n# Evaluate with optimal thresholds\ny_true_opt, y_pred_opt, y_probs_opt = evaluate_with_thresholds(\n    model, test_loader, config.DEVICE, optimal_thresholds\n)\n\n# Calculate metrics\nexact_match_acc_opt = accuracy_score(y_true_opt, y_pred_opt)\nmacro_f1_opt = f1_score(y_true_opt, y_pred_opt, average='macro', zero_division=0)\nmicro_f1_opt = f1_score(y_true_opt, y_pred_opt, average='micro', zero_division=0)\n\n# Calculate improvements\nacc_improvement = (exact_match_acc_opt - exact_match_acc) * 100\nf1_improvement = (macro_f1_opt - macro_f1) * 100\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"OPTIMAL THRESHOLDS TEST PERFORMANCE\")\nprint(\"=\"*60)\nprint(f\"Exact Match Accuracy:  {exact_match_acc_opt*100:6.2f}% ({acc_improvement:+.2f}%)\")\nprint(f\"Macro F1 Score:         {macro_f1_opt:6.4f} ({f1_improvement:+.2f}%)\")\nprint(f\"Micro F1 Score:         {micro_f1_opt:6.4f}\")\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:50.517821Z","iopub.status.idle":"2026-01-04T13:53:50.518325Z","shell.execute_reply":"2026-01-04T13:53:50.518074Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 21. DETAILED CLASSIFICATION REPORT (OPTIMAL THRESHOLDS) ---\nprint(\"\\nDetailed Classification Report (Optimal Thresholds):\")\nprint(\"=\"*80)\nreport_optimal = classification_report(\n    y_true_opt, y_pred_opt,\n    target_names=config.LABELS,\n    zero_division=0\n)\nprint(report_optimal)\nprint(\"=\"*80)\n\n# Save to file\nwith open('optimal_threshold_report.txt', 'w') as f:\n    f.write(f\"Optimal Threshold Performance\\n\")\n    f.write(f\"=\"*60 + \"\\n\")\n    f.write(f\"Exact Match Accuracy: {exact_match_acc_opt:.4f}\\n\")\n    f.write(f\"Macro F1: {macro_f1_opt:.4f}\\n\")\n    f.write(f\"Micro F1: {micro_f1_opt:.4f}\\n\\n\")\n    \n    f.write(f\"Optimal Thresholds:\\n\")\n    for name, thresh in threshold_dict.items():\n        f.write(f\"  {name}: {thresh:.3f}\\n\")\n    f.write(f\"\\nDetailed Classification Report:\\n\")\n    f.write(report_optimal)\n\nprint(\"\\nOptimal threshold report saved to optimal_threshold_report.txt\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:50.519302Z","iopub.status.idle":"2026-01-04T13:53:50.519828Z","shell.execute_reply":"2026-01-04T13:53:50.519577Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 21. DETAILED CLASSIFICATION REPORT (OPTIMAL THRESHOLDS) ---\nprint(\"\\nDetailed Classification Report (Optimal Thresholds):\")\nprint(\"=\"*80)\nreport_optimal = classification_report(\n    y_true_opt, y_pred_opt,\n    target_names=config.LABELS,\n    zero_division=0\n)\nprint(report_optimal)\nprint(\"=\"*80)\n\n# Save to file\nwith open('optimal_threshold_report.txt', 'w') as f:\n    f.write(f\"Optimal Threshold Performance\\n\")\n    f.write(f\"=\"*60 + \"\\n\")\n    f.write(f\"Exact Match Accuracy: {exact_match_acc_opt:.4f}\\n\")\n    f.write(f\"Macro F1: {macro_f1_opt:.4f}\\n\")\n    f.write(f\"Micro F1: {micro_f1_opt:.4f}\\n\\n\")\n\n    f.write(f\"Optimal Thresholds:\\n\")\n    for name, thresh in threshold_dict.items():\n        f.write(f\"  {name}: {thresh:.3f}\\n\")\n    f.write(f\"\\nDetailed Classification Report:\\n\")\n    f.write(report_optimal)\n\nprint(\"\\nOptimal threshold report saved to optimal_threshold_report.txt\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:50.520872Z","iopub.status.idle":"2026-01-04T13:53:50.521360Z","shell.execute_reply":"2026-01-04T13:53:50.521117Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 22. THRESHOLD COMPARISON VISUALIZATION ---\n# Per-class recall comparison\nrecall_baseline = []\nrecall_optimal = []\n\nfor i in range(len(config.LABELS)):\n    # Baseline recall\n    tp = ((y_true[:, i] == 1) & (y_pred[:, i] == 1)).sum()\n    fn = ((y_true[:, i] == 1) & (y_pred[:, i] == 0)).sum()\n    rec_baseline = tp / (tp + fn) if (tp + fn) > 0 else 0\n    recall_baseline.append(rec_baseline)\n\n    # Optimal recall\n    tp = ((y_true_opt[:, i] == 1) & (y_pred_opt[:, i] == 1)).sum()\n    fn = ((y_true_opt[:, i] == 1) & (y_pred_opt[:, i] == 0)).sum()\n    rec_optimal = tp / (tp + fn) if (tp + fn) > 0 else 0\n    recall_optimal.append(rec_optimal)\n\n# Create comparison plot\nfig, axes = plt.subplots(1, 2, figsize=(16, 6))\n\n# Plot 1: Recall comparison\nx = np.arange(len(config.LABELS))\nwidth = 0.35\n\naxes[0].bar(\n    x - width/2,\n    [r*100 for r in recall_baseline],\n    width,\n    label='Baseline (0.5)',\n    alpha=0.8,\n    color='coral'\n)\naxes[0].bar(\n    x + width/2,\n    [r*100 for r in recall_optimal],\n    width,\n    label='Optimal Thresholds',\n    alpha=0.8,\n    color='steelblue'\n)\naxes[0].set_xlabel('Class', fontsize=12)\naxes[0].set_ylabel('Recall (%)', fontsize=12)\naxes[0].set_title('Recall: Baseline vs Optimal Thresholds', fontsize=14, fontweight='bold')\naxes[0].set_xticks(x)\naxes[0].set_xticklabels(config.LABELS, rotation=45, ha='right')\naxes[0].legend(fontsize=11)\naxes[0].grid(axis='y', alpha=0.3)\naxes[0].set_ylim(0, 105)\n\n# Add improvement annotations\nfor i, (before, after) in enumerate(zip(recall_baseline, recall_optimal)):\n    delta = (after - before) * 100\n    if abs(delta) > 3:\n        axes[0].annotate(\n            f'{delta:+.1f}%',\n            xy=(i + width/2, after*100),\n            xytext=(i + width/2, after*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\ncolors = [\n    'green' if t < 0.5\n    else 'orange' if t < 0.6\n    else 'red'\n    for t in optimal_thresholds\n]\nbars = axes[1].bar(config.LABELS, optimal_thresholds, color=colors, alpha=0.7)\naxes[1].axhline(y=0.5, color='black', linestyle='--', linewidth=2, label='Default (0.5)')\naxes[1].set_xlabel('Class', fontsize=12)\naxes[1].set_ylabel('Optimal Threshold', fontsize=12)\naxes[1].set_title('Optimal Threshold per Class', fontsize=14, fontweight='bold')\naxes[1].set_xticks(x)\naxes[1].set_xticklabels(config.LABELS, rotation=45, ha='right')\naxes[1].legend(fontsize=11)\naxes[1].grid(axis='y', alpha=0.3)\naxes[1].set_ylim(0, 1.0)\n\n# Add threshold values on bars\nfor bar, thresh in zip(bars, optimal_thresholds):\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\nplt.tight_layout()\nplt.savefig('threshold_optimization_comparison.png', dpi=150, bbox_inches='tight')\nplt.show()\n\nprint(\"Threshold comparison plot saved to threshold_optimization_comparison.png\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ========================================\n# PHASE 5: VISUAL PREDICTIONS ON TEST IMAGES\n# ========================================\n\nLet's visualize the model's predictions on random test images to see how well it performs in practice.","metadata":{}},{"cell_type":"code","source":"# --- 26. VISUALIZE PREDICTIONS ON RANDOM TEST IMAGES ---\ndef plot_predictions(model, test_df, num_images=15, thresholds=None):\n    \"\"\"\n    Plot random test images with predictions and ground truth labels.\n    \"\"\"\n    # Sample random images\n    sample = test_df.sample(n=num_images, random_state=42)\n\n    rows = int(np.ceil(num_images / 5))\n    fig, axes = plt.subplots(rows, 5, figsize=(20, rows * 4))\n    axes = axes.ravel() if num_images > 1 else [axes]\n\n    model.eval()\n\n    with torch.no_grad():\n        for i, (idx, row) in enumerate(sample.iterrows()):\n            # Load image\n            image_path = row['image']\n            image = Image.open(image_path).convert(\"RGB\")\n\n            # Preprocess\n            img_tensor = test_transforms(image).unsqueeze(0).to(config.DEVICE)\n\n            # Predict\n            outputs = model(img_tensor)\n            probs = torch.sigmoid(outputs)\n\n            # Apply thresholds\n            if thresholds is not None:\n                preds = (probs > torch.tensor(thresholds).to(config.DEVICE)).int()\n            else:\n                preds = (probs > 0.5).int()\n\n            # Get labels\n            true_labels = row['labels'].split(' ')\n            pred_labels = [config.LABELS[j] for j, val in enumerate(preds[0]) if val == 1]\n            if not pred_labels:\n                pred_labels = ['healthy']\n\n            # Plot\n            axes[i].imshow(image)\n\n            # Color code based on correctness\n            correct = set(true_labels) == set(pred_labels)\n            color = 'green' if correct else 'red'\n\n            axes[i].set_title(\n                f\"True: {', '.join(true_labels)}\\nPred: {', '.join(pred_labels)}\",\n                fontsize=9,\n                color=color,\n                fontweight='bold'\n            )\n            axes[i].axis('off')\n\n    plt.suptitle(\n        f'Model Predictions on Test Images\\n(Optimal Thresholds)',\n        fontsize=16,\n        fontweight='bold'\n    )\n    plt.tight_layout()\n    plt.savefig('visual_predictions_optimal.png', dpi=150, bbox_inches='tight')\n    plt.show()\n\nprint(\"\\nGenerating predictions on random test images...\")\nprint(\"(Green title = correct, Red title = incorrect)\")\nplot_predictions(model, test_df, num_images=15, thresholds=optimal_thresholds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:50.525670Z","iopub.status.idle":"2026-01-04T13:53:50.526163Z","shell.execute_reply":"2026-01-04T13:53:50.525920Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 27. PER-CLASS ACCURACY ANALYSIS ---\n# Calculate per-class accuracy with optimal thresholds\nresults_per_class = []\n\nfor i, label in enumerate(config.LABELS):\n    # Baseline\n    acc_base = accuracy_score(y_true[:, i], y_pred[:, i])\n    f1_base = f1_score(y_true[:, i], y_pred[:, i], zero_division=0)\n\n    # Optimal\n    acc_opt = accuracy_score(y_true_opt[:, i], y_pred_opt[:, i])\n    f1_opt = f1_score(y_true_opt[:, i], y_pred_opt[:, i], zero_division=0)\n\n    results_per_class.append({\n        'Class': label,\n        'Optimal Threshold': optimal_thresholds[i],\n        'Baseline Acc': acc_base,\n        'Optimal Acc': acc_opt,\n        'Baseline F1': f1_base,\n        'Optimal F1': f1_opt\n    })\n\ndf_per_class = pd.DataFrame(results_per_class)\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"PER-CLASS PERFORMANCE ANALYSIS\")\nprint(\"=\"*80)\nprint(df_per_class.to_string(index=False))\nprint(\"=\"*80)\n\n# Save to CSV\ndf_per_class.to_csv('per_class_performance.csv', index=False)\nprint(\"\\nPer-class performance saved to per_class_performance.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T13:53:50.527228Z","iopub.status.idle":"2026-01-04T13:53:50.527757Z","shell.execute_reply":"2026-01-04T13:53:50.527486Z"}},"outputs":[],"execution_count":null}]}