{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":25563,"databundleVersionId":2094376,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":121906917,"sourceType":"kernelVersion"},{"sourceId":670827,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":508097,"modelId":522766}],"dockerImageVersionId":30408,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%capture\n!pip install imutils","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T07:47:33.840855Z","iopub.execute_input":"2025-12-03T07:47:33.841981Z","iopub.status.idle":"2025-12-03T07:47:42.536782Z","shell.execute_reply.started":"2025-12-03T07:47:33.841938Z","shell.execute_reply":"2025-12-03T07:47:42.535581Z"}},"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\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":{"execution":{"iopub.status.busy":"2025-12-03T07:47:42.539196Z","iopub.execute_input":"2025-12-03T07:47:42.539611Z","iopub.status.idle":"2025-12-03T07:47:42.547728Z","shell.execute_reply.started":"2025-12-03T07:47:42.539576Z","shell.execute_reply":"2025-12-03T07:47:42.546518Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class config:\n    # specify the paths to datasets\n    DATA_DIR = Path('../input/plant-pathology-2021-fgvc8/train_images')\n    ROOT_DIR = Path('./data')\n    TRAIN_DIR = ROOT_DIR.joinpath('train')\n    TEST_DIR = ROOT_DIR.joinpath('test')\n    VAL_DIR = ROOT_DIR.joinpath('val')\n\n    # set the input height and width\n    INPUT_HEIGHT = 224\n    INPUT_WIDTH = 224\n\n    # set the input heig/ht and width\n    IMAGENET_MEAN = [0.485, 0.456, 0.406]\n    IMAGENET_STD = [0.229, 0.224, 0.225]\n    \n    IMAGE_TYPE = '.jpg'\n    BATCH_SIZE = 32\n    # will use the vision transformer\n    MODEL_NAME = 'vit_base'\n    \n    DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n    TRAINING_PARAMS = 'training_hyperparams/default_train_params'\n    LABELS = ['complex', 'frog_eye_leaf_spot', 'healthy', 'powdery_mildew', 'rust', 'scab']\n    NUM_CLASSES = len(LABELS)\n    CHECKPOINT_DIR = 'checkpoints'\n","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:47:42.548850Z","iopub.execute_input":"2025-12-03T07:47:42.549122Z","iopub.status.idle":"2025-12-03T07:47:42.559882Z","shell.execute_reply.started":"2025-12-03T07:47:42.549096Z","shell.execute_reply":"2025-12-03T07:47:42.558870Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Resize images and save to disk\n+ Use https://www.kaggle.com/code/harpdeci/multi-label-classification-plant-pathology/output?select=__notebook_source__.ipynb to import the output resized image to use when testing for faster result\n+ If you want to run resized code yourself => Uncomment below code","metadata":{}},{"cell_type":"code","source":"!mkdir data","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:47:42.562163Z","iopub.execute_input":"2025-12-03T07:47:42.562734Z","iopub.status.idle":"2025-12-03T07:47:43.618683Z","shell.execute_reply.started":"2025-12-03T07:47:42.562701Z","shell.execute_reply":"2025-12-03T07:47:43.617624Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Set the desired output dimensions\n# output_size = (224, 224)\n\n# # Get a list of all image file paths in the input directory\n# image_paths = list(paths.list_images(config.DATA_DIR))\n\n# # Create a progress bar object\n# progress_bar = tqdm(total=len(image_paths), desc='Resizing images')\n\n# # Loop over all image file paths\n# for image_path in image_paths:\n#     # Load the image with PIL\n#     image_path=Path(image_path)\n#     image = Image.open(image_path)\n\n#     # Resize the image\n#     resized_image = image.resize(output_size)\n\n#     # Get the output file path\n#     output_path = config.ROOT_DIR / image_path.name\n\n#     # Save the resized image to disk\n#     resized_image.save(output_path)\n#     # Update the progress bar\n#     progress_bar.update(1)\n    \n# # Close the progress bar\n# progress_bar.close()","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:47:43.620130Z","iopub.execute_input":"2025-12-03T07:47:43.620398Z","iopub.status.idle":"2025-12-03T07:47:43.625495Z","shell.execute_reply.started":"2025-12-03T07:47:43.620371Z","shell.execute_reply":"2025-12-03T07:47:43.624533Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import shutil\n# import os\n\n# # Name of the output file\n# output_filename = \"my_kaggle_data\"\n\n# # Directory to zip (usually '/kaggle/working')\n# directory_to_zip = \"/kaggle/working\"\n\n# # Create the zip file\n# shutil.make_archive(output_filename, 'zip', directory_to_zip)\n\n# print(f\"Created {output_filename}.zip\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T07:47:43.627067Z","iopub.execute_input":"2025-12-03T07:47:43.627689Z","iopub.status.idle":"2025-12-03T07:47:43.635363Z","shell.execute_reply.started":"2025-12-03T07:47:43.627646Z","shell.execute_reply":"2025-12-03T07:47:43.634384Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# from IPython.display import FileLink\n\n# # This will create a clickable link in the output area\n# FileLink(r'my_kaggle_data.zip')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T07:47:43.636770Z","iopub.execute_input":"2025-12-03T07:47:43.637412Z","iopub.status.idle":"2025-12-03T07:47:43.648318Z","shell.execute_reply.started":"2025-12-03T07:47:43.637372Z","shell.execute_reply":"2025-12-03T07:47:43.647427Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def 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: './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":{"execution":{"iopub.status.busy":"2025-12-03T07:47:43.649406Z","iopub.execute_input":"2025-12-03T07:47:43.649682Z","iopub.status.idle":"2025-12-03T07:47:43.663089Z","shell.execute_reply.started":"2025-12-03T07:47:43.649659Z","shell.execute_reply":"2025-12-03T07:47:43.662257Z"},"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":{"execution":{"iopub.status.busy":"2025-12-03T07:47:43.664326Z","iopub.execute_input":"2025-12-03T07:47:43.664866Z","iopub.status.idle":"2025-12-03T07:47:43.707426Z","shell.execute_reply.started":"2025-12-03T07:47:43.664839Z","shell.execute_reply":"2025-12-03T07:47:43.706641Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"(train_df['labels'].value_counts()).plot(kind='bar')","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:47:43.710406Z","iopub.execute_input":"2025-12-03T07:47:43.710694Z","iopub.status.idle":"2025-12-03T07:47:43.990716Z","shell.execute_reply.started":"2025-12-03T07:47:43.710667Z","shell.execute_reply":"2025-12-03T07:47:43.989725Z"},"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":{"execution":{"iopub.status.busy":"2025-12-03T07:47:43.992439Z","iopub.execute_input":"2025-12-03T07:47:43.992828Z","iopub.status.idle":"2025-12-03T07:47:44.017862Z","shell.execute_reply.started":"2025-12-03T07:47:43.992790Z","shell.execute_reply":"2025-12-03T07:47:44.017077Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_labels.value_counts().plot(kind='bar')","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:47:44.019378Z","iopub.execute_input":"2025-12-03T07:47:44.020076Z","iopub.status.idle":"2025-12-03T07:47:44.180133Z","shell.execute_reply.started":"2025-12-03T07:47:44.020033Z","shell.execute_reply":"2025-12-03T07:47:44.179154Z"},"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":"def 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":{"execution":{"iopub.status.busy":"2025-12-03T07:47:44.181370Z","iopub.execute_input":"2025-12-03T07:47:44.181770Z","iopub.status.idle":"2025-12-03T07:47:44.190081Z","shell.execute_reply.started":"2025-12-03T07:47:44.181731Z","shell.execute_reply":"2025-12-03T07:47:44.188890Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"examine_images(train_df, num_images=20)","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:47:44.191397Z","iopub.execute_input":"2025-12-03T07:47:44.191751Z","iopub.status.idle":"2025-12-03T07:47:50.443254Z","shell.execute_reply.started":"2025-12-03T07:47:44.191721Z","shell.execute_reply":"2025-12-03T07:47:50.441370Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 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":{"execution":{"iopub.status.busy":"2025-12-03T07:47:50.444578Z","iopub.execute_input":"2025-12-03T07:47:50.444878Z","iopub.status.idle":"2025-12-03T07:47:50.456346Z","shell.execute_reply.started":"2025-12-03T07:47:50.444851Z","shell.execute_reply":"2025-12-03T07:47:50.455352Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def 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":{"execution":{"iopub.status.busy":"2025-12-03T07:47:50.457698Z","iopub.execute_input":"2025-12-03T07:47:50.458018Z","iopub.status.idle":"2025-12-03T07:47:51.105295Z","shell.execute_reply.started":"2025-12-03T07:47:50.457986Z","shell.execute_reply":"2025-12-03T07:47:51.104406Z"},"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":{"execution":{"iopub.status.busy":"2025-12-03T07:47:51.106917Z","iopub.execute_input":"2025-12-03T07:47:51.107237Z","iopub.status.idle":"2025-12-03T07:47:51.113255Z","shell.execute_reply.started":"2025-12-03T07:47:51.107210Z","shell.execute_reply":"2025-12-03T07:47:51.112219Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class 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":{"execution":{"iopub.status.busy":"2025-12-03T07:47:51.114603Z","iopub.execute_input":"2025-12-03T07:47:51.115011Z","iopub.status.idle":"2025-12-03T07:47:51.125129Z","shell.execute_reply.started":"2025-12-03T07:47:51.114979Z","shell.execute_reply":"2025-12-03T07:47:51.124244Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_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":{"execution":{"iopub.status.busy":"2025-12-03T07:47:51.126535Z","iopub.execute_input":"2025-12-03T07:47:51.126895Z","iopub.status.idle":"2025-12-03T07:47:51.136785Z","shell.execute_reply.started":"2025-12-03T07:47:51.126849Z","shell.execute_reply":"2025-12-03T07:47:51.135818Z"},"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":{"execution":{"iopub.status.busy":"2025-12-03T07:47:51.138145Z","iopub.execute_input":"2025-12-03T07:47:51.138411Z","iopub.status.idle":"2025-12-03T07:47:51.153528Z","shell.execute_reply.started":"2025-12-03T07:47:51.138386Z","shell.execute_reply":"2025-12-03T07:47:51.152662Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 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":{"execution":{"iopub.status.busy":"2025-12-03T07:47:51.154507Z","iopub.execute_input":"2025-12-03T07:47:51.154755Z","iopub.status.idle":"2025-12-03T07:47:51.162813Z","shell.execute_reply.started":"2025-12-03T07:47:51.154732Z","shell.execute_reply":"2025-12-03T07:47:51.161895Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_weights","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:47:51.164146Z","iopub.execute_input":"2025-12-03T07:47:51.164972Z","iopub.status.idle":"2025-12-03T07:47:51.173348Z","shell.execute_reply.started":"2025-12-03T07:47:51.164943Z","shell.execute_reply":"2025-12-03T07:47:51.172501Z"},"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,"execution":{"iopub.status.busy":"2025-12-03T07:47:51.174549Z","iopub.execute_input":"2025-12-03T07:47:51.174901Z","iopub.status.idle":"2025-12-03T07:47:51.184505Z","shell.execute_reply.started":"2025-12-03T07:47:51.174862Z","shell.execute_reply":"2025-12-03T07:47:51.183595Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\nimport torch.nn.functional as F\n\nclass SimpleResBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, stride=1):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False)\n        self.bn1 = nn.BatchNorm2d(out_channels)\n        self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False)\n        self.bn2 = nn.BatchNorm2d(out_channels)\n        \n        self.shortcut = nn.Sequential()\n        if stride != 1 or in_channels != out_channels:\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False),\n                nn.BatchNorm2d(out_channels)\n            )\n\n    def forward(self, x):\n        out = F.relu(self.bn1(self.conv1(x)))\n        out = self.bn2(self.conv2(out))\n        out += self.shortcut(x)\n        out = F.relu(out)\n        return out\n\nclass CustomPlantCNN(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n        # Initial convolution\n        self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        self.bn1 = nn.BatchNorm2d(64)\n        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n        \n        # ResBlocks (Feature Extraction)\n        self.layer1 = SimpleResBlock(64, 64, stride=1)\n        self.layer2 = SimpleResBlock(64, 128, stride=2)\n        self.layer3 = SimpleResBlock(128, 256, stride=2)\n        self.layer4 = SimpleResBlock(256, 512, stride=2)\n        \n        # Classification Head\n        self.avgpool = nn.AdaptiveAvgPool2d((1, 1))\n        # Dropout helps reduce overfitting\n        self.dropout = nn.Dropout(p=0.5) \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.maxpool(x)\n        \n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n        \n        x = self.avgpool(x)\n        x = torch.flatten(x, 1)\n        x = self.dropout(x)\n        x = self.fc(x)\n        return x\n\n# Initialize Model\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel = CustomPlantCNN(num_classes=config.NUM_CLASSES).to(device)\nprint(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T07:47:51.186008Z","iopub.execute_input":"2025-12-03T07:47:51.186344Z","iopub.status.idle":"2025-12-03T07:47:51.244435Z","shell.execute_reply.started":"2025-12-03T07:47:51.186310Z","shell.execute_reply":"2025-12-03T07:47:51.243431Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 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\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 pass\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        \n        # Backward pass\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        # Statistics\n        running_loss += loss.item()\n        \n        # Calculate simplistic accuracy for progress bar (threshold 0.5)\n        preds = (torch.sigmoid(outputs) > 0.5).float()\n        correct_preds += (preds == labels).float().sum()\n        total_preds += labels.numel() # Total number of individual label predictions\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    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,"execution":{"iopub.status.busy":"2025-12-03T07:47:51.245974Z","iopub.execute_input":"2025-12-03T07:47:51.246444Z","iopub.status.idle":"2025-12-03T07:47:51.258525Z","shell.execute_reply.started":"2025-12-03T07:47:51.246402Z","shell.execute_reply":"2025-12-03T07:47:51.257402Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# num_epochs = 20\n# history = {'train_loss': [], 'val_loss': [], 'train_acc': [], 'val_acc': []}\n# best_val_loss = float('inf')\n\n# for epoch in range(num_epochs):\n#     print(f\"Epoch {epoch+1}/{num_epochs}\")\n    \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#     # Store history for monitoring\n#     history['train_loss'].append(train_loss)\n#     history['val_loss'].append(val_loss)\n#     history['train_acc'].append(train_acc.item())\n#     history['val_acc'].append(val_acc.item())\n    \n#     # \"Fine-tune\" aspect: Learning Rate Scheduling\n#     scheduler.step(val_loss)\n    \n#     print(f\"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.4f}\")\n    \n#     # Save best model\n#     if val_loss < best_val_loss:\n#         best_val_loss = val_loss\n#         torch.save(model.state_dict(), 'best_model_scratch.pth')\n#         print(\"Saved Best Model!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T07:47:51.259614Z","iopub.execute_input":"2025-12-03T07:47:51.259836Z","iopub.status.idle":"2025-12-03T07:47:51.270260Z","shell.execute_reply.started":"2025-12-03T07:47:51.259814Z","shell.execute_reply":"2025-12-03T07:47:51.269504Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# plt.figure(figsize=(12, 5))\n\n# # Plot Loss\n# plt.subplot(1, 2, 1)\n# plt.plot(history['train_loss'], label='Train Loss')\n# plt.plot(history['val_loss'], label='Validation Loss')\n# plt.title('Training Process: Loss')\n# plt.legend()\n\n# # Plot Accuracy\n# plt.subplot(1, 2, 2)\n# plt.plot(history['train_acc'], label='Train Acc')\n# plt.plot(history['val_acc'], label='Validation Acc')\n# plt.title('Training Process: Accuracy')\n# plt.legend()\n\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T07:47:51.271214Z","iopub.execute_input":"2025-12-03T07:47:51.271433Z","iopub.status.idle":"2025-12-03T07:47:51.282261Z","shell.execute_reply.started":"2025-12-03T07:47:51.271411Z","shell.execute_reply":"2025-12-03T07:47:51.281346Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Some tricks to fine-tune the model","metadata":{}},{"cell_type":"code","source":"# --- 1. MixUp Helpers ---\ndef mixup_data(x, y, alpha=0.4):\n    if alpha > 0:\n        lam = np.random.beta(alpha, alpha)\n    else:\n        lam = 1\n    batch_size = x.size()[0]\n    index = torch.randperm(batch_size).to(x.device)\n    mixed_x = lam * x + (1 - lam) * x[index, :]\n    y_a, y_b = y, y[index]\n    return mixed_x, y_a, y_b, lam\n\ndef mixup_criterion(criterion, pred, y_a, y_b, lam):\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n\n# --- 2. Threshold Calibration ---\ndef find_optimal_thresholds(model, val_loader, device):\n    model.eval()\n    val_probs = []\n    val_targets = []\n    \n    with torch.no_grad():\n        for images, labels in val_loader:\n            images = images.to(device)\n            outputs = model(images)\n            val_probs.append(torch.sigmoid(outputs).cpu())\n            val_targets.append(labels.cpu())\n            \n    val_probs = torch.cat(val_probs)\n    val_targets = torch.cat(val_targets)\n    best_thresholds = []\n    threshold_range = torch.arange(0.1, 0.95, 0.05)\n    \n    for i in range(val_probs.shape[1]):\n        best_f1 = 0\n        best_thresh = 0.5\n        for thresh in threshold_range:\n            preds = (val_probs[:, i] > thresh).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.item()\n        best_thresholds.append(best_thresh)\n        \n    return torch.tensor(best_thresholds).to(device)\n\n# --- 3. Robust Evaluation Function ---\ndef evaluate_model(model, loader, criterion, device, class_names, thresholds=None, use_tta=False, verbose=True):\n    model.eval()\n    if thresholds is None:\n        thresholds = torch.tensor([0.5] * len(class_names)).to(device)\n    else:\n        thresholds = thresholds.to(device)\n        \n    all_targets = []\n    all_preds = []\n    running_loss = 0.0\n    \n    if verbose:\n        print(f\"Evaluating... TTA={'ON' if use_tta else 'OFF'}\")\n    \n    with torch.no_grad():\n        for images, labels in loader:\n            images, labels = images.to(device), labels.to(device)\n            \n            logits_orig = model(images)\n            \n            if use_tta:\n                # 3-Way TTA\n                logits_h = model(torch.flip(images, dims=[3]))\n                logits_v = model(torch.flip(images, dims=[2]))\n                probs = (torch.sigmoid(logits_orig) + torch.sigmoid(logits_h) + torch.sigmoid(logits_v)) / 3.0\n            else:\n                probs = torch.sigmoid(logits_orig)\n                \n            loss = criterion(logits_orig, labels)\n            running_loss += loss.item()\n            \n            preds = (probs > thresholds).float()\n            all_targets.append(labels.cpu().numpy())\n            all_preds.append(preds.cpu().numpy())\n\n    all_targets = np.vstack(all_targets)\n    all_preds = np.vstack(all_preds)\n    \n    avg_loss = running_loss / len(loader)\n    macro_f1 = f1_score(all_targets, all_preds, average='macro', zero_division=0)\n    exact_acc = (all_targets == all_preds).all(axis=1).mean()\n    \n    return avg_loss, exact_acc, macro_f1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T07:47:51.287417Z","iopub.execute_input":"2025-12-03T07:47:51.287959Z","iopub.status.idle":"2025-12-03T07:47:51.302376Z","shell.execute_reply.started":"2025-12-03T07:47:51.287932Z","shell.execute_reply":"2025-12-03T07:47:51.301511Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Configuration ---\nEPOCHS = 30 # Adjust as needed (30 recommended)\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ncriterion = nn.BCEWithLogitsLoss(pos_weight=class_weights.to(device))\n\n# ==========================================\n# 1. TRAIN BASELINE MODEL\n# ==========================================\nprint(\"--- Starting Baseline Training ---\")\nbaseline_model = CustomPlantCNN(num_classes=config.NUM_CLASSES).to(device)\noptimizer = torch.optim.AdamW(baseline_model.parameters(), lr=1e-3, weight_decay=1e-4)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=3)\n\nbest_baseline_f1 = 0.0\n\n# --- METRIC LISTS ---\nbase_history = {\n    'train_loss': [], 'train_acc': [],\n    'val_loss': [], 'val_acc': [], 'val_f1': []\n}\n\nfor epoch in range(EPOCHS):\n    baseline_model.train()\n    running_loss = 0.0\n    running_acc = 0.0\n    total_samples = 0\n    \n    # Train Loop\n    for images, labels in tqdm(train_loader, desc=f\"Baseline Epoch {epoch+1}\", leave=False):\n        images, labels = images.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = baseline_model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        # Calculate Stats\n        batch_size = images.size(0)\n        running_loss += loss.item() * batch_size\n        preds = (torch.sigmoid(outputs) > 0.5).float()\n        acc = (preds == labels).float().mean().item()\n        running_acc += acc * batch_size\n        total_samples += batch_size\n\n    # Compute Epoch Averages\n    epoch_train_loss = running_loss / total_samples\n    epoch_train_acc = running_acc / total_samples\n    \n    # Store in history dict\n    base_history['train_loss'].append(epoch_train_loss)\n    base_history['train_acc'].append(epoch_train_acc)\n        \n    # Validation\n    val_loss, val_acc, val_f1 = evaluate_model(baseline_model, val_loader, criterion, device, config.LABELS, verbose=False)\n    \n    base_history['val_loss'].append(val_loss)\n    base_history['val_acc'].append(val_acc)\n    base_history['val_f1'].append(val_f1)\n    \n    scheduler.step(val_loss)\n    \n    print(f\"  Epoch {epoch+1}: Train Loss: {epoch_train_loss:.4f} | Val Loss: {val_loss:.4f} | Val F1: {val_f1:.4f}\")\n\n    # --- SAVE CHECKPOINT & HISTORY ---\n    if val_f1 > best_baseline_f1:\n        best_baseline_f1 = val_f1\n        \n        # 1. Save Model\n        torch.save(baseline_model.state_dict(), 'baseline_model.pth')\n        \n        # 2. Save History to JSON (Easier to read/download)\n        with open('baseline_history.json', 'w') as f:\n            json.dump(base_history, f)\n            \n        print(f\"  >> Saved New Best Baseline (F1: {val_f1:.4f}) and history.\")\n\n# ==========================================\n# 2. TRAIN ADVANCED MODEL\n# (MixUp + Cosine Annealing + COOLDOWN)\n# ==========================================\nprint(\"\\n--- Starting Advanced Training (MixUp + Cooldown) ---\")\nadvanced_model = CustomPlantCNN(num_classes=config.NUM_CLASSES).to(device)\noptimizer = torch.optim.AdamW(advanced_model.parameters(), lr=1e-3, weight_decay=1e-4)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-6)\n\nbest_advanced_f1 = 0.0\n\n# --- METRIC LISTS ---\nadv_history = {\n    'train_loss': [], 'train_acc': [],\n    'val_loss': [], 'val_acc': [], 'val_f1': []\n}\n\ncooldown_start = EPOCHS - 5\n\nfor epoch in range(EPOCHS):\n    advanced_model.train()\n    running_loss = 0.0\n    running_acc = 0.0\n    total_samples = 0\n    \n    use_mixup = True\n    if epoch >= cooldown_start:\n        use_mixup = False\n        \n    loop_desc = f\"Adv Epoch {epoch+1} {'(MixUp)' if use_mixup else '(Cooldown)'}\"\n    \n    for images, labels in tqdm(train_loader, desc=loop_desc, leave=False):\n        images, labels = images.to(device), labels.to(device)\n        optimizer.zero_grad()\n        batch_size = images.size(0)\n        \n        if use_mixup:\n            images, labels_a, labels_b, lam = mixup_data(images, labels, alpha=0.4)\n            outputs = advanced_model(images)\n            loss = mixup_criterion(criterion, outputs, labels_a, labels_b, lam)\n            \n            preds = (torch.sigmoid(outputs) > 0.5).float()\n            acc = lam * (preds == labels_a).float().mean().item() + (1 - lam) * (preds == labels_b).float().mean().item()\n        else:\n            outputs = advanced_model(images)\n            loss = criterion(outputs, labels)\n            preds = (torch.sigmoid(outputs) > 0.5).float()\n            acc = (preds == labels).float().mean().item()\n            \n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item() * batch_size\n        running_acc += acc * batch_size\n        total_samples += batch_size\n        \n    scheduler.step()\n    \n    epoch_train_loss = running_loss / total_samples\n    epoch_train_acc = running_acc / total_samples\n    \n    adv_history['train_loss'].append(epoch_train_loss)\n    adv_history['train_acc'].append(epoch_train_acc)\n    \n    # Validation\n    val_loss, val_acc, val_f1 = evaluate_model(advanced_model, val_loader, criterion, device, config.LABELS, verbose=False)\n    \n    adv_history['val_loss'].append(val_loss)\n    adv_history['val_acc'].append(val_acc)\n    adv_history['val_f1'].append(val_f1)\n    \n    print(f\"  Epoch {epoch+1}: Train Loss: {epoch_train_loss:.4f} | Val Loss: {val_loss:.4f} | Val F1: {val_f1:.4f}\")\n\n    # --- SAVE CHECKPOINT & HISTORY ---\n    if val_f1 > best_advanced_f1:\n        best_advanced_f1 = val_f1\n        \n        # 1. Save Model\n        torch.save(advanced_model.state_dict(), 'advanced_model.pth')\n        \n        # 2. Save History\n        with open('advanced_history.json', 'w') as f:\n            json.dump(adv_history, f)\n            \n        print(f\"  >> Saved New Best Advanced (F1: {val_f1:.4f}) and history.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T07:47:51.303938Z","iopub.execute_input":"2025-12-03T07:47:51.304232Z","iopub.status.idle":"2025-12-03T07:50:05.773562Z","shell.execute_reply.started":"2025-12-03T07:47:51.304207Z","shell.execute_reply":"2025-12-03T07:50:05.772049Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load the saved JSON files\nwith open('baseline_history.json', 'r') as f:\n    base_history = json.load(f)\n\nwith open('advanced_history.json', 'r') as f:\n    adv_history = json.load(f)\n\ndef plot_saved_metrics(history, title):\n    epochs = range(1, len(history['train_loss']) + 1)\n    \n    plt.figure(figsize=(12, 5))\n    \n    # Loss\n    plt.subplot(1, 2, 1)\n    plt.plot(epochs, history['train_loss'], label='Train Loss')\n    plt.plot(epochs, history['val_loss'], label='Val Loss')\n    plt.title(f'{title} - Loss')\n    plt.xlabel('Epochs')\n    plt.legend()\n    plt.grid(True, alpha=0.3)\n    \n    # Accuracy\n    plt.subplot(1, 2, 2)\n    plt.plot(epochs, history['train_acc'], label='Train Acc')\n    plt.plot(epochs, history['val_acc'], label='Val Acc')\n    plt.title(f'{title} - Accuracy')\n    plt.xlabel('Epochs')\n    plt.legend()\n    plt.grid(True, alpha=0.3)\n    \n    plt.show()\n\nplot_saved_metrics(base_history, \"Baseline Model\")\nplot_saved_metrics(adv_history, \"Advanced Model\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T07:50:15.473621Z","iopub.execute_input":"2025-12-03T07:50:15.474344Z","iopub.status.idle":"2025-12-03T07:50:15.495733Z","shell.execute_reply.started":"2025-12-03T07:50:15.474310Z","shell.execute_reply":"2025-12-03T07:50:15.494437Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"base_model_path = \"baseline_model.pth\"\nadvanced_model_path = \"advanced_model.pth\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T07:50:24.610118Z","iopub.execute_input":"2025-12-03T07:50:24.611193Z","iopub.status.idle":"2025-12-03T07:50:24.615769Z","shell.execute_reply.started":"2025-12-03T07:50:24.611116Z","shell.execute_reply":"2025-12-03T07:50:24.614784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- REPORT GENERATION ---\nresults = []\n\n# 1. Load Best Baseline\nmodel = CustomPlantCNN(num_classes=config.NUM_CLASSES).to(device)\nmodel.load_state_dict(torch.load(base_model_path))\n\n# Baseline: Standard Eval\nloss, acc, f1 = evaluate_model(model, test_loader, criterion, device, config.LABELS, use_tta=False)\nresults.append({'Model': 'Baseline', 'Method': 'Standard', 'Test Acc': acc, 'F1 Score': f1})\n\n# Baseline: + Thresholds\nthresholds = find_optimal_thresholds(model, val_loader, device)\nloss, acc, f1 = evaluate_model(model, test_loader, criterion, device, config.LABELS, thresholds=thresholds, use_tta=False)\nresults.append({'Model': 'Baseline', 'Method': '+ Threshold Calibration', 'Test Acc': acc, 'F1 Score': f1})\n\n# Baseline: + Thresholds + TTA\nloss, acc, f1 = evaluate_model(model, test_loader, criterion, device, config.LABELS, thresholds=thresholds, use_tta=True)\nresults.append({'Model': 'Baseline', 'Method': '+ TTA (Horizontal+Vertical)', 'Test Acc': acc, 'F1 Score': f1})\n\n\n# 2. Load Best Advanced\nmodel = CustomPlantCNN(num_classes=config.NUM_CLASSES).to(device)\nmodel.load_state_dict(torch.load(advanced_model_path))\n\n# Advanced: Standard Eval\nloss, acc, f1 = evaluate_model(model, test_loader, criterion, device, config.LABELS, use_tta=False)\nresults.append({'Model': 'Advanced (MixUp+Cos)', 'Method': 'Standard', 'Test Acc': acc, 'F1 Score': f1})\n\n# Advanced: + Thresholds\nthresholds = find_optimal_thresholds(model, val_loader, device)\nloss, acc, f1 = evaluate_model(model, test_loader, criterion, device, config.LABELS, thresholds=thresholds, use_tta=False)\nresults.append({'Model': 'Advanced (MixUp+Cos)', 'Method': '+ Threshold Calibration', 'Test Acc': acc, 'F1 Score': f1})\n\n# Advanced: + Thresholds + TTA\nloss, acc, f1 = evaluate_model(model, test_loader, criterion, device, config.LABELS, thresholds=thresholds, use_tta=True)\nresults.append({'Model': 'Advanced (MixUp+Cos)', 'Method': '+ TTA (Full Pipeline)', 'Test Acc': acc, 'F1 Score': f1})\n\n# --- Display Final DataFrame ---\ndf_results = pd.DataFrame(results)\nprint(\"\\n\" + \"=\"*40)\nprint(\"FINAL MODEL COMPARISON REPORT\")\nprint(\"=\"*40)\ndisplay(df_results)\n\n# Optional: Plot Comparison\nplt.figure(figsize=(10,6))\nplt.barh(df_results['Method'] + ' (' + df_results['Model'] + ')', df_results['Test Acc'], color='skyblue')\nplt.xlabel('Test Accuracy')\nplt.title('Impact of Improvements on Model Performance')\nplt.axvline(x=0.84, color='r', linestyle='--', label='Target (0.84)')\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T07:50:05.778214Z","iopub.status.idle":"2025-12-03T07:50:05.778594Z","shell.execute_reply.started":"2025-12-03T07:50:05.778399Z","shell.execute_reply":"2025-12-03T07:50:05.778415Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 1. UTILITY FUNCTIONS (Decoding & Plotting) ---\n\ndef decode_label(encoded_label, class_list):\n    \"\"\"Converts binary tensor/list to list of class strings\"\"\"\n    return [class_list[i] for i, val in enumerate(encoded_label) if val == 1]\n\ndef pred_and_plot_image(model, image_path, subplot, ground_truth=None, \n                        class_names=config.LABELS, thresholds=None, device=config.DEVICE):\n    \"\"\"Predicts and plots a single image with threshold support\"\"\"\n    \n    # Load Image\n    if isinstance(image_path, pathlib.PosixPath) or isinstance(image_path, str):\n        # Handle local path or string path\n        if str(image_path).startswith('http'):\n             img = Image.open(requests.get(image_path, stream=True).raw).convert('RGB')\n        else:\n             img = Image.open(image_path).convert('RGB')\n    \n    # Preprocess\n    transform = transforms.Compose([\n        transforms.Resize((config.INPUT_HEIGHT, config.INPUT_WIDTH)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=config.IMAGENET_MEAN, std=config.IMAGENET_STD)\n    ])\n    \n    img_tensor = transform(img).unsqueeze(0).to(device)\n    \n    # Predict\n    model.eval()\n    with torch.no_grad():\n        logits = model(img_tensor)\n        probs = torch.sigmoid(logits)\n        \n        if thresholds is not None:\n            # Broadcast thresholds to match batch size\n            preds = (probs > thresholds).int()\n        else:\n            preds = (probs > 0.5).int()\n            \n    preds = preds.cpu().squeeze().tolist()\n    predicted_labels = decode_label(preds, class_names)\n    \n    if not predicted_labels:\n        predicted_labels = [\"No Disease\"]\n\n    # Plot\n    plt.subplot(*subplot)\n    plt.imshow(img)\n    \n    if ground_truth:\n        title = f\"True: {ground_truth}\\nPred: {', '.join(predicted_labels)}\"\n    else:\n        title = f\"Pred: {', '.join(predicted_labels)}\"\n        \n    plt.title(title, fontsize=10)\n    plt.axis('off')\n\ndef plot_random_test_images(model, test_df, thresholds=None):\n    \"\"\"Plots a grid of random test images with predictions\"\"\"\n    num_images = 15\n    sample = test_df.sample(n=num_images, random_state=42)\n    image_paths = sample['image'].tolist()\n    labels = sample['labels'].tolist()\n    \n    rows = int(np.ceil(num_images / 5))\n    plt.figure(figsize=(20, rows * 4))\n    \n    for i, img_path in enumerate(image_paths):\n        pred_and_plot_image(\n            model=model,\n            image_path=img_path,\n            subplot=(rows, 5, i+1),\n            ground_truth=labels[i],\n            class_names=config.LABELS,\n            thresholds=thresholds\n        )\n    plt.tight_layout()\n    plt.show()\n\ndef analyze_model_performance(model, model_name, train_loader, val_loader, test_loader, thresholds=None):\n    \"\"\"\n    Runs a full analysis: Train/Val/Test Loss, Detailed Metrics, Confusion Matrix, and Visuals\n    \"\"\"\n    print(f\"\\n{'='*20} ANALYZING: {model_name} {'='*20}\")\n    model.to(device)\n    \n    # 1. Calculate Final Losses (Snapshot)\n    print(\"Calculating final Train/Val/Test stats...\")\n    train_loss, train_acc, _ = evaluate_model(model, train_loader, criterion, device, config.LABELS, thresholds=thresholds, verbose=False)\n    val_loss, val_acc, _ = evaluate_model(model, val_loader, criterion, device, config.LABELS, thresholds=thresholds, verbose=False)\n    test_loss, test_acc, test_f1 = evaluate_model(model, test_loader, criterion, device, config.LABELS, thresholds=thresholds, verbose=False)\n    \n    print(f\"\\n--- Overall Performance ---\")\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f}\")\n    print(f\"Val Loss:   {val_loss:.4f} | Val Acc:   {val_acc:.4f}\")\n    print(f\"Test Loss:  {test_loss:.4f} | Test Acc:  {test_acc:.4f}\")\n    print(f\"Test F1:    {test_f1:.4f}\")\n\n    # 2. Detailed Classification Report (Test Set)\n    print(f\"\\n--- Detailed Classification Report (Test Set) ---\")\n    \n    # Get all predictions for reporting\n    all_preds = []\n    all_targets = []\n    model.eval()\n    with torch.no_grad():\n        for images, labels in test_loader:\n            images, labels = images.to(device), labels.to(device)\n            logits = model(images)\n            probs = torch.sigmoid(logits)\n            if thresholds is not None:\n                preds = (probs > thresholds).float()\n            else:\n                preds = (probs > 0.5).float()\n            all_preds.append(preds.cpu().numpy())\n            all_targets.append(labels.cpu().numpy())\n            \n    all_preds = np.vstack(all_preds)\n    all_targets = np.vstack(all_targets)\n    \n    print(classification_report(all_targets, all_preds, target_names=config.LABELS, zero_division=0))\n    \n    # 3. Dominant Class Confusion Matrix\n    # (Converts multi-label to single-label based on max probability for cleaner visualization)\n    print(f\"\\n--- Confusion Matrix (Dominant Class) ---\")\n    pred_indices = np.argmax(all_preds, axis=1)\n    target_indices = np.argmax(all_targets, axis=1)\n    \n    cm = confusion_matrix(target_indices, pred_indices)\n    plt.figure(figsize=(10, 8))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=config.LABELS, yticklabels=config.LABELS)\n    plt.xlabel('Predicted')\n    plt.ylabel('True')\n    plt.title(f'{model_name} Confusion Matrix')\n    plt.show()\n    \n    # 4. Visual Predictions\n    print(f\"\\n--- Visual Predictions on Test Set ---\")\n    plot_random_test_images(model, test_df, thresholds=thresholds)\n\n\n# --- 2. MAIN EXECUTION BLOCK ---\n\nresults = []\n\n# ==================================\n# A. ANALYZE BASELINE MODEL\n# ==================================\nbaseline_model = CustomPlantCNN(num_classes=config.NUM_CLASSES).to(device)\nbaseline_model.load_state_dict(torch.load(base_model_path))\n\n# 1. Get Baseline Stats for Summary Table\nthresholds_base = find_optimal_thresholds(baseline_model, val_loader, device)\nloss, acc, f1 = evaluate_model(baseline_model, test_loader, criterion, device, config.LABELS, thresholds=thresholds_base, use_tta=True)\nresults.append({'Model': 'Baseline', 'Test Acc': acc, 'F1 Score': f1, 'Technique': 'Thresholds + TTA'})\n\n# 2. Run Deep Dive Analysis\nanalyze_model_performance(baseline_model, \"Baseline Model\", train_loader, val_loader, test_loader, thresholds=thresholds_base)\n\n\n# ==================================\n# B. ANALYZE ADVANCED MODEL\n# ==================================\nadvanced_model = CustomPlantCNN(num_classes=config.NUM_CLASSES).to(device)\nadvanced_model.load_state_dict(torch.load(advanced_model_path))\n\n# 1. Get Advanced Stats for Summary Table\nthresholds_adv = find_optimal_thresholds(advanced_model, val_loader, device)\nloss, acc, f1 = evaluate_model(advanced_model, test_loader, criterion, device, config.LABELS, thresholds=thresholds_adv, use_tta=True)\nresults.append({'Model': 'Advanced (MixUp+Cos)', 'Test Acc': acc, 'F1 Score': f1, 'Technique': 'Thresholds + TTA'})\n\n# 2. Run Deep Dive Analysis\nanalyze_model_performance(advanced_model, \"Advanced Model\", train_loader, val_loader, test_loader, thresholds=thresholds_adv)\n\n\n# --- 3. FINAL SUMMARY TABLE ---\ndf_results = pd.DataFrame(results)\nprint(\"\\n\" + \"=\"*40)\nprint(\"FINAL REPORT SUMMARY\")\nprint(\"=\"*40)\ndisplay(df_results)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T07:50:05.779994Z","iopub.status.idle":"2025-12-03T07:50:05.780300Z","shell.execute_reply.started":"2025-12-03T07:50:05.780144Z","shell.execute_reply":"2025-12-03T07:50:05.780159Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import requests\nfrom PIL import Image\nfrom io import BytesIO\nimport textwrap\n\ndef predict_from_url(url, model, transform, class_names, device):\n    # 1. Download the image\n    try:\n        response = requests.get(url)\n        response.raise_for_status()\n        img = Image.open(BytesIO(response.content)).convert('RGB')\n    except Exception as e:\n        print(f\"Error loading image: {e}\")\n        return\n\n    # 2. Preprocess\n    # We use the same 'val_transforms' (Resize + Normalize)\n    img_tensor = transform(img).unsqueeze(0) # Add batch dimension -> [1, 3, 224, 224]\n    img_tensor = img_tensor.to(device)\n\n    # 3. Predict\n    model.eval()\n    with torch.no_grad():\n        output = model(img_tensor)\n        probs = torch.sigmoid(output) # Convert to 0-1 range\n        \n        # Apply threshold to get binary labels\n        preds = (probs > 0.5).int().cpu().numpy()[0]\n        \n    # 4. Decode Labels\n    predicted_labels = [class_names[i] for i, val in enumerate(preds) if val == 1]\n    \n    if not predicted_labels:\n        predicted_labels = [\"Uncertain/Healthy\"] # Fallback if no class > 0.5\n        \n    prediction_text = \" | \".join(predicted_labels)\n\n    # 5. Visualize\n    plt.figure(figsize=(6, 6))\n    plt.imshow(img)\n    plt.axis('off')\n    \n    # Add title with wrapping to avoid cutting off text\n    title = f\"Prediction: {prediction_text}\"\n    plt.title(\"\\n\".join(textwrap.wrap(title, width=30)), fontsize=14, color='darkblue')\n    plt.show()\n\n# --- Run Predictions on your URLs ---\n\nurls = [\n    'https://www.planetnatural.com/wp-content/uploads/2012/12/common-rust-disease.jpg',\n    'https://www.greenlife.co.ke/wp-content/uploads/2022/04/powdery_mildew.jpg',\n    'https://c8.alamy.com/comp/PB2H05/frog-eye-leaf-spot-or-cercospora-diseases-on-leaves-of-suicide-tree-PB2H05.jpg'\n]\n\nprint(\"Running Predictions on Real Images...\")\nfor url in urls:\n    predict_from_url(url, model, val_transforms, config.LABELS, device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T07:50:05.781338Z","iopub.status.idle":"2025-12-03T07:50:05.781664Z","shell.execute_reply.started":"2025-12-03T07:50:05.781507Z","shell.execute_reply":"2025-12-03T07:50:05.781523Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# training_params =  training_hyperparams.get(config.TRAINING_PARAMS)","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:50:05.782850Z","iopub.status.idle":"2025-12-03T07:50:05.783139Z","shell.execute_reply.started":"2025-12-03T07:50:05.782993Z","shell.execute_reply":"2025-12-03T07:50:05.783009Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # To reduce clutter in the notebook I've turned the verbosity off, you can turn it on to see the full output\n# training_params[\"train_metrics_list\"] = ['my_accuracy']\n# training_params[\"valid_metrics_list\"] = ['my_accuracy']\n# training_params[\"metric_to_watch\"] = \"my_accuracy\"\n\n# # Set the silent mode to True to reduce clutter in the notebook, you can turn it on to see the full output\n# training_params[\"silent_mode\"] = True\n# training_params[\"optimizer\"] = 'AdamW'\n# training_params['average_best_models'] = True\n# training_params['ema'] = True\n# training_params[\"criterion_params\"] = {'smooth_eps': 0.20}\n# training_params[\"max_epochs\"] = 30\n# training_params[\"initial_lr\"] = 0.00001\n# training_params[\"loss\"] = criterion","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:50:05.784781Z","iopub.status.idle":"2025-12-03T07:50:05.785106Z","shell.execute_reply.started":"2025-12-03T07:50:05.784958Z","shell.execute_reply":"2025-12-03T07:50:05.784973Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model = models.get(config.MODEL_NAME, num_classes=config.NUM_CLASSES, pretrained_weights='imagenet')","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:50:05.786836Z","iopub.status.idle":"2025-12-03T07:50:05.787326Z","shell.execute_reply.started":"2025-12-03T07:50:05.787061Z","shell.execute_reply":"2025-12-03T07:50:05.787085Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# best_full_model = models.get(config.MODEL_NAME,\n#                         num_classes=config.NUM_CLASSES,\n#                         checkpoint_path=os.path.join(full_model_trainer.checkpoints_dir_path, \"average_model.pth\"))","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:50:05.788421Z","iopub.status.idle":"2025-12-03T07:50:05.788914Z","shell.execute_reply.started":"2025-12-03T07:50:05.788663Z","shell.execute_reply":"2025-12-03T07:50:05.788687Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# full_model_trainer.test(model=best_full_model,\n#             test_loader=test_loader,\n#             test_metrics_list=['my_accuracy'])","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:50:05.790579Z","iopub.status.idle":"2025-12-03T07:50:05.791050Z","shell.execute_reply.started":"2025-12-03T07:50:05.790805Z","shell.execute_reply":"2025-12-03T07:50:05.790830Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# plot_random_test_images(best_full_model, test_df)","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:50:05.792947Z","iopub.status.idle":"2025-12-03T07:50:05.793278Z","shell.execute_reply.started":"2025-12-03T07:50:05.793112Z","shell.execute_reply":"2025-12-03T07:50:05.793129Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# pred_and_plot_image(image_path='https://www.planetnatural.com/wp-content/uploads/2012/12/common-rust-disease.jpg', subplot=(1, 1, 1))","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:50:05.794701Z","iopub.status.idle":"2025-12-03T07:50:05.795178Z","shell.execute_reply.started":"2025-12-03T07:50:05.794918Z","shell.execute_reply":"2025-12-03T07:50:05.794942Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# pred_and_plot_image(image_path='https://www.greenlife.co.ke/wp-content/uploads/2022/04/powdery_mildew.jpg', subplot=(1, 1, 1))","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:50:05.796535Z","iopub.status.idle":"2025-12-03T07:50:05.796994Z","shell.execute_reply.started":"2025-12-03T07:50:05.796752Z","shell.execute_reply":"2025-12-03T07:50:05.796776Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# pred_and_plot_image(image_path='https://soybeanresearchinfo.com/wp-content/uploads/2020/05/Frogeye-leaf-spot-Daren-Mueller-17-1300x867.jpg', subplot=(1, 1, 1))\n","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:50:05.798282Z","iopub.status.idle":"2025-12-03T07:50:05.798706Z","shell.execute_reply.started":"2025-12-03T07:50:05.798520Z","shell.execute_reply":"2025-12-03T07:50:05.798545Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# pred_and_plot_image(image_path='https://c8.alamy.com/comp/PB2H05/frog-eye-leaf-spot-or-cercospora-diseases-on-leaves-of-suicide-tree-PB2H05.jpg', subplot=(1, 1, 1))","metadata":{"execution":{"iopub.status.busy":"2025-12-03T07:50:05.800102Z","iopub.status.idle":"2025-12-03T07:50:05.800439Z","shell.execute_reply.started":"2025-12-03T07:50:05.800271Z","shell.execute_reply":"2025-12-03T07:50:05.800287Z"},"trusted":true},"outputs":[],"execution_count":null}]}