{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","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,"sourceType":"competition"},{"sourceId":121906917,"sourceType":"kernelVersion"},{"sourceId":670827,"sourceType":"modelInstanceVersion","modelInstanceId":508097,"modelId":522766}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\n!pip install imutils","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:42.817250Z","iopub.execute_input":"2026-01-04T03:33:42.817554Z","iopub.status.idle":"2026-01-04T03:33:45.987459Z","shell.execute_reply.started":"2026-01-04T03:33:42.817522Z","shell.execute_reply":"2026-01-04T03:33:45.986727Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport math\nimport random\nfrom typing import Dict, List,Tuple\nimport requests\nfrom pathlib import Path\n\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport glob\nfrom pathlib import Path, PurePath\nimport pathlib\nimport pandas as pd\n\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn as nn\nimport json\nimport torchvision\nfrom torchvision import datasets\nfrom torchvision import transforms\nimport torch.nn.functional as F\n\n\nimport seaborn as sns\nfrom sklearn.metrics import classification_report, multilabel_confusion_matrix, confusion_matrix, f1_score, precision_score\n\nfrom PIL import Image\n\nfrom sklearn.model_selection import train_test_split\n\nfrom imutils import paths\n\nimport textwrap\nfrom tqdm import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:45.989030Z","iopub.execute_input":"2026-01-04T03:33:45.989266Z","iopub.status.idle":"2026-01-04T03:33:45.995520Z","shell.execute_reply.started":"2026-01-04T03:33:45.989243Z","shell.execute_reply":"2026-01-04T03:33:45.994916Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\nimport torch\n\nclass config:\n    # specify the paths to datasets\n    DATA_DIR = Path('../input/plant-pathology-2021-fgvc8/train_images')\n    ROOT_DIR = Path('./data')\n    CSV_DIR = Path('../input/plant-pathology-2021-fgvc8/train.csv')\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    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\n    MODEL_NAME = 'resnet_scratch'  \n\n    DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n    LABELS = ['complex', 'frog_eye_leaf_spot', 'healthy', 'powdery_mildew', 'rust', 'scab']\n    NUM_CLASSES = len(LABELS)\n\n    CHECKPOINT_DIR = 'checkpoints'\n    print(\"Using device:\", DEVICE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:45.996089Z","iopub.execute_input":"2026-01-04T03:33:45.996252Z","iopub.status.idle":"2026-01-04T03:33:46.017808Z","shell.execute_reply.started":"2026-01-04T03:33:45.996239Z","shell.execute_reply":"2026-01-04T03:33:46.017142Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nfrom sklearn.model_selection import train_test_split\n\ndef split_df(csv_dir):\n    '''\n    This function take csv file and split it into train, valid, and test\n    '''\n    df = pd.read_csv(csv_dir)\n    #df['image'] =  df['image'].apply(lambda x: '/kaggle/working/data' + x) # Use when run resized by yourself\n    df['image'] =  df['image'].apply(lambda x: '../input/multi-label-classification-plant-pathology/data/' + x)\n\n    # train dataframe\n    train_df, dummy_df = train_test_split(df,  train_size= 0.7, shuffle= True, random_state= 42)\n\n    # valid and test dataframe\n    valid_df, test_df = train_test_split(dummy_df,  train_size= 0.5, shuffle= True, random_state= 42)\n\n    return train_df, valid_df, test_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:46.018498Z","iopub.execute_input":"2026-01-04T03:33:46.018743Z","iopub.status.idle":"2026-01-04T03:33:46.033056Z","shell.execute_reply.started":"2026-01-04T03:33:46.018723Z","shell.execute_reply":"2026-01-04T03:33:46.032422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df, valid_df, test_df = split_df('../input/plant-pathology-2021-fgvc8/train.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:46.035022Z","iopub.execute_input":"2026-01-04T03:33:46.035217Z","iopub.status.idle":"2026-01-04T03:33:46.077083Z","shell.execute_reply.started":"2026-01-04T03:33:46.035201Z","shell.execute_reply":"2026-01-04T03:33:46.076574Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"(train_df['labels'].value_counts()).plot(kind='bar')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:46.077802Z","iopub.execute_input":"2026-01-04T03:33:46.078109Z","iopub.status.idle":"2026-01-04T03:33:46.273164Z","shell.execute_reply.started":"2026-01-04T03:33:46.078092Z","shell.execute_reply":"2026-01-04T03:33:46.272560Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_labels = train_df['labels'].str.split(expand=True).stack().reset_index(drop=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:46.273917Z","iopub.execute_input":"2026-01-04T03:33:46.274190Z","iopub.status.idle":"2026-01-04T03:33:46.291791Z","shell.execute_reply.started":"2026-01-04T03:33:46.274172Z","shell.execute_reply":"2026-01-04T03:33:46.291173Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_labels.value_counts().plot(kind='bar')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:46.292887Z","iopub.execute_input":"2026-01-04T03:33:46.293189Z","iopub.status.idle":"2026-01-04T03:33:46.442288Z","shell.execute_reply.started":"2026-01-04T03:33:46.293170Z","shell.execute_reply":"2026-01-04T03:33:46.441719Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"There's some class imbalance happening here. \n\nThis will needed to be handled when we define our loss function.","metadata":{}},{"cell_type":"code","source":"import math\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom pathlib import Path\ndef examine_images(df, num_images=20):\n    image_paths = df['image'].sample(n=num_images, random_state=42)\n    labels = df['labels'].loc[image_paths.index]\n\n    num_rows = int(math.ceil(num_images/5))\n    num_cols = 5\n    fig, axs = plt.subplots(num_rows, num_cols, figsize=(30, 30),tight_layout=True)\n    axs = axs.ravel()\n\n    for i, image_path in enumerate(image_paths):\n        image = Image.open(Path(image_path))\n        label = labels.iloc[i]\n        axs[i].imshow(image)\n        axs[i].set_title(f\"Label: {label}\", fontsize=25)\n        axs[i].axis('off')\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:46.443051Z","iopub.execute_input":"2026-01-04T03:33:46.443349Z","iopub.status.idle":"2026-01-04T03:33:46.449002Z","shell.execute_reply.started":"2026-01-04T03:33:46.443332Z","shell.execute_reply":"2026-01-04T03:33:46.448248Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"examine_images(train_df, num_images=20)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:46.449840Z","iopub.execute_input":"2026-01-04T03:33:46.450105Z","iopub.status.idle":"2026-01-04T03:33:52.072375Z","shell.execute_reply.started":"2026-01-04T03:33:46.450087Z","shell.execute_reply":"2026-01-04T03:33:52.071070Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision import transforms\n\n# initialize our data augmentation functions\nresize = transforms.Resize(size=(config.INPUT_HEIGHT,config.INPUT_WIDTH))\nmake_tensor = transforms.ToTensor()\nnormalize = transforms.Normalize(mean=config.IMAGENET_MEAN, std=config.IMAGENET_STD)\ncenter_cropper = transforms.CenterCrop((config.INPUT_HEIGHT,config.INPUT_WIDTH))\nrandom_resized_crop = transforms.RandomResizedCrop(size=(config.INPUT_HEIGHT, config.INPUT_WIDTH), scale=(0.6, 1.0))\nrandom_horizontal_flip = transforms.RandomHorizontalFlip(p=0.75)\nrandom_vertical_flip = transforms.RandomVerticalFlip(p=0.75)\nrandom_rotation = transforms.RandomRotation(degrees=90)\nrandom_crop = transforms.RandomCrop(size=(200,200))\naugmix = transforms.AugMix(severity = 3, mixture_width=3, alpha=0.2)\nauto_augment = transforms.AutoAugment()\nrandom_augment = transforms.RandAugment()\ncolor_jitter = transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2)\n\n# Stronger augmentations for training from scratch\ntrain_transforms = transforms.Compose([\n    # Randomly crop and resize (simulates different distances/zoom)\n    random_resized_crop,\n    random_horizontal_flip,\n    random_vertical_flip,\n    random_rotation,\n    # Jitter brightness/contrast/saturation (simulates different lighting)\n    color_jitter,\n    make_tensor,\n    normalize\n])\n\nval_transforms = transforms.Compose([resize, make_tensor, normalize])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:52.073699Z","iopub.execute_input":"2026-01-04T03:33:52.074052Z","iopub.status.idle":"2026-01-04T03:33:52.089023Z","shell.execute_reply.started":"2026-01-04T03:33:52.074019Z","shell.execute_reply":"2026-01-04T03:33:52.088004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport torchvision\nfrom PIL import Image\n\ndef apply_transform(img: Image, transform) -> np.ndarray:\n    \"\"\"\n    Applies a transform to a PIL Image and returns a numpy array of the transformed image.\n\n    Args:\n        img (PIL.Image): The input image to transform.\n        transform (torchvision.transforms.Compose): The transform to apply to the image.\n\n    Returns:\n        np.ndarray: A numpy array representing the transformed image.\n    \"\"\"\n    # Apply the transform to the image\n    if isinstance(transform, torchvision.transforms.Compose):\n        # Apply PyTorch transform to image array\n        transformed_image = train_transforms(img)\n\n    elif isinstance(transform, A.Compose):\n        # Apply Albumentations transform to image array\n        img_array = np.array(img)\n        transformed_image = transform(image=img_array)[\"image\"]\n\n    # Convert the image tensor to a numpy array and transpose the axes to (height, width, channels)\n    img_array = transformed_image.numpy().transpose((1, 2, 0))\n\n    # Clip the pixel values to the range [0, 1]\n    img_array = np.clip(img_array, 0, 1)\n\n    return img_array\n\n\ndef visualize_transform(image: np.ndarray, original_image: np.ndarray = None) -> None:\n    \"\"\"\n    Visualize the transformed image.\n\n    Args:\n        image (np.ndarray): A NumPy array representing the transformed image.\n        original_image (np.ndarray, optional): A NumPy array representing the original image. Defaults to None.\n    \"\"\"\n    fontsize = 18\n    \n    if original_image is None:\n        # Create a plot with 1 row and 2 columns.\n        f, ax = plt.subplots(1, 2, figsize=(12, 12))\n\n        # Show the transformed image in the first column.\n        ax[0].imshow(image)\n    else:\n        # Create a plot with 1 row and 2 columns.\n        f, ax = plt.subplots(1, 2, figsize=(12, 12))\n\n        # Show the original image in the first column.\n        ax[0].imshow(original_image)\n        ax[0].set_title('Original image', fontsize=fontsize)\n        \n        # Show the transformed image in the second column.\n        ax[1].imshow(image)\n        ax[1].set_title('Transformed image', fontsize=fontsize)\n        \nimg = Image.open(train_df['image'].sample(n=1).iloc[0])\nimg_array = apply_transform(img, train_transforms)\nvisualize_transform(img_array, original_image=img)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:52.090927Z","iopub.execute_input":"2026-01-04T03:33:52.091115Z","iopub.status.idle":"2026-01-04T03:33:52.689092Z","shell.execute_reply.started":"2026-01-04T03:33:52.091100Z","shell.execute_reply":"2026-01-04T03:33:52.688347Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def encode_label(labels, class_list):\n    \"\"\"Encode a list of labels using one-hot encoding.\n\n    Args:\n        label: A list of labels to encode.\n        class_list: A list of all possible labels. Defaults to DEFAULT_LABELS.\n\n    Returns:\n        A tensor representing the one-hot encoding of the input labels.\n    \"\"\"\n    # Create a tensor of zeros with the same length as the class list\n    target = torch.zeros(len(class_list))\n    for label in labels:\n        # Find the index of the current label in the class list\n        idx = class_list.index(label)\n        # Set the corresponding index in the target tensor to 1\n        target[idx] = 1\n    return target\n\n\n\ndef decode_label(encoded_label, class_list):\n    \"\"\"Decode a one-hot encoded label into its original label(s).\n\n    Args:\n        encoded_label: A tensor representing the one-hot encoding of a label.\n        class_list: A list of all possible labels. Defaults to DEFAULT_LABELS.\n\n    Returns:\n        A list of the decoded label(s).\n    \"\"\"\n    # Use a list comprehension to create the decoded list\n    decoded = [class_list[i] for i, val in enumerate(encoded_label) if val == 1]\n\n    # Return the list of decoded label(s)\n    return decoded","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:52.689952Z","iopub.execute_input":"2026-01-04T03:33:52.690302Z","iopub.status.idle":"2026-01-04T03:33:52.695438Z","shell.execute_reply.started":"2026-01-04T03:33:52.690277Z","shell.execute_reply":"2026-01-04T03:33:52.694638Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import Dataset\nfrom PIL import Image\n\nclass PlantDataset(Dataset):\n    def __init__(self, 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        image_path = self.dataframe['image'].iloc[idx]\n        image = Image.open(image_path)\n        labels = self.dataframe.iloc[idx]['labels'].split(' ')\n        encoded_labels = encode_label(labels, config.LABELS)\n        if self.transform:\n            image = self.transform(image)\n        return image, encoded_labels","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:52.698551Z","iopub.execute_input":"2026-01-04T03:33:52.698806Z","iopub.status.idle":"2026-01-04T03:33:52.715164Z","shell.execute_reply.started":"2026-01-04T03:33:52.698789Z","shell.execute_reply":"2026-01-04T03:33:52.714373Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\ntrain_dataset = PlantDataset(train_df, transform = train_transforms)\nval_dataset = PlantDataset(valid_df, transform = val_transforms)\ntest_dataset = PlantDataset(test_df, transform = val_transforms)\n\ntrain_loader = DataLoader(train_dataset, batch_size=config.BATCH_SIZE, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=config.BATCH_SIZE)\ntest_loader = DataLoader(test_dataset, batch_size=config.BATCH_SIZE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:52.715911Z","iopub.execute_input":"2026-01-04T03:33:52.716141Z","iopub.status.idle":"2026-01-04T03:33:52.732042Z","shell.execute_reply.started":"2026-01-04T03:33:52.716119Z","shell.execute_reply":"2026-01-04T03:33:52.731329Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Multi-label classification loss function\n\nFor multilabel classification problems, where each input can have multiple possible output labels, the BCEWithLogitsLoss is a good choice for the loss function.\n\nBCEWithLogitsLoss combines the sigmoid activation function and the binary cross-entropy loss into a single function, which makes the computation more numerically stable and efficient. It takes as input the logits tensor and the true targets tensor, both of the same shape. It first applies the sigmoid function to the logit tensor to obtain the predicted probabilities. \n\nThen, it computes the binary cross-entropy loss between the predicted and true targets, which measures the dissimilarity between the predicted and true probabilities.\n\nDuring training, the goal is to minimize this loss by adjusting the model parameters using backpropagation so that the model can make more accurate predictions. In PyTorch, you can use the BCEWithLogitsLoss function, which combines the sigmoid and BCE loss functions to efficiently compute both the activation and the loss in a single forward pass. The BCEWithLogitsLoss function applies the sigmoid activation function to the logits, which are the unbounded real-valued outputs of the model. \n\nThen it computes the BCE loss between the sigmoid activations and the target labels.\n\nAnother option is the MultiLabelSoftMarginLoss function, which works well for multilabel classification problems. This loss function applies the soft-margin version of the sigmoid function to the logits. It computes the negative log-likelihood of the targets under the predicted probability distribution.\n\nThese choices come down to the specific characteristics of your problem and the architecture of your model. Experimenting with different loss functions and seeing which performs best for your task may be helpful.\n\n## Handling class imbalance\n\n\nFor this example I'll employ class weighting.\n\nThis is an approach for handling class imbalance in multilabel classification. This method assigns higher weights to the minority classes and lower weights to the majority classes to balance their representation in the loss function. \n\n\n### How to choose the class weights\n\nLike eveyrthing in depe learning, how to weight the classes depends on your specific problem and just how imbalanced your classes are. \n\nYou've got a few options available to you for figuring out how to weight the classes:\n\n1) **Manual assignment**: For example, if there are three classes with class frequencies of 0.2, 0.3, and 0.5, you could assign arbitrary class weights of [3.5, 2.25, 1.23] to balance their representation in the loss function.\n\n2) **Inverse class frequency**: You could also weight the classes based on the inverse of their frequency in the training set. This gives classes with fewer examples higher weights and the classes with more example lower weights. For example, if there are three classes with frequencies of 0.2, 0.3, and 0.5, you assign class weights of [2.5, 1.67, 1] based on their inverse frequencies.\n\n3) **Automatic assignment**: You could just let the machines decide for you. Some machine learning frameworks and libraries have built-in functions to automatically calculate the class weights based on the class frequencies or other metrics. For example, `scikit-learn`'s compute_class_weight function can calculate the class weights based on the inverse of their frequency or the square root of their frequency, among other methods.","metadata":{}},{"cell_type":"code","source":"#get the class counts\nclass_counts = all_labels.value_counts()[config.LABELS]\nclass_counts","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:52.732605Z","iopub.execute_input":"2026-01-04T03:33:52.732825Z","iopub.status.idle":"2026-01-04T03:33:52.753146Z","shell.execute_reply.started":"2026-01-04T03:33:52.732810Z","shell.execute_reply":"2026-01-04T03:33:52.752542Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\n\n# Compute inverse class frequency\nclass_weights = torch.reciprocal(torch.tensor(class_counts.values).float()) # invert the counts and convert them to floats\nclass_weights /= torch.max(class_weights) # normalize the weights by the maximum weight\n\n# Define loss function using class weights\ncriterion = nn.BCEWithLogitsLoss(pos_weight=class_weights)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:52.754307Z","iopub.execute_input":"2026-01-04T03:33:52.754851Z","iopub.status.idle":"2026-01-04T03:33:52.766877Z","shell.execute_reply.started":"2026-01-04T03:33:52.754834Z","shell.execute_reply":"2026-01-04T03:33:52.766221Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_weights","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:52.767582Z","iopub.execute_input":"2026-01-04T03:33:52.767826Z","iopub.status.idle":"2026-01-04T03:33:52.781989Z","shell.execute_reply.started":"2026-01-04T03:33:52.767807Z","shell.execute_reply":"2026-01-04T03:33:52.781302Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"In the above code:\n\n- Use pandas `Series.value_counts()` function to compute the frequency of each class and use the list `config.LABELS` list to order the counts.\n- Invert the class counts and normalize them by the maximum weight to obtain the class weights inversely proportional to their frequency.\n- Define a `BCEWithLogitsLoss` loss function using the `pos_weight` argument to pass the class weights.\n\nThe `pos_weight` argument of the `BCEWithLogitsLoss` loss function in PyTorch represents the weight of positive examples in the loss calculation. In binary classification problems, the `pos_weight` can be used to address class imbalance by giving more weight to positive examples than negative examples. In multilabel classification problems, the `pos_weight` can be used to address class imbalance by giving more weight to less frequent classes than more frequent classes.\n\nThe `pos_weight` argument is a tensor of weights that has the same shape as the target tensor. \n\nThe weights are applied to each element of the target tensor proportionally to the weight assigned to its corresponding class. In our example, the `pos_weight` tensor has values `[0.5968, 0.2913, 0.2767, 1.0000, 0.6120, 0.2216]`, where each value corresponds to a class.","metadata":{}},{"cell_type":"code","source":"import torchvision.transforms as transforms\n\n# Stronger augmentations for training from scratch\ntrain_transforms = transforms.Compose([\n    # Randomly crop and resize (simulates different distances/zoom)\n    transforms.RandomResizedCrop(size=(config.INPUT_HEIGHT, config.INPUT_WIDTH), scale=(0.6, 1.0)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(degrees=30),\n    # Jitter brightness/contrast/saturation (simulates different lighting)\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=config.IMAGENET_MEAN, std=config.IMAGENET_STD)\n])\n\nval_transforms = transforms.Compose([\n    transforms.Resize(size=(config.INPUT_HEIGHT, config.INPUT_WIDTH)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=config.IMAGENET_MEAN, std=config.IMAGENET_STD)\n])\n\n# Re-initialize datasets with these new transforms\ntrain_dataset = PlantDataset(train_df, transform=train_transforms)\nval_dataset = PlantDataset(valid_df, transform=val_transforms)\ntest_dataset = PlantDataset(test_df, transform=val_transforms)\n\n# FIXED: Set num_workers=0 for Kaggle compatibility\n# Kaggle has limited CPU cores; multiple workers can cause multiprocessing issues and crashes\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=0)\nval_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:52.782672Z","iopub.execute_input":"2026-01-04T03:33:52.783594Z","iopub.status.idle":"2026-01-04T03:33:52.793990Z","shell.execute_reply.started":"2026-01-04T03:33:52.783577Z","shell.execute_reply":"2026-01-04T03:33:52.793407Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(config.DEVICE)\n\nclass ResidualBlock(nn.Module):\n    def __init__(self, in_ch, out_ch, stride=1):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_ch, out_ch, 3, stride, 1, bias=False)\n        self.bn1 = nn.BatchNorm2d(out_ch)\n        self.conv2 = nn.Conv2d(out_ch, out_ch, 3, 1, 1, bias=False)\n        self.bn2 = nn.BatchNorm2d(out_ch)\n        self.shortcut = nn.Sequential()\n        if stride != 1 or in_ch != out_ch:\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_ch, out_ch, 1, stride, bias=False),\n                nn.BatchNorm2d(out_ch)\n            )\n    def forward(self, x):\n        identity = x\n        out = F.relu(self.bn1(self.conv1(x)))\n        out = self.bn2(self.conv2(out))\n        out += self.shortcut(identity)\n        return F.relu(out)\ndef make_layer(block, in_ch, out_ch, num_blocks, stride):\n    layers = [block(in_ch, out_ch, stride)]\n    for _ in range(1, num_blocks):\n        layers.append(block(out_ch, out_ch))\n    return nn.Sequential(*layers)\n\nclass ResNetScratch(nn.Module):\n    def __init__(self, num_classes=config.NUM_CLASSES):\n        super().__init__()\n        self.conv1 = nn.Conv2d(3, 64, 7, 2, 3, bias=False)\n        self.bn1 = nn.BatchNorm2d(64)\n        self.pool = nn.MaxPool2d(3, 2, 1)\n\n        self.layer1 = make_layer(ResidualBlock, 64, 64, 2, 1)\n        self.layer2 = make_layer(ResidualBlock, 64, 128, 2, 2)\n        self.layer3 = make_layer(ResidualBlock, 128, 256, 2, 2)\n        self.layer4 = make_layer(ResidualBlock, 256, 512, 2, 2)\n\n        self.gap = nn.AdaptiveAvgPool2d((1,1))\n        self.fc = nn.Linear(512, num_classes)\n\n    def forward(self, x):\n        x = F.relu(self.bn1(self.conv1(x)))\n        x = self.pool(x)\n        x = self.layer1(x); x = self.layer2(x)\n        x = self.layer3(x); x = self.layer4(x)\n        x = self.gap(x).flatten(1)\n        return self.fc(x)\n        \n# FIXED: Initialize model and move to device consistently\nmodel = ResNetScratch(num_classes=config.NUM_CLASSES).to(device)\nprint(f\"Model initialized and moved to device: {device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:52.794653Z","iopub.execute_input":"2026-01-04T03:33:52.794974Z","iopub.status.idle":"2026-01-04T03:33:52.944961Z","shell.execute_reply.started":"2026-01-04T03:33:52.794957Z","shell.execute_reply":"2026-01-04T03:33:52.944316Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# Use the loss function you already defined in your notebook (handling class imbalance)\ncriterion = nn.BCEWithLogitsLoss(pos_weight=class_weights.to(device))\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-4)\n\n# Use get_last_lr() if you need to access the learning rate during training\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=3)\n\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    UPDATED: Now returns F1 score for proper model selection in multi-label classification\n    \"\"\"\n    model.train()\n    running_loss = 0.0\n    correct_preds = 0\n    total_preds = 0\n    \n    # Collect all predictions and labels for F1 calculation\n    all_preds = []\n    all_labels = []\n\n    loop = tqdm(loader, leave=True)\n    for images, labels in loop:\n        images, labels = images.to(device), labels.to(device)\n\n        # Forward\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n\n        # Backprop\n        optimizer.zero_grad()\n        loss.backward()\n\n        # This was accumulating gradient norms every batch, causing Kaggle memory issues\n\n        optimizer.step()\n\n        # Metrics\n        running_loss += loss.item()\n        preds = (torch.sigmoid(outputs) > 0.5).float()\n        correct_preds += (preds == labels).float().sum()\n        total_preds += labels.numel()\n        \n        # Store for F1 calculation\n        all_preds.append(preds.cpu())\n        all_labels.append(labels.cpu())\n\n        loop.set_description(f\"Loss: {loss.item():.4f}\")\n\n    avg_loss = running_loss / len(loader)\n    avg_acc = correct_preds / total_preds\n    \n    # Calculate macro F1 score\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    \n    return avg_loss, avg_acc, avg_f1\n\ndef validate(model, loader, criterion, device):\n    \"\"\"\n    Validate the model.\n    Returns: (avg_loss, avg_accuracy, avg_f1_macro)\n    \n    UPDATED: Now returns F1 score for proper model selection in multi-label classification\n    \"\"\"\n    model.eval()\n    running_loss = 0.0\n    correct_preds = 0\n    total_preds = 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()\n            preds = (torch.sigmoid(outputs) > 0.5).float()\n            correct_preds += (preds == labels).float().sum()\n            total_preds += labels.numel()\n            \n            # Store for F1 calculation\n            all_preds.append(preds.cpu())\n            all_labels.append(labels.cpu())\n\n            loop.set_postfix(loss=loss.item())\n\n    avg_loss = running_loss / len(loader)\n    avg_acc = correct_preds / total_preds\n    \n    # Calculate macro F1 score\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    \n    return avg_loss, avg_acc, avg_f1\n\nprint(\"Training functions defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:52.945612Z","iopub.execute_input":"2026-01-04T03:33:52.945869Z","iopub.status.idle":"2026-01-04T03:33:52.956657Z","shell.execute_reply.started":"2026-01-04T03:33:52.945846Z","shell.execute_reply":"2026-01-04T03:33:52.956125Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_epochs = 10\n\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')\ncheckpoint_path = 'best_model.pth'\n\nprint(\"=\"*60)\nprint(f\"UPDATED: Model selection now based on validation F1 score\")\nprint(f\"This ensures we select the best model for multi-label classification\")\nprint(f\"=\"*60)\n\nfor epoch in range(num_epochs):\n    train_loss, train_acc, train_f1 = train_one_epoch(model, train_loader, criterion, optimizer, device)\n    val_loss, val_acc, val_f1 = validate(model, val_loader, criterion, device)\n\n    scheduler.step(val_loss)\n\n    # Track history\n    history['train_loss'].append(train_loss)\n    history['train_acc'].append(train_acc.item() if torch.is_tensor(train_acc) else train_acc)\n    history['train_f1'].append(train_f1)  # UPDATED: Track F1\n    history['val_loss'].append(val_loss)\n    history['val_acc'].append(val_acc.item() if torch.is_tensor(val_acc) else val_acc)\n    history['val_f1'].append(val_f1)      # UPDATED: Track F1\n\n    # UPDATED: Save best model based on validation F1 (not loss)\n    if val_f1 > best_val_f1:\n        best_val_f1 = val_f1\n        torch.save(model.state_dict(), checkpoint_path)\n        print(f\"-> Saved best model with val_f1: {val_f1:.4f}\")\n    else:\n        print(f\"-> Val F1 ({val_f1:.4f}) did not improve from {best_val_f1:.4f}\")\n\n    print(f\"\\nEpoch [{epoch+1}/{num_epochs}] \"\n          f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | Train F1: {train_f1:.4f}\\n\"\n          f\"Val Loss: {val_loss:.4f}     | Val Acc: {val_acc:.4f}     | Val F1: {val_f1:.4f}\")\n\nprint(f\"\\nTraining complete! Best validation F1: {best_val_f1:.4f}\")  # UPDATED: Show F1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-04T03:33:52.957489Z","iopub.execute_input":"2026-01-04T03:33:52.957814Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ==========================================\n## PHASE 3: TRAINING VISUALIZATION\n## ==========================================","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport json\n\ndef plot_training_history(history, save_path=None):\n    \"\"\"\n    Plot training and validation metrics\n\n    Args:\n        history: Dictionary with 'train_loss', 'train_acc', 'train_f1', 'val_loss', 'val_acc', 'val_f1'\n        save_path: Optional path to save the figure\n    \"\"\"\n    epochs = range(1, len(history['train_loss']) + 1)\n\n    fig, axes = plt.subplots(1, 3, figsize=(20, 5))\n\n    # Plot Loss\n    axes[0].plot(epochs, history['train_loss'], 'b-o', label='Train Loss', linewidth=2, markersize=6)\n    axes[0].plot(epochs, history['val_loss'], 'r-s', label='Val Loss', linewidth=2, markersize=6)\n    axes[0].set_title('Training and Validation Loss', fontsize=14, fontweight='bold')\n    axes[0].set_xlabel('Epoch', fontsize=12)\n    axes[0].set_ylabel('Loss (BCEWithLogits)', fontsize=12)\n    axes[0].legend(fontsize=11)\n    axes[0].grid(True, alpha=0.3, linestyle='--')\n\n    # Plot Accuracy\n    axes[1].plot(epochs, history['train_acc'], 'b-o', label='Train Acc', linewidth=2, markersize=6)\n    axes[1].plot(epochs, history['val_acc'], 'r-s', label='Val Acc', linewidth=2, markersize=6)\n    axes[1].set_title('Training and Validation Accuracy', fontsize=14, fontweight='bold')\n    axes[1].set_xlabel('Epoch', fontsize=12)\n    axes[1].set_ylabel('Accuracy', fontsize=12)\n    axes[1].legend(fontsize=11)\n    axes[1].grid(True, alpha=0.3, linestyle='--')\n\n    # Plot F1 Score \n    axes[2].plot(epochs, history['train_f1'], 'b-o', label='Train F1', linewidth=2, markersize=6)\n    axes[2].plot(epochs, history['val_f1'], 'r-s', label='Val F1', linewidth=2, markersize=6)\n    axes[2].set_title('Training and Validation F1 Score (Macro)', fontsize=14, fontweight='bold')\n    axes[2].set_xlabel('Epoch', fontsize=12)\n    axes[2].set_ylabel('F1 Score', fontsize=12)\n    axes[2].legend(fontsize=11)\n    axes[2].grid(True, alpha=0.3, linestyle='--')\n\n    plt.tight_layout()\n\n    if save_path:\n        plt.savefig(save_path, dpi=300, bbox_inches='tight')\n        print(f\"Saved training plot to {save_path}\")\n\n    plt.show()\n\n    # Print summary statistics\n    print(\"\\n\" + \"=\"*60)\n    print(\"TRAINING SUMMARY\")\n    print(\"=\"*60)\n    print(f\"Final Train Loss: {history['train_loss'][-1]:.4f} | Final Train Acc: {history['train_acc'][-1]:.4f} | Final Train F1: {history['train_f1'][-1]:.4f}\")\n    print(f\"Final Val Loss:   {history['val_loss'][-1]:.4f} | Final Val Acc:   {history['val_acc'][-1]:.4f} | Final Val F1:   {history['val_f1'][-1]:.4f}\")\n    print(f\"Best Val Loss:    {min(history['val_loss']):.4f} (epoch {history['val_loss'].index(min(history['val_loss']))+1})\")\n    print(f\"Best Val Acc:     {max(history['val_acc']):.4f} (epoch {history['val_acc'].index(max(history['val_acc']))+1})\")\n    print(f\"Best Val F1:      {max(history['val_f1']):.4f} (epoch {history['val_f1'].index(max(history['val_f1']))+1})\")\n    print(\"=\"*60)\n\n# Plot the training history\nplot_training_history(history, save_path='training_curves.png')\n\n# Optional: Save history to JSON for later analysis\nwith open('training_history.json', 'w') as f:\n    # Convert tensors to lists for JSON serialization\n    history_serializable = {\n        key: [float(x) if torch.is_tensor(x) else x for x in values]\n        for key, values in history.items()\n    }\n    json.dump(history_serializable, f, indent=2)\nprint(\"\\nSaved training history to training_history.json\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ==========================================\n## PHASE 4: INITIAL MODEL EVALUATION (BASELINE)\n## ==========================================\n## This evaluates the model with DEFAULT 0.5 thresholds (before optimization)\n## Run this FIRST to establish baseline performance","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\nimport numpy as np\nimport torch\nimport warnings\nfrom tqdm import tqdm\nfrom sklearn.metrics import (\n    classification_report,\n    accuracy_score,\n    multilabel_confusion_matrix,\n    f1_score\n)\n\nwarnings.filterwarnings('ignore')\n\nprint(\"=\"*80)\nprint(\"FINAL MODEL EVALUATION ON TEST SET\")\nprint(\"=\"*80)\nprint(\"\\nNOTE: Test set is 15% of original dataset (70/15/15 train/val/test split)\")\nprint(\"=\"*80)\n\n# --------------------------------------------------\n# 1. LOAD BEST MODEL\n# --------------------------------------------------\nif os.path.exists(\"best_model.pth\"):\n    model.load_state_dict(torch.load(\"best_model.pth\", map_location=config.DEVICE))\n    print(\"\\n✓ Loaded best model weights from best_model.pth\")\nelse:\n    print(\"\\n⚠ Warning: best_model.pth not found. Using current model weights.\")\n\nmodel.eval()\n\n# --------------------------------------------------\n# 2. COLLECT PREDICTIONS\n# --------------------------------------------------\ny_true = []\ny_pred = []\ny_probs = []  # Store probabilities for analysis\n\nTHRESHOLD = 0.5\n\nwith torch.no_grad():\n    for images, labels in tqdm(test_loader, desc=\"Evaluating on test set\"):\n        images = images.to(config.DEVICE)\n        labels = labels.to(config.DEVICE)\n\n        outputs = model(images)\n        probs = torch.sigmoid(outputs)\n        preds = (probs > THRESHOLD).int()\n\n        y_true.append(labels.cpu().numpy())\n        y_pred.append(preds.cpu().numpy())\n        y_probs.append(probs.cpu().numpy())\n\ny_true = np.vstack(y_true)\ny_pred = np.vstack(y_pred)\ny_probs = np.vstack(y_probs)\n\n# --------------------------------------------------\n# 3. CALCULATE COMPREHENSIVE METRICS\n# --------------------------------------------------\n# Exact match accuracy (all labels must match)\nexact_match_acc = accuracy_score(y_true, y_pred)\n\n# Per-label accuracy\nper_label_acc = []\nfor i in range(y_true.shape[1]):\n    acc = accuracy_score(y_true[:, i], y_pred[:, i])\n    per_label_acc.append(acc)\n\n# Macro and micro F1 scores\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\" + \"=\"*80)\nprint(\"TEST SET PERFORMANCE SUMMARY\")\nprint(\"=\"*80)\nprint(f\"Test Set Size:          {len(y_true)} images\")\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(\"=\"*80)\n\n# --------------------------------------------------\n# 4. PER-LABEL PERFORMANCE\n# --------------------------------------------------\nprint(\"\\nPer-Label Accuracy:\")\nprint(\"-\" * 40)\nfor label, acc in zip(config.LABELS, per_label_acc):\n    print(f\"{label:25s}: {acc*100:6.2f}%\")\n\n# --------------------------------------------------\n# 5. DETAILED CLASSIFICATION REPORT\n# --------------------------------------------------\nprint(\"\\n\" + \"=\"*80)\nprint(\"DETAILED CLASSIFICATION REPORT\")\nprint(\"=\"*80)\nprint(classification_report(\n    y_true,\n    y_pred,\n    target_names=config.LABELS,\n    zero_division=0\n))\n\n# --------------------------------------------------\n# 6. MULTI-LABEL CONFUSION MATRICES\n# --------------------------------------------------\nprint(\"\\n\" + \"=\"*80)\nprint(\"CONFUSION MATRICES (Per Label)\")\nprint(\"=\"*80)\n\nmcm = multilabel_confusion_matrix(y_true, y_pred)\n\nfor idx, label in enumerate(config.LABELS):\n    tn, fp, fn, tp = mcm[idx].ravel()\n\n    # Calculate metrics for this label\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    plt.figure(figsize=(5, 4))\n    sns.heatmap(\n        [[tp, fp], [fn, tn]],\n        annot=True,\n        fmt='d',\n        cmap='Greens',\n        xticklabels=['Pred Positive', 'Pred Negative'],\n        yticklabels=['True Positive', 'True Negative'],\n        cbar_kws={'label': 'Count'}\n    )\n\n    plt.title(f'Confusion Matrix - {label}\\nPrecision: {precision:.3f} | Recall: {recall:.3f} | F1: {f1:.3f}',\n              fontsize=11, fontweight='bold')\n    plt.ylabel('True Label', fontsize=10)\n    plt.xlabel('Predicted Label', fontsize=10)\n    plt.tight_layout()\n    plt.savefig(f'confusion_matrix_{label}.png', dpi=150, bbox_inches='tight')\n    plt.show()\n\n# --------------------------------------------------\n# 7. PREDICTION CONFIDENCE ANALYSIS\n# --------------------------------------------------\nprint(\"\\n\" + \"=\"*80)\nprint(\"PREDICTION CONFIDENCE ANALYSIS\")\nprint(\"=\"*80)\n\nmean_confidence_correct = []\nmean_confidence_incorrect = []\n\nfor i in range(y_true.shape[1]):\n    # Correct predictions\n    correct_mask = (y_true[:, i] == y_pred[:, i]) & (y_true[:, i] == 1)\n    if correct_mask.sum() > 0:\n        mean_confidence_correct.append(y_probs[correct_mask, i].mean())\n    else:\n        mean_confidence_correct.append(0.0)\n\n    # Incorrect predictions (false positives)\n    fp_mask = (y_true[:, i] == 0) & (y_pred[:, i] == 1)\n    if fp_mask.sum() > 0:\n        mean_confidence_incorrect.append(y_probs[fp_mask, i].mean())\n    else:\n        mean_confidence_incorrect.append(0.0)\n\nplt.figure(figsize=(12, 6))\nx = np.arange(len(config.LABELS))\nwidth = 0.35\n\nplt.bar(x - width/2, mean_confidence_correct, width, label='Correct Predictions (True Pos)', alpha=0.8)\nplt.bar(x + width/2, mean_confidence_incorrect, width, label='False Positives', alpha=0.8)\n\nplt.xlabel('Labels', fontsize=12)\nplt.ylabel('Mean Confidence', fontsize=12)\nplt.title('Prediction Confidence Analysis', fontsize=14, fontweight='bold')\nplt.xticks(x, config.LABELS, rotation=45, ha='right')\nplt.legend()\nplt.grid(axis='y', alpha=0.3)\nplt.tight_layout()\nplt.savefig('confidence_analysis.png', dpi=150, bbox_inches='tight')\nplt.show()\n\nprint(\"\\n✓ Evaluation complete! All plots saved.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ==========================================\n## PHASE 5: THRESHOLD OPTIMIZATION\n## ==========================================\n## This finds optimal thresholds for each class to improve recall","metadata":{}},{"cell_type":"code","source":"# ====================================================================\n# NEW CELL: OPTIMAL THRESHOLD OPTIMIZATION\n# ====================================================================\n# This cell finds class-specific optimal thresholds to improve recall\n# for poorly performing classes (scab, frog_eye_leaf_spot)\n\nimport numpy as np\nimport torch\nfrom sklearn.metrics import f1_score\nfrom tqdm import tqdm\n\nprint(\"=\"*80)\nprint(\"THRESHOLD OPTIMIZATION - FINDING CLASS-SPECIFIC THRESHOLDS\")\nprint(\"=\"*80)\nprint(\"\\nProblem: Using 0.5 threshold for all classes is suboptimal.\")\nprint(\"Solution: Find optimal threshold for each class based on F1 score.\")\nprint(\"=\"*80)\n\n# Load best model\nmodel.load_state_dict(torch.load(\"best_model.pth\", map_location=config.DEVICE))\nmodel.eval()\n\n# --------------------------------------------------\n# STEP 1: Collect validation probabilities\n# --------------------------------------------------\nprint(\"\\n[Step 1/3] Collecting validation predictions...\")\n\nval_probs = []\nval_targets = []\n\nwith torch.no_grad():\n    for images, labels in tqdm(val_loader, desc=\"Collecting val predictions\"):\n        images = images.to(config.DEVICE)\n        outputs = model(images)\n        val_probs.append(torch.sigmoid(outputs).cpu())\n        val_targets.append(labels.cpu())\n\nval_probs = torch.cat(val_probs)\nval_targets = torch.cat(val_targets)\n\nprint(f\"✓ Collected {len(val_probs)} validation samples\")\n\n# --------------------------------------------------\n# STEP 2: Find optimal threshold for each class\n# --------------------------------------------------\nprint(\"\\n[Step 2/3] Finding optimal thresholds for each class...\")\n\noptimal_thresholds = []\nthreshold_range = np.arange(0.15, 0.85, 0.02)  # Test from 0.15 to 0.85\n\nprint(\"\\n\" + \"-\"*80)\nprint(f\"{'Class':<25} {'Opt Thresh':<10} {'Best F1':<10} {'Default F1':<10} {'Improvement':<12}\")\nprint(\"-\"*80)\n\nfor i, label in enumerate(config.LABELS):\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).float()\n        f1 = f1_score(val_targets[:, i], preds.numpy(), zero_division=0)\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 = (val_probs[:, i] > 0.5).float()\n    default_f1 = f1_score(val_targets[:, i], default_preds.numpy(), zero_division=0)\n    \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        print(f\"{label:<25} {best_thresh:<10.3f} {best_f1:<10.4f} {default_f1:<10.4f} +{improvement:>10.2f}% ⭐\")\n    elif improvement > 0:\n        print(f\"{label:<25} {best_thresh:<10.3f} {best_f1:<10.4f} {default_f1:<10.4f} +{improvement:>10.2f}%\")\n    else:\n        print(f\"{label:<25} {best_thresh:<10.3f} {best_f1:<10.4f} {default_f1:<10.4f} {improvement:>10.2f}%\")\n\nprint(\"-\"*80)\n\n# Convert to tensor\noptimal_thresholds = torch.tensor(optimal_thresholds)\n\nprint(f\"\\n✓ Optimal thresholds found!\")\nprint(f\"\\nDefault thresholds:     {torch.tensor([0.5, 0.5, 0.5, 0.5, 0.5, 0.5])}\")\nprint(f\"Optimal thresholds:     {optimal_thresholds}\")\n\n# --------------------------------------------------\n# STEP 3: Re-evaluate on TEST set with optimal thresholds\n# --------------------------------------------------\nprint(\"\\n[Step 3/3] Re-evaluating on TEST set with optimal thresholds...\")\n\n# Collect test predictions\ntest_probs = []\ntest_targets = []\n\nwith torch.no_grad():\n    for images, labels in tqdm(test_loader, desc=\"Collecting test predictions\"):\n        images = images.to(config.DEVICE)\n        outputs = model(images)\n        test_probs.append(torch.sigmoid(outputs).cpu())\n        test_targets.append(labels.cpu())\n\ntest_probs = torch.cat(test_probs)\ntest_targets = torch.cat(test_targets)\n\n# Predictions with default threshold (0.5)\npreds_default = (test_probs > 0.5).int()\n\n# Predictions with optimal thresholds\npreds_optimal = torch.zeros_like(test_probs)\nfor i in range(len(config.LABELS)):\n    preds_optimal[:, i] = (test_probs[:, i] > optimal_thresholds[i]).int()\n\n# Calculate metrics for both\nexact_match_default = (test_targets == preds_default).all(axis=1).float().mean().item()\nexact_match_optimal = (test_targets == preds_optimal).all(axis=1).float().mean().item()\n\n# Per-class F1 scores\nf1_default = []\nf1_optimal = []\nrecall_default = []\nrecall_optimal = []\n\nfor i in range(len(config.LABELS)):\n    f1_default.append(f1_score(test_targets[:, i], preds_default[:, i], zero_division=0))\n    f1_optimal.append(f1_score(test_targets[:, i], preds_optimal[:, i], zero_division=0))\n    \n    # Recall calculation\n    tp = ((test_targets[:, i] == 1) & (preds_default[:, i] == 1)).sum().item()\n    fn = ((test_targets[:, i] == 1) & (preds_default[:, i] == 0)).sum().item()\n    recall_default.append(tp / (tp + fn) if (tp + fn) > 0 else 0)\n    \n    tp = ((test_targets[:, i] == 1) & (preds_optimal[:, i] == 1)).sum().item()\n    fn = ((test_targets[:, i] == 1) & (preds_optimal[:, i] == 0)).sum().item()\n    recall_optimal.append(tp / (tp + fn) if (tp + fn) > 0 else 0)\n\n# --------------------------------------------------\n# COMPARISON RESULTS\n# --------------------------------------------------\nprint(\"\\n\" + \"=\"*80)\nprint(\"BEFORE vs AFTER: OPTIMAL THRESHOLDS\")\nprint(\"=\"*80)\n\nprint(f\"\\n{'Class':<25} {'Thresh':<8} {'Recall Bef':<11} {'Recall Aft':<11} {'F1 Bef':<9} {'F1 Aft':<9} {'Recall Δ':<10}\")\nprint(\"-\"*80)\n\nfor i, label in enumerate(config.LABELS):\n    recall_delta = (recall_optimal[i] - recall_default[i]) * 100\n    f1_delta = (f1_optimal[i] - f1_default[i]) * 100\n    \n    if abs(recall_delta) > 5:\n        marker = \" ⭐\" if recall_delta > 0 else \" ⚠\"\n    else:\n        marker = \"\"\n    \n    print(f\"{label:<25} {optimal_thresholds[i]:<8.3f} \"\n          f\"{recall_default[i]*100:<11.2f} {recall_optimal[i]*100:<11.2f} \"\n          f\"{f1_default[i]:<9.4f} {f1_optimal[i]:<9.4f} \"\n          f\"{recall_delta:>+6.2f}%{marker}\")\n\nprint(\"-\"*80)\n\n# Overall improvement\nprint(f\"\\nOverall Metrics:\")\nprint(f\"  Exact Match Accuracy (Before): {exact_match_default*100:.2f}%\")\nprint(f\"  Exact Match Accuracy (After):  {exact_match_optimal*100:.2f}%\")\nprint(f\"  Improvement:                    {(exact_match_optimal - exact_match_default)*100:+.2f}%\")\n\nmacro_f1_default = np.mean(f1_default)\nmacro_f1_optimal = np.mean(f1_optimal)\nprint(f\"  Macro F1 (Before):              {macro_f1_default:.4f}\")\nprint(f\"  Macro F1 (After):               {macro_f1_optimal:.4f}\")\nprint(f\"  Improvement:                    {macro_f1_optimal - macro_f1_default:+.4f}\")\n\n# --------------------------------------------------\n# VISUALIZATION: Threshold Comparison\n# --------------------------------------------------\nprint(\"\\n\" + \"=\"*80)\nprint(\"CREATING VISUALIZATION...\")\nprint(\"=\"*80)\n\nimport matplotlib.pyplot as plt\n\nfig, axes = plt.subplots(1, 2, figsize=(15, 6))\n\n# Plot 1: Recall Comparison\nx = np.arange(len(config.LABELS))\nwidth = 0.35\n\naxes[0].bar(x - width/2, [r*100 for r in recall_default], width, \n            label='Default (0.5)', alpha=0.8, color='coral')\naxes[0].bar(x + width/2, [r*100 for r in recall_optimal], width, \n            label='Optimal Thresholds', alpha=0.8, color='steelblue')\naxes[0].set_xlabel('Class', fontsize=12)\naxes[0].set_ylabel('Recall (%)', fontsize=12)\naxes[0].set_title('Recall: Default 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_default, recall_optimal)):\n    delta = (after - before) * 100\n    if abs(delta) > 3:\n        axes[0].annotate(f'{delta:+.1f}%', \n                        xy=(i + width/2, after*100), \n                        xytext=(i + width/2, after*100 + 5),\n                        ha='center', fontsize=9, fontweight='bold',\n                        color='green' if delta > 0 else 'red')\n\n# Plot 2: Optimal Thresholds per Class\ncolors = ['green' if t < 0.5 else 'orange' if t < 0.6 else 'red' for t in optimal_thresholds]\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 i, (bar, thresh) in enumerate(zip(bars, optimal_thresholds)):\n    height = bar.get_height()\n    axes[1].text(bar.get_x() + bar.get_width()/2., height,\n                f'{thresh:.3f}', ha='center', va='bottom', fontsize=9)\n\nplt.tight_layout()\nplt.savefig('threshold_optimization_comparison.png', dpi=150, bbox_inches='tight')\nplt.show()\n\n# Save thresholds to file\nimport json\nthreshold_dict = {label: float(thresh) for label, thresh in zip(config.LABELS, optimal_thresholds)}\nwith open('optimal_thresholds.json', 'w') as f:\n    json.dump(threshold_dict, f, indent=2)\n\nprint(\"\\n✓ Threshold optimization complete!\")\nprint(f\"✓ Saved optimal thresholds to 'optimal_thresholds.json'\")\nprint(f\"✓ Saved comparison plot to 'threshold_optimization_comparison.png'\")\nprint(\"\\n\" + \"=\"*80)\nprint(\"KEY FINDINGS:\")\nprint(\"=\"*80)\nprint(f\"• Classes that benefit most from threshold tuning:\")\nfor i, label in enumerate(config.LABELS):\n    delta = (recall_optimal[i] - recall_default[i]) * 100\n    if delta > 5:\n        print(f\"  - {label}: {delta:+.1f}% recall improvement (thresh: {optimal_thresholds[i]:.3f})\")\nprint(\"=\"*80)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ==========================================\n## PHASE 6: FINAL EVALUATION WITH OPTIMAL THRESHOLDS\n## ==========================================\n## This evaluates the model using the optimized thresholds","metadata":{}},{"cell_type":"code","source":"# ====================================================================\n# FINAL EVALUATION WITH OPTIMAL THRESHOLDS\n# ====================================================================\n# This cell performs comprehensive evaluation using the optimized thresholds\n# and creates a combined confusion matrix (dominant class style like plant-pathology.ipynb)\n\nimport numpy as np\nimport torch\nimport json\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import (\n    classification_report,\n    accuracy_score,\n    multilabel_confusion_matrix,\n    f1_score,\n    precision_score,\n    recall_score,\n    confusion_matrix\n)\nfrom tqdm import tqdm\n\nprint(\"=\"*80)\nprint(\"COMPREHENSIVE TEST SET EVALUATION WITH OPTIMAL THRESHOLDS\")\nprint(\"=\"*80)\n\n# --------------------------------------------------\n# 1. LOAD BEST MODEL AND OPTIMAL THRESHOLDS\n# --------------------------------------------------\nmodel.load_state_dict(torch.load(\"best_model.pth\", map_location=config.DEVICE))\nmodel.eval()\n\n# Load optimal thresholds from the threshold optimization step\nif os.path.exists(\"optimal_thresholds.json\"):\n    with open(\"optimal_thresholds.json\", \"r\") as f:\n        threshold_dict = json.load(f)\n    optimal_thresholds = torch.tensor([threshold_dict[label] for label in config.LABELS])\n    print(f\"\\n✓ Loaded optimal thresholds: {optimal_thresholds}\")\nelse:\n    print(\"\\n⚠ Warning: optimal_thresholds.json not found. Using default 0.5 thresholds\")\n    optimal_thresholds = torch.tensor([0.5] * len(config.LABELS))\n\n# --------------------------------------------------\n# 2. COLLECT TEST PREDICTIONS WITH OPTIMAL THRESHOLDS\n# --------------------------------------------------\ny_true = []\ny_pred = []\ny_probs = []\n\nprint(\"\\nCollecting test predictions with optimal thresholds...\")\n\nwith torch.no_grad():\n    for images, labels in tqdm(test_loader, desc=\"Predicting on test set\"):\n        images = images.to(config.DEVICE)\n        labels = labels.to(config.DEVICE)\n\n        outputs = model(images)\n        probs = torch.sigmoid(outputs)\n        \n        # Apply optimal thresholds per class\n        preds = torch.zeros_like(probs)\n        for i in range(len(config.LABELS)):\n            preds[:, i] = (probs[:, i] > optimal_thresholds[i]).float()\n\n        y_true.append(labels.cpu().numpy())\n        y_pred.append(preds.cpu().numpy())\n        y_probs.append(probs.cpu().numpy())\n\ny_true = np.vstack(y_true)\ny_pred = np.vstack(y_pred)\ny_probs = np.vstack(y_probs)\n\n# --------------------------------------------------\n# 3. CALCULATE COMPREHENSIVE METRICS\n# --------------------------------------------------\n# Exact Match Accuracy (all labels must match exactly)\nexact_match_acc = accuracy_score(y_true, y_pred)\n\n# Per-label metrics\nper_label_precision = []\nper_label_recall = []\nper_label_f1 = []\nper_label_accuracy = []\n\nfor i in range(len(config.LABELS)):\n    per_label_precision.append(precision_score(y_true[:, i], y_pred[:, i], zero_division=0))\n    per_label_recall.append(recall_score(y_true[:, i], y_pred[:, i], zero_division=0))\n    per_label_f1.append(f1_score(y_true[:, i], y_pred[:, i], zero_division=0))\n    per_label_accuracy.append(accuracy_score(y_true[:, i], y_pred[:, i]))\n\n# Macro and Micro averages\nmacro_precision = np.mean(per_label_precision)\nmacro_recall = np.mean(per_label_recall)\nmacro_f1 = np.mean(per_label_f1)\nmicro_f1 = f1_score(y_true, y_pred, average='micro', zero_division=0)\n\n# --------------------------------------------------\n# 4. PRINT SUMMARY TABLE\n# --------------------------------------------------\nprint(\"\\n\" + \"=\"*80)\nprint(\"PERFORMANCE SUMMARY WITH OPTIMAL THRESHOLDS\")\nprint(\"=\"*80)\n\nmetrics_df = pd.DataFrame({\n    'Class': config.LABELS,\n    'Optimal Thresh': optimal_thresholds.numpy(),\n    'Precision': [f\"{p:.4f}\" for p in per_label_precision],\n    'Recall': [f\"{r:.4f}\" for r in per_label_recall],\n    'F1 Score': [f\"{f:.4f}\" for f in per_label_f1],\n    'Accuracy': [f\"{a:.4f}\" for a in per_label_accuracy]\n})\n\nprint(metrics_df.to_string(index=False))\n\nprint(\"\\n\" + \"-\"*80)\nprint(\"OVERALL METRICS:\")\nprint(\"-\"*80)\nprint(f\"Exact Match Accuracy:  {exact_match_acc*100:6.2f}%\")\nprint(f\"Macro Precision:      {macro_precision:6.4f}\")\nprint(f\"Macro Recall:         {macro_recall:6.4f}\")\nprint(f\"Macro F1:             {macro_f1:6.4f}\")\nprint(f\"Micro F1:             {micro_f1:6.4f}\")\nprint(\"=\"*80)\n\n# --------------------------------------------------\n# 5. COMBINED CONFUSION MATRIX (DOMINANT CLASS STYLE)\n# --------------------------------------------------\n# For multi-label, we convert to single-label by taking the dominant class\n# (the class with highest probability for each sample)\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"COMBINED CONFUSION MATRIX (DOMINANT CLASS)\")\nprint(\"=\"*80)\nprint(\"\\nConverting multi-label predictions to dominant class (argmax)...\")\n\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 combined confusion matrix\ncm_combined = confusion_matrix(true_dominant, pred_dominant)\n\n# Normalize by row (true labels) for better visualization\ncm_normalized = cm_combined.astype('float') / cm_combined.sum(axis=1)[:, np.newaxis]\n\n# Plot combined confusion matrix\nfig, axes = plt.subplots(1, 2, figsize=(18, 8))\n\n# Raw counts\nsns.heatmap(cm_combined, annot=True, fmt='d', cmap='Blues', \n            xticklabels=config.LABELS, yticklabels=config.LABELS,\n            cbar_kws={'label': 'Count'}, ax=axes[0])\naxes[0].set_title('Confusion Matrix (Raw Counts)', fontsize=14, fontweight='bold')\naxes[0].set_xlabel('Predicted Class', fontsize=12)\naxes[0].set_ylabel('True Class', fontsize=12)\n\n# Normalized\nsns.heatmap(cm_normalized, annot=True, fmt='.2%', cmap='Blues', \n            xticklabels=config.LABELS, yticklabels=config.LABELS,\n            cbar_kws={'label': 'Proportion'}, ax=axes[1])\naxes[1].set_title('Confusion Matrix (Normalized by Row)', fontsize=14, fontweight='bold')\naxes[1].set_xlabel('Predicted Class', fontsize=12)\naxes[1].set_ylabel('True Class', fontsize=12)\n\nplt.tight_layout()\nplt.savefig('combined_confusion_matrix_dominant_class.png', dpi=150, bbox_inches='tight')\nplt.show()\n\n# --------------------------------------------------\n# 6. PER-CLASS CONFUSION MATRICES (MULTI-LABEL)\n# --------------------------------------------------\nprint(\"\\n\" + \"=\"*80)\nprint(\"PER-CLASS MULTI-LABEL CONFUSION MATRICES\")\nprint(\"=\"*80)\n\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    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(f'{label}\\nPrec: {precision:.3f} | Rec: {recall:.3f} | F1: {f1:.3f}',\n                       fontsize=11, fontweight='bold')\n\nplt.tight_layout()\nplt.savefig('per_class_confusion_matrices.png', dpi=150, bbox_inches='tight')\nplt.show()\n\n# --------------------------------------------------\n# 7. DETAILED CLASSIFICATION REPORT\n# --------------------------------------------------\nprint(\"\\n\" + \"=\"*80)\nprint(\"DETAILED CLASSIFICATION REPORT\")\nprint(\"=\"*80)\nprint(classification_report(\n    y_true, y_pred,\n    target_names=config.LABELS,\n    zero_division=0\n))\n\n# --------------------------------------------------\n# 8. SAVE METRICS TO FILE\n# --------------------------------------------------\nresults_summary = {\n    'exact_match_accuracy': float(exact_match_acc),\n    'macro_precision': float(macro_precision),\n    'macro_recall': float(macro_recall),\n    'macro_f1': float(macro_f1),\n    'micro_f1': float(micro_f1),\n    'per_class_metrics': {\n        label: {\n            'optimal_threshold': float(optimal_thresholds[i]),\n            'precision': float(per_label_precision[i]),\n            'recall': float(per_label_recall[i]),\n            'f1': float(per_label_f1[i]),\n            'accuracy': float(per_label_accuracy[i])\n        }\n        for i, label in enumerate(config.LABELS)\n    }\n}\n\nwith open('test_evaluation_with_optimal_thresholds.json', 'w') as f:\n    json.dump(results_summary, f, indent=2)\n\nprint(\"\\n✓ Evaluation complete!\")\nprint(\"✓ Saved results to 'test_evaluation_with_optimal_thresholds.json'\")\nprint(\"✓ Saved combined confusion matrix to 'combined_confusion_matrix_dominant_class.png'\")\nprint(\"✓ Saved per-class confusion matrices to 'per_class_confusion_matrices.png'\")\nprint(\"=\"*80)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}