{"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":"none","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":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\n!pip install imutils","metadata":{"trusted":true},"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},"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},"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},"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"(train_df['labels'].value_counts()).plot(kind='bar')","metadata":{"trusted":true},"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_labels.value_counts().plot(kind='bar')","metadata":{"trusted":true},"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"examine_images(train_df, num_images=20)","metadata":{"trusted":true},"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},"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},"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},"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},"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},"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},"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_weights","metadata":{"trusted":true},"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# Dataloaders (Keep your existing code, but ensure shuffle=True for train)\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=2)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device=torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\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)\nmodel = ResNetScratch(num_classes=config.NUM_CLASSES).to(device)\n","metadata":{"trusted":true},"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# Scheduler for \"Fine Tuning\" aspect (see Phase 4)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=3, verbose=True)\n\ngrad_log = []  \n\ndef train_one_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct_preds = 0\n    total_preds = 0\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        # Log gradient norm \n        total_grad = sum(p.grad.norm().item() for p in model.parameters() if p.grad is not None)\n        grad_log.append(total_grad)\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        loop.set_description(f\"Loss: {loss.item():.4f} | GradNorm: {total_grad:.3f}\")\n\n    avg_loss = running_loss / len(loader)\n    avg_acc = correct_preds / total_preds\n    return avg_loss, avg_acc\n\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct_preds = 0\n    total_preds = 0\n    \n    with torch.no_grad():\n        for images, labels in loader:\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    return running_loss / len(loader), correct_preds / total_preds","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_epochs = 10\nfor epoch in range(num_epochs):\n    train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, device)\n    val_loss, val_acc = validate(model, val_loader, criterion, device)\n\n    scheduler.step(val_loss)\n\n    print(f\"\\nEpoch [{epoch+1}/{num_epochs}] \"\n          f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} \"\n          f\"Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.4f}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"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)\n\nwarnings.filterwarnings('ignore')\n\nprint(\"=== RUNNING MODEL EVALUATION ON TEST SET ===\")\n\n# --------------------------------------------------\n# 1. LABEL DEFINITIONS\n# --------------------------------------------------\nLABELS = config.LABELS\nNUM_CLASSES = len(LABELS)\n\n# --------------------------------------------------\n# 2. LOAD BEST MODEL (IF EXISTS)\n# --------------------------------------------------\nif os.path.exists(\"best_model.pth\"):\n    model.load_state_dict(torch.load(\"best_model.pth\", map_location=config.DEVICE))\n    print(\"-> Loaded best model weights.\")\nelse:\n    print(\"-> Warning: Using current model weights.\")\n\nmodel.eval()\n\n# --------------------------------------------------\n# 3. COLLECT PREDICTIONS\n# --------------------------------------------------\ny_true = []\ny_pred = []\n\nTHRESHOLD = 0.5\n\nwith torch.no_grad():\n    for images, labels in tqdm(test_loader, desc=\"Predicting\"):\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\ny_true = np.vstack(y_true)\ny_pred = np.vstack(y_pred)\n\n# --------------------------------------------------\n# 4. EXACT MATCH ACCURACY (STRICT METRIC)\n# --------------------------------------------------\nexact_match_acc = accuracy_score(y_true, y_pred)\n\nprint(\"\\n\" + \"#\" * 60)\nprint(f\"### EXACT MATCH ACCURACY (TEST): {exact_match_acc * 100:.2f}% ###\")\nprint(\"#\" * 60 + \"\\n\")\n\n# --------------------------------------------------\n# 5. CLASSIFICATION REPORT\n# --------------------------------------------------\nprint(\"DETAILED CLASSIFICATION REPORT:\")\nprint(\n    classification_report(\n        y_true,\n        y_pred,\n        target_names=LABELS,\n        zero_division=0\n    )\n)\n\n# --------------------------------------------------\n# 6. MULTI-LABEL CONFUSION MATRICES\n# --------------------------------------------------\nmcm = multilabel_confusion_matrix(y_true, y_pred)\n\nfor idx, label in enumerate(LABELS):\n    tn, fp, fn, tp = mcm[idx].ravel()\n\n    plt.figure(figsize=(4, 3))\n    sns.heatmap(\n        [[tp, fp], [fn, tn]],\n        annot=True,\n        fmt='d',\n        cmap='Greens',\n        xticklabels=['Predicted 1', 'Predicted 0'],\n        yticklabels=['Actual 1', 'Actual 0']\n    )\n\n    plt.title(f'Confusion Matrix - {label}')\n    plt.ylabel('Actual')\n    plt.xlabel('Predicted')\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}