{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<img src='https://github.com/Deci-AI/super-gradients/blob/master/documentation/assets/SG_img/SG%20-%20Horizontal%20Glow.png?raw=true'>\n\n## If you want to learn more about the EfficientNet family of model architectures, be sure to check out [my FREE course on Udemy](https://www.udemy.com/course/supergradients-efficientnet/).\n\nIn this notebook, you'll use the SuperGradients training library to classify plant pathology.\n\nBefore you dive into the code, it's worth talking about something: The difference between <span style=\"color:red\">multi-class</span>, <span style=\"color:green\">multi-task</span>, and <span style=\"color:orange\">multi-label learning</span>,.\n\n<span style=\"color:red\">\n\nIn multi-class learning, you're predicting a single label for each input, but each label is a single element from a set of possible labels. \n\nFor example, imagine you're working with a dataset of clothing images, and you want to predict whether each item is a shirt, a pair of pants, or a pair of shoes. To add another dimension, you should also predict the colour of each item, such as whether it's red, blue, or green. With multi-class learning, you can predict both the clothing item and its colour at the same time. For example, \"red shirt\" vs \"blue shirt\" vs \"brown show\" vs \"black shoe,\" etc. etc. \n\nYou typically use the Cross-entropy loss function to train a neural network on this problem.\n</span>\n\n<span style=\"color:green\">\n\nIn multi-task learning, you have multiple problems that need to be solved simultaneously. \n\nFor instance, you could predict the clothing item and its colour as separate tasks. The idea here is that solving one task could help solve the other task - for example, certain colours might be more common for certain types of clothing. In this case, you can assign each output (clothing item and colour) loss function to train a neural network. \n\nYou can then combine the loss functions by summing them up (or averaging) and using weights to balance the importance of each task.\n</span>\n\n<span style=\"color:orange\">\n\nMulti-label learning is a special case of multi-task learning. \n\nIn this scenario, you should label an image with multiple clothing items and their colours. You can break down the task into multiple binary classification problems to solve this. If the possible labels are \"shirt 👔\", \"pants 👖\", and \"shoes 👟\", and the possible colours are <span style=\"color:red\"> \"red\",</span> <span style=\"color:lightblue\"> \"blue\",</span> and <span style=\"color:green\">\"green\",</span> you would need to train the network to answer questions like \"is there a shirt in the image?\" and \"is the shirt in the image red?\" You want to use multi-label learning in the scenario where your labels are not mutually exclusive. \n\nYou would use the `BCEWithLogisLoss` for multi-label learning since it combines a sigmoid activation function and binary cross-entropy loss into a single function, making it efficient and numerically stable.\n</span>\n\n## What type of learning are we going to perfom here?\n\nIn this notebook we will perform multi-label classification.\n\nEach leaf can have one of many different pathologies which are non mutually exclusive. This means we need to think about how we'd define our loss function and how we would define a custom accuracy metric.\n\n#### What is SuperGradients?\n\n[SuperGradients](https://github.com/Deci-AI/super-gradients) is an open-source PyTorch based training library that has a number of pre-trained models for you to use, training recipies that will get you amazing accuracy, and many [training tricks](https://www.deeplearningdaily.community/t/tips-for-training-your-neural-networks/307) that you can use with just the \"flip of a switch\". For this example you'll use an EfficientNetB0 to perfom the classification. You can check out our [model zoo](https://github.com/Deci-AI/super-gradients/blob/master/src/super_gradients/training/Computer_Vision_Models_Pretrained_Checkpoints.md) and use any of the pretrained models we have available.\n\nFeel free to reach out to me on my community forum, [Deep Learning Daily (free and open to all)](https://www.deeplearningdaily.community/), should you have any questions.\n","metadata":{}},{"cell_type":"code","source":"%%capture\n!pip install imutils\n!pip install super-gradients==3.0.7","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"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\n\nimport torchvision\nfrom torchvision import datasets\nfrom torchvision import transforms\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\nimport super_gradients\nfrom super_gradients.common.object_names import Models\nfrom super_gradients.training import Trainer\nfrom super_gradients.training import training_hyperparams\nfrom super_gradients.training.metrics.classification_metrics import Accuracy, Top5\nfrom super_gradients.training.utils.early_stopping import EarlyStop\nfrom super_gradients.training import models\nfrom super_gradients.training.utils.callbacks import Phase\nfrom super_gradients.common.registry import register_metric, register_model,register_loss","metadata":{"execution":{"iopub.status.busy":"2023-03-12T14:34:24.694689Z","iopub.execute_input":"2023-03-12T14:34:24.695469Z","iopub.status.idle":"2023-03-12T14:34:39.316294Z","shell.execute_reply.started":"2023-03-12T14:34:24.695431Z","shell.execute_reply":"2023-03-12T14:34:39.315423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-03-12T14:34:39.330827Z","iopub.execute_input":"2023-03-12T14:34:39.331114Z","iopub.status.idle":"2023-03-12T14:34:39.340791Z","shell.execute_reply.started":"2023-03-12T14:34:39.331079Z","shell.execute_reply":"2023-03-12T14:34:39.339915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Resize images and save to disk\n\nThe below code will resize all the images and copy to the working directory.\n\nThis helps with GPU utilization, since this dataset is massive and my first attempts led to slow training. I decided to resize and save to disk, this helps with GPU utilization and speeds up training.\n\nNote that it will take at least 90 minutes to resize all the times since there are A LOT of them.","metadata":{}},{"cell_type":"code","source":"!mkdir data","metadata":{"execution":{"iopub.status.busy":"2023-03-12T14:34:39.345699Z","iopub.execute_input":"2023-03-12T14:34:39.346827Z","iopub.status.idle":"2023-03-12T14:34:40.477166Z","shell.execute_reply.started":"2023-03-12T14:34:39.346789Z","shell.execute_reply":"2023-03-12T14:34:40.476190Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Set the desired output dimensions\noutput_size = (224, 224)\n\n# Get a list of all image file paths in the input directory\nimage_paths = list(paths.list_images(config.DATA_DIR))\n\n# Create a progress bar object\nprogress_bar = tqdm(total=len(image_paths), desc='Resizing images')\n\n# Loop over all image file paths\nfor 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\nprogress_bar.close()","metadata":{"execution":{"iopub.status.busy":"2023-03-12T14:34:40.484514Z","iopub.execute_input":"2023-03-12T14:34:40.484952Z","iopub.status.idle":"2023-03-12T15:42:10.568860Z","shell.execute_reply.started":"2023-03-12T14:34:40.484907Z","shell.execute_reply":"2023-03-12T15:42:10.568479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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)\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":"2023-03-12T15:42:10.570685Z","iopub.execute_input":"2023-03-12T15:42:10.571092Z","iopub.status.idle":"2023-03-12T15:42:10.576012Z","shell.execute_reply.started":"2023-03-12T15:42:10.571052Z","shell.execute_reply":"2023-03-12T15:42:10.575625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df, valid_df, test_df = split_df('../input/plant-pathology-2021-fgvc8/train.csv')","metadata":{"execution":{"iopub.status.busy":"2023-03-12T15:42:10.577675Z","iopub.execute_input":"2023-03-12T15:42:10.578044Z","iopub.status.idle":"2023-03-12T15:42:10.654188Z","shell.execute_reply.started":"2023-03-12T15:42:10.578008Z","shell.execute_reply":"2023-03-12T15:42:10.653905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"(train_df['labels'].value_counts()).plot(kind='bar')","metadata":{"execution":{"iopub.status.busy":"2023-03-12T15:42:10.656471Z","iopub.execute_input":"2023-03-12T15:42:10.657506Z","iopub.status.idle":"2023-03-12T15:42:10.966594Z","shell.execute_reply.started":"2023-03-12T15:42:10.657468Z","shell.execute_reply":"2023-03-12T15:42:10.966240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"You can see that some examples belong to multiple labels. Let's get a more accurate count of each class.\n\nThis is the line of code you will use `train_df['labels'].str.split(expand=True).stack().reset_index(drop=True)`\n\nLet me break this down for you step-by-step\n\n1. `train_df['labels']`: Gets the 'labels' column from the train_df dataframe.\n2. `str.split(expand=True)`: Applies the `split()` function to each string in the 'labels' column to split the string into a list of substrings using whitespace as the delimiter. The `expand=True` argument ensures that the returned object is a DataFrame where each substring gets a separate column in the DataFrame.\n3. `stack()`: Reshapes the resulting DataFrame by stacking the columns into rows so that all substrings are represented as a single column. Each row in the resulting Series object contains a single substring.\n4. `reset_index(drop=True)`: Resets the index of the resulting Series object so that it starts from 0 and drops the old index.\n\nThe result is a pandas Series object containing all the substrings from the 'labels' column of the train_df dataframe. \n\nEach element in the Series corresponds to a single label from the original 'labels' column, and the elements are ordered the same way they appeared in the original column. \n\nBy counting the frequency of each element in this Series, you can obtain the value counts for all classes.","metadata":{}},{"cell_type":"code","source":"all_labels = train_df['labels'].str.split(expand=True).stack().reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T15:42:10.968014Z","iopub.execute_input":"2023-03-12T15:42:10.968592Z","iopub.status.idle":"2023-03-12T15:42:10.994256Z","shell.execute_reply.started":"2023-03-12T15:42:10.968513Z","shell.execute_reply":"2023-03-12T15:42:10.993991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_labels.value_counts().plot(kind='bar')","metadata":{"execution":{"iopub.status.busy":"2023-03-12T15:42:10.996442Z","iopub.execute_input":"2023-03-12T15:42:10.997330Z","iopub.status.idle":"2023-03-12T15:42:11.214420Z","shell.execute_reply.started":"2023-03-12T15:42:10.997293Z","shell.execute_reply":"2023-03-12T15:42:11.214062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-03-12T15:42:11.215914Z","iopub.execute_input":"2023-03-12T15:42:11.216507Z","iopub.status.idle":"2023-03-12T15:42:11.223021Z","shell.execute_reply.started":"2023-03-12T15:42:11.216469Z","shell.execute_reply":"2023-03-12T15:42:11.222764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"examine_images(train_df, num_images=20)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T15:42:11.224581Z","iopub.execute_input":"2023-03-12T15:42:11.225237Z","iopub.status.idle":"2023-03-12T15:42:17.332829Z","shell.execute_reply.started":"2023-03-12T15:42:11.225201Z","shell.execute_reply":"2023-03-12T15:42:17.332434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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_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()\n\n# initialize our training and validation set data augmentation pipeline\ntrain_transforms = transforms.Compose([\n#   resize, \n  auto_augment,\n#   random_augment,\n  make_tensor,\n  normalize\n])\n\nval_transforms = transforms.Compose([resize, make_tensor, normalize])","metadata":{"execution":{"iopub.status.busy":"2023-03-12T15:42:17.333964Z","iopub.execute_input":"2023-03-12T15:42:17.334343Z","iopub.status.idle":"2023-03-12T15:42:17.340815Z","shell.execute_reply.started":"2023-03-12T15:42:17.334296Z","shell.execute_reply":"2023-03-12T15:42:17.340455Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-03-12T15:42:17.347973Z","iopub.execute_input":"2023-03-12T15:42:17.348271Z","iopub.status.idle":"2023-03-12T15:42:18.104124Z","shell.execute_reply.started":"2023-03-12T15:42:17.348244Z","shell.execute_reply":"2023-03-12T15:42:18.103577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-03-12T15:42:18.106120Z","iopub.execute_input":"2023-03-12T15:42:18.106744Z","iopub.status.idle":"2023-03-12T15:42:18.112022Z","shell.execute_reply.started":"2023-03-12T15:42:18.106694Z","shell.execute_reply":"2023-03-12T15:42:18.111291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-03-12T15:42:18.113656Z","iopub.execute_input":"2023-03-12T15:42:18.114320Z","iopub.status.idle":"2023-03-12T15:42:18.122642Z","shell.execute_reply.started":"2023-03-12T15:42:18.114284Z","shell.execute_reply":"2023-03-12T15:42:18.122260Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-03-12T15:42:18.124168Z","iopub.execute_input":"2023-03-12T15:42:18.124907Z","iopub.status.idle":"2023-03-12T15:42:18.137269Z","shell.execute_reply.started":"2023-03-12T15:42:18.124863Z","shell.execute_reply":"2023-03-12T15:42:18.136760Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-03-12T15:42:18.138821Z","iopub.execute_input":"2023-03-12T15:42:18.139468Z","iopub.status.idle":"2023-03-12T15:42:18.150751Z","shell.execute_reply.started":"2023-03-12T15:42:18.139430Z","shell.execute_reply":"2023-03-12T15:42:18.150388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-03-12T15:42:18.152360Z","iopub.execute_input":"2023-03-12T15:42:18.152795Z","iopub.status.idle":"2023-03-12T15:42:18.164565Z","shell.execute_reply.started":"2023-03-12T15:42:18.152758Z","shell.execute_reply":"2023-03-12T15:42:18.164260Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_weights","metadata":{"execution":{"iopub.status.busy":"2023-03-12T15:42:18.165831Z","iopub.execute_input":"2023-03-12T15:42:18.166101Z","iopub.status.idle":"2023-03-12T15:42:18.190087Z","shell.execute_reply.started":"2023-03-12T15:42:18.166054Z","shell.execute_reply":"2023-03-12T15:42:18.189752Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"markdown","source":"# Model output\n\nThe output of the EfficientNet model for multilabel classification will be a tensor of shape (batch_size, num_classes), where batch_size is the number of examples in the batch and num_classes is the number of classes in the dataset. \n\nEach element of the tensor will be a real number between -infinity and +infinity, representing the model's confidence that the corresponding class is present in the input image.\n\nTo obtain the predicted labels, you need to pass the output through a sigmoid function and then apply a threshold to the output. \n\nFor example, you can use a threshold of 0.5 and consider all classes with output values greater than 0.5 as present, and those with values less than 0.5 as absent. \n\n# Custom accuracy metric\n\nYou need to define a custom accuracy metric for multi-label classification.\n\nDefined below, you have a PyTorch Metric subclass called `MyAccuracy,` which calculates the accuracy of binary classification predictions. The metric is registered to SuperGradients using a decorator called `register_metric,` which takes a string argument representing the metric's name.\n\nThis metric allows for us to correctly handle multi-hot encoded labels using the `all` method to determine whether all elements in a row of the predicted and target tensors match. \n\nThe `update` method takes in the predictions and targets as input in this implementation. It applies a threshold of 0.5 to the predictions using the sigmoid function. It then converts the predictions to integers using the `int()` function. \n\nNext, it computes the number of correct predictions for each sample in the batch by comparing the predicted and target tensors using the `all` method. \n\nFinally, it updates the \"correct\" and \"total\" states of the class by adding the number of correct predictions and the total number of predictions, respectively.\n\nThe `compute` method calculates and returns the accuracy as the ratio of correct predictions to the total number of predictions.","metadata":{}},{"cell_type":"code","source":"from torchmetrics import Metric\nimport torch\nfrom super_gradients.common.registry import register_metric\n\n@register_metric('my_accuracy')\n# Define a new class named MyAccuracy, which inherits from the Metric class in the PyTorch library\nclass MyAccuracy(Metric):\n    # Constructor method that takes in a number of classes as an argument\n    def __init__(self, num_classes=config.NUM_CLASSES):\n        # Calls the constructor of the parent class Metric\n        super().__init__()\n        # Adds two states to the instance of the class: \"correct\" and \"total\"\n        self.add_state(\"correct\", default=torch.tensor(0), dist_reduce_fx=\"sum\")\n        self.add_state(\"total\", default=torch.tensor(0), dist_reduce_fx=\"sum\")\n\n    # A method that takes in predictions and target values and updates the state of the class\n    def update(self, preds: torch.Tensor, target: torch.Tensor):\n        # Applies a threshold of 0.5 to the predictions and converts them to integers\n        preds = (torch.sigmoid(preds) > 0.50).int()\n        self.correct += torch.sum((preds == target).all(dim=1))\n        self.total += target.shape[0]\n        \n    # A method that calculates and returns the accuracy\n    def compute(self):\n        # Calculates the accuracy as the ratio of the number of correct predictions to the total number of predictions\n        return self.correct.float() / self.total","metadata":{"execution":{"iopub.status.busy":"2023-03-12T15:42:18.191597Z","iopub.execute_input":"2023-03-12T15:42:18.192250Z","iopub.status.idle":"2023-03-12T15:42:18.199938Z","shell.execute_reply.started":"2023-03-12T15:42:18.192209Z","shell.execute_reply":"2023-03-12T15:42:18.199660Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_params =  training_hyperparams.get(config.TRAINING_PARAMS)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T15:42:18.203190Z","iopub.execute_input":"2023-03-12T15:42:18.203781Z","iopub.status.idle":"2023-03-12T15:42:18.347673Z","shell.execute_reply.started":"2023-03-12T15:42:18.203746Z","shell.execute_reply":"2023-03-12T15:42:18.347404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\ntraining_params[\"train_metrics_list\"] = ['my_accuracy']\ntraining_params[\"valid_metrics_list\"] = ['my_accuracy']\ntraining_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\ntraining_params[\"silent_mode\"] = True\ntraining_params[\"optimizer\"] = 'AdamW'\ntraining_params['average_best_models'] = True\ntraining_params['ema'] = True\ntraining_params[\"criterion_params\"] = {'smooth_eps': 0.20}\ntraining_params[\"max_epochs\"] = 30\ntraining_params[\"initial_lr\"] = 0.00001\ntraining_params[\"loss\"] = criterion","metadata":{"execution":{"iopub.status.busy":"2023-03-12T15:42:18.349637Z","iopub.execute_input":"2023-03-12T15:42:18.350222Z","iopub.status.idle":"2023-03-12T15:42:18.355019Z","shell.execute_reply.started":"2023-03-12T15:42:18.350186Z","shell.execute_reply":"2023-03-12T15:42:18.354770Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = models.get(config.MODEL_NAME, num_classes=config.NUM_CLASSES, pretrained_weights='imagenet')","metadata":{"execution":{"iopub.status.busy":"2023-03-12T15:42:18.356429Z","iopub.execute_input":"2023-03-12T15:42:18.356974Z","iopub.status.idle":"2023-03-12T15:42:36.754114Z","shell.execute_reply.started":"2023-03-12T15:42:18.356934Z","shell.execute_reply":"2023-03-12T15:42:36.753644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"full_model_trainer = Trainer(experiment_name='0_Baseline_Experiment', ckpt_root_dir=config.CHECKPOINT_DIR)\n\nfull_model_trainer.train(model=model, \n              training_params=training_params, \n              train_loader=train_loader,\n              valid_loader=val_loader)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T15:42:36.757696Z","iopub.execute_input":"2023-03-12T15:42:36.758009Z","iopub.status.idle":"2023-03-12T18:33:12.606772Z","shell.execute_reply.started":"2023-03-12T15:42:36.757980Z","shell.execute_reply":"2023-03-12T18:33:12.605510Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-03-12T18:33:12.611791Z","iopub.execute_input":"2023-03-12T18:33:12.612599Z","iopub.status.idle":"2023-03-12T18:33:14.591403Z","shell.execute_reply.started":"2023-03-12T18:33:12.612547Z","shell.execute_reply":"2023-03-12T18:33:14.590925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-03-12T18:33:14.596124Z","iopub.execute_input":"2023-03-12T18:33:14.598918Z","iopub.status.idle":"2023-03-12T18:33:38.280686Z","shell.execute_reply.started":"2023-03-12T18:33:14.598876Z","shell.execute_reply":"2023-03-12T18:33:38.280196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import Tuple\nimport requests\nimport torchvision\nimport random\nimport textwrap\n\ndef pred_and_plot_image(image_path: str, \n                        subplot: Tuple[int, int, int],  # subplot tuple for `subplot()` function\n                        ground_truth:str = None,\n                        model: torch.nn.Module = best_full_model,\n                        image_size: Tuple[int, int] = (config.INPUT_HEIGHT, config.INPUT_WIDTH),\n                        transform: torchvision.transforms = None,\n                        device: torch.device=config.DEVICE):\n\n    if isinstance(image_path, pathlib.PosixPath):\n        img = Image.open(image_path)\n    else: \n        img = Image.open(requests.get(image_path, stream=True).raw)\n\n    # create transformation for image (if one doesn't exist)\n    if transform is None:\n        transform = transforms.Compose([\n            transforms.Resize(image_size),\n            transforms.ToTensor(),\n            transforms.Normalize(mean=config.IMAGENET_MEAN,\n                                 std=config.IMAGENET_STD),\n        ])\n    transformed_image = transform(img)\n\n    # make sure the model is on the target device\n    model.to(device)\n    # turn on model evaluation mode and inference mode\n    model.eval()\n    with torch.inference_mode():\n        # add an extra dimension to image (model requires samples in [batch_size, color_channels, height, width])\n        transformed_image = transformed_image.unsqueeze(dim=0)\n        # make a prediction on image with an extra dimension and send it to the target device\n        target_image_pred = model(transformed_image.to(device))\n        # apply sigmoid to predictions, return 1 at each index where greater than threshold         \n    preds = (torch.sigmoid(target_image_pred) > 0.50).int()\n    # from tensor to list         \n    preds = torch.round(preds).squeeze().tolist()\n    # convert float to ints        \n    preds = [int(i) for i in preds]\n    predicted_labels = decode_label(preds, config.LABELS)\n\n    # plot image with predicted label \n    plt.subplot(*subplot)\n    plt.imshow(img)\n    if isinstance(image_path, pathlib.PosixPath):\n        # actual label\n        title = f\"Ground Truth: {ground_truth} | Pred: {' '.join(predicted_labels)}\"\n    else:\n        title = f\"Pred: {' '.join(predicted_labels)}\"\n    plt.title(\"\\n\".join(textwrap.wrap(title, width=20)))  # wrap text using textwrap.wrap() function\n    plt.axis(False)\n    \n\ndef plot_random_test_images(model, test_images):\n    num_images_to_plot = 30\n\n    # extract image paths and labels from the dataframe\n    test_image_paths = test_images['image'].tolist()\n    test_image_labels = test_images['labels'].tolist()\n\n    # sample k image paths and labels\n    random_indices = random.sample(range(len(test_image_paths)), k=num_images_to_plot)\n    test_image_path_sample = [pathlib.PosixPath(test_image_paths[i]) for i in random_indices]\n    test_image_label_sample = [test_image_labels[i] for i in random_indices]\n\n    # set up subplots\n    num_rows = int(np.ceil(num_images_to_plot / 5))\n    fig, ax = plt.subplots(num_rows, 5, figsize=(15, num_rows * 3))\n    ax = ax.flatten()\n\n    # Make predictions on and plot the images\n    for i, image_path in enumerate(test_image_path_sample):\n        label = test_image_label_sample[i]\n        pred_and_plot_image(model=model,\n                            image_path=image_path,\n                            ground_truth=label,\n                            subplot=(num_rows, 5, i+1),\n                            image_size=(config.INPUT_HEIGHT, config.INPUT_WIDTH))\n\n    # adjust spacing between subplots\n    plt.subplots_adjust(wspace=1)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-03-12T20:39:42.636222Z","iopub.execute_input":"2023-03-12T20:39:42.636612Z","iopub.status.idle":"2023-03-12T20:39:42.653275Z","shell.execute_reply.started":"2023-03-12T20:39:42.636575Z","shell.execute_reply":"2023-03-12T20:39:42.652802Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_random_test_images(best_full_model, test_df)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T20:39:43.581175Z","iopub.execute_input":"2023-03-12T20:39:43.582324Z","iopub.status.idle":"2023-03-12T20:39:46.073506Z","shell.execute_reply.started":"2023-03-12T20:39:43.582281Z","shell.execute_reply":"2023-03-12T20:39:46.072654Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-03-12T20:39:46.076172Z","iopub.execute_input":"2023-03-12T20:39:46.076849Z","iopub.status.idle":"2023-03-12T20:39:46.519953Z","shell.execute_reply.started":"2023-03-12T20:39:46.076812Z","shell.execute_reply":"2023-03-12T20:39:46.519426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-03-12T20:40:35.337047Z","iopub.execute_input":"2023-03-12T20:40:35.337518Z","iopub.status.idle":"2023-03-12T20:40:36.462239Z","shell.execute_reply.started":"2023-03-12T20:40:35.337473Z","shell.execute_reply":"2023-03-12T20:40:36.461785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-03-12T20:42:25.808314Z","iopub.execute_input":"2023-03-12T20:42:25.808941Z","iopub.status.idle":"2023-03-12T20:42:26.665466Z","shell.execute_reply.started":"2023-03-12T20:42:25.808899Z","shell.execute_reply":"2023-03-12T20:42:26.664977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-03-12T20:44:20.273564Z","iopub.execute_input":"2023-03-12T20:44:20.274570Z","iopub.status.idle":"2023-03-12T20:44:20.987437Z","shell.execute_reply.started":"2023-03-12T20:44:20.274528Z","shell.execute_reply":"2023-03-12T20:44:20.986948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Your homework\n\nCopy/fork this notebook and try some different architectures.\n\nIf you have a question you can leave a comment on this notebook, or visit the community and post it in the [Q&A section](https://www.deeplearningdaily.community/c/qanda/8).\n\n## Use a different pretrained model\n\nYou can change the model you use. Take a look at the [SG model zoo](https://github.com/Deci-AI/super-gradients/blob/master/src/super_gradients/training/Computer_Vision_Models_Pretrained_Checkpoints.md)\n\nFor example, if you wanted to use RegNet you would do the following:\n\n```\nresnet_imagenet_model = models.get(model_name='regnetY800', num_classes=NUM_CLASSES, pretrained_weights='imagenet)\nresnet_params =  training_hyperparams.get('training_hyperparams/imagenet_regnetY_train_params')\n```\n\nNote you can also pass 'model_name=regnetY200', 'model_name=regnetY400', 'model_name=regnetY600' to try a variety of the architecture\n\nFor ResNet50, you would do:\n\n```\nresnet_imagenet_model = models.get(model_name='resnet50', num_classes=NUM_CLASSES, pretrained_weights='imagenet)\nresnet_params =  training_hyperparams.get('training_hyperparams/imagenet_resnet50_train_params')\n```\n\nNote you can also pass 'model_name=resnet18' or 'model_name=resnet34' to try a variety of the architecture\n\nFor MobileNetV2, you would do:\n\n```\nmobilenet_imagenet_model = models.get(model_name='mobilenet_v2', num_classes=NUM_CLASSES, pretrained_weights='imagenet)\nresnet_params =  training_hyperparams.get('training_hyperparams/imagenet_mobilenetv2_train_params')\n```\n\nFor MobileNetV3, you would do:\n\n```\nmobilenet_imagenet_model = models.get(model_name='mobilenet_v3_large', num_classes=NUM_CLASSES, pretrained_weights='imagenet)\nresnet_params =  training_hyperparams.get('training_hyperparams/imagenet_mobilenetv3_train_params')\n```\n\nNote you can also pass 'model_name=mobilenet_v3_small' to try a variety of the architecture\n\n\nFor ViT, you would do:\n\n\n```\nvit_imagenet_model = models.get(model_name='vit_base', num_classes=NUM_CLASSES, pretrained_weights='imagenet')\nvit_params =  training_hyperparams.get(\"training_hyperparams/imagenet_vit_train_params\")\n```\n\nNote you can also pass 'model_name=vit_large' to try a variety of the architecture\n\n\nI encourage you play around with different optimizers, all you have to do is change the value of `training_params[\"optimizer\"]`. You can use one of ['Adam','SGD','RMSProp'] out of the box. You can play around with the optimizer params as well.\n\nIn general, play and tweak around the training recipies...\n\n## Training recipes\n\nSuperGradients has a number of [training recipes](https://github.com/Deci-AI/super-gradients/tree/master/src/super_gradients/recipes) you can use. [See here](https://github.com/Deci-AI/super-gradients/blob/master/src/super_gradients/recipes/training_hyperparams/default_train_params.yaml) for more information about the training params.\n","metadata":{}}]}