{"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":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nimport os\nimport pytorch_lightning as pl\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn import model_selection\nimport torchvision.transforms as transforms\nimport torchvision.io \nimport librosa\nfrom PIL import Image\nimport albumentations as alb\nimport torch.multiprocessing as mp\nimport warnings\n\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:15:20.329220Z","iopub.execute_input":"2022-04-22T05:15:20.329984Z","iopub.status.idle":"2022-04-22T05:15:28.542098Z","shell.execute_reply.started":"2022-04-22T05:15:20.329844Z","shell.execute_reply":"2022-04-22T05:15:28.540907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q torchtoolbox timm","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:17:45.185088Z","iopub.execute_input":"2022-04-22T05:17:45.185412Z","iopub.status.idle":"2022-04-22T05:17:57.640208Z","shell.execute_reply.started":"2022-04-22T05:17:45.185364Z","shell.execute_reply":"2022-04-22T05:17:57.639028Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TransformationType:\n    TORCHVISION = \"torchvision\"\n    ALB = \"albumentations\"\n\nclass Models:\n    RESNET34 = \"resnet34\"\n    RESNET50 = \"resnet50\"\n    RESNEXT50 = \"resnext50_32x4d\"\n    EFFNET_B4 = \"tf_efficientnet_b4_ns\"\n    EFFNET_B0 = \"tf_efficientnet_b0_ns\"\n\nclass ImgStats:\n    IMAGENET_MEAN = [0.485, 0.456, 0.406]\n    IMAGENET_STD = [0.229, 0.224, 0.225]    \n\nclass WandbConfig:\n    WANDB_KEY = \"\"\n    WANDB_RUN_NAME = \"melspec_resnext_run4\"\n    WANDB_PROJECT = \"Pog_MusicClf_Train\"\n    USE_WANDB = False    \n\nclass SchedulerConfig:\n    SCHEDULER_PATIENCE = 4\n    SCHEDULER = 'LinearWithWarmup'\n    T_0 = 10 # for CosineAnnealingWarmRestarts\n    MIN_LR = 5e-7 # for CosineAnnealingWarmRestarts\n    MAX_LR = 1e-2\n    STEPS_PER_EPOCH = 0    \n    \n# CONSTANTS\nclass Config:\n    # whether to use mel spectrograms generated using audio augmentations ( multiple mel spec for one audio)\n    USE_MEL_SPEC_AUG = False\n    RUNTIME = \"KAGGLE\"\n    RESUME_FROM_CHKPT = None\n    NUM_CLASSES = 5\n    BATCH_SIZE = 128\n    NUM_FOLDS = 5\n    UNFREEZE_EPOCH_NO = 1\n    NUM_EPOCHS = 40\n    NUM_WORKERS = mp.cpu_count()\n    INPUT_IMAGE_SIZE = (128,128)\n    IMG_MEAN = ImgStats.IMAGENET_MEAN\n    IMG_STD = ImgStats.IMAGENET_STD\n    FAST_DEV_RUN = False\n    PRECISION = 16    \n    PATIENCE = 10    \n    SUBSET_ROWS_FRAC = 0.05\n    TRAIN_ON_SUBSET = False\n    RANDOM_SEED = 42\n    MODEL_TO_USE = Models.RESNEXT50\n    PRETRAINED = True            \n    FINE_TUNE = True    \n    FIND_LR = False\n    WEIGHT_DECAY = 1e-6\n    USE_MIXUP = False\n    # Parameter used to sample lambda values from the beta distribution. Recommended value 0.2   \n    MIXUP_ALPHA = 0.2   \n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')    \n    # model hyperparameters\n    MODEL_PARAMS = {    \n        \"drop_out\": 0.25,\n        \"lr\": 0.001,\n        \"warmup_prop\": 0.05\n    }","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:21:37.333299Z","iopub.execute_input":"2022-04-22T05:21:37.333657Z","iopub.status.idle":"2022-04-22T05:21:37.347433Z","shell.execute_reply.started":"2022-04-22T05:21:37.333625Z","shell.execute_reply":"2022-04-22T05:21:37.346322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if Config.RUNTIME == \"COLAB\":\n    Config.DATA_ROOT_FOLDER = \"/content/gdrive/MyDrive/Kaggle/Pog_Music_Classification/data/\"\n    Config.IMG_ROOT_FOLDER = \"/content/kaggle/processed_train/mel_spec/\"\nelif Config.RUNTIME == \"KAGGLE\":\n    Config.DATA_ROOT_FOLDER = \"/kaggle/input/kaggle-pog-series-s01e02/\"\n    Config.IMG_ROOT_FOLDER = \"/kaggle/input/pog-music-melspec-iter2/mel_spec_train/kaggle/processed_train/mel_spec/\"\nelse:\n    Config.DATA_ROOT_FOLDER = \"./data/\"\n    Config.IMG_ROOT_FOLDER = \"./data/processed_train/mel_spec/\"","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:21:37.349848Z","iopub.execute_input":"2022-04-22T05:21:37.350993Z","iopub.status.idle":"2022-04-22T05:21:37.365258Z","shell.execute_reply.started":"2022-04-22T05:21:37.350924Z","shell.execute_reply":"2022-04-22T05:21:37.363981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def config_to_dict(cfg):\n    # dir is an inbuilt python function that returns the list of attributes and methods of any object\n    return dict((name, getattr(cfg, name)) for name in dir(cfg) if not name.startswith('__'))","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:21:37.367085Z","iopub.execute_input":"2022-04-22T05:21:37.367528Z","iopub.status.idle":"2022-04-22T05:21:37.380145Z","shell.execute_reply.started":"2022-04-22T05:21:37.367482Z","shell.execute_reply":"2022-04-22T05:21:37.379128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# split the training dataframe into kfolds for cross validation. We do this before any processing is done\n# on the data. We use stratified group kfold if the target distribution is unbalanced and multiple records belong to\n# a single group (song_id). To prevent data leakage we need to ensure all records of a group are assigned to one fold only\n# i.e. there is no leakage between train and validation sets\ndef strat_group_kfold_dataframe(df, target_col_name, num_folds=Config.NUM_FOLDS):\n    # we create a new column called kfold and fill it with -1\n    df[\"kfold\"] = -1\n    # randomize of shuffle the rows of dataframe before splitting is done\n    df = df.sample(frac=1, random_state=Config.RANDOM_SEED).reset_index(drop=True)\n    # get the target data\n    y = df[target_col_name].values\n    if Config.USE_MEL_SPEC_AUG:\n        groups = df.song_id.values\n        skf = model_selection.StratifiedGroupKFold(n_splits=num_folds, shuffle=True, random_state=Config.RANDOM_SEED)\n        for fold, (train_index, val_index) in enumerate(skf.split(X=df, y=y, groups=groups)):\n            df.loc[val_index, \"kfold\"] = fold    \n    else:\n        skf = model_selection.StratifiedKFold(n_splits=num_folds, shuffle=True, random_state=Config.RANDOM_SEED)\n        for fold, (train_index, val_index) in enumerate(skf.split(X=df, y=y)):\n            df.loc[val_index, \"kfold\"] = fold \n    return df     \n\nif Config.USE_MEL_SPEC_AUG:\n    df_train = pd.read_csv(\"/kaggle/input/pog-musicclf-melspec-aug/df_train_aug.csv\")\n    df_train = df_train.drop([\"Unnamed: 0\"], axis=1)\n    df_train[\"row_id\"] = df_train.index\nelse:\n    df_train = pd.read_csv(Config.DATA_ROOT_FOLDER + \"train.csv\")\n    # filter out records without any corresponding mel spectrogram image\n    df_train[\"mspec_exists\"] = df_train.filename.map(\n        lambda fp: os.path.exists(Config.IMG_ROOT_FOLDER + fp.split(\".\")[0] + \".jpg\")\n    )\n    df_train = df_train[df_train.mspec_exists]\nif Config.TRAIN_ON_SUBSET:\n    print(f\"Selecting {Config.SUBSET_ROWS_FRAC * 100}% training data\")\n    df_train = df_train.sample(frac=Config.SUBSET_ROWS_FRAC, random_state=Config.RANDOM_SEED).reset_index(drop=True)\n    \ndf_train = strat_group_kfold_dataframe(df_train, target_col_name=\"genre_id\")\ndf_train.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:21:37.382776Z","iopub.execute_input":"2022-04-22T05:21:37.383216Z","iopub.status.idle":"2022-04-22T05:22:36.531243Z","shell.execute_reply.started":"2022-04-22T05:21:37.383144Z","shell.execute_reply":"2022-04-22T05:22:36.530188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if Config.USE_MEL_SPEC_AUG:\n    # if a song_id is assigned to multiple folds, the below df would return more than 19909 records\n    df_train.groupby([\"song_id\", \"kfold\"], as_index=False)[\"mel_spec\"].count()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:22:50.760918Z","iopub.execute_input":"2022-04-22T05:22:50.761779Z","iopub.status.idle":"2022-04-22T05:22:50.767023Z","shell.execute_reply.started":"2022-04-22T05:22:50.761728Z","shell.execute_reply":"2022-04-22T05:22:50.766016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Config.NUM_CLASSES = len(df_train.genre_id.unique())","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:22:50.777745Z","iopub.execute_input":"2022-04-22T05:22:50.779102Z","iopub.status.idle":"2022-04-22T05:22:50.786541Z","shell.execute_reply.started":"2022-04-22T05:22:50.779056Z","shell.execute_reply":"2022-04-22T05:22:50.785461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# A dataset contains the logic to fetch, load and if required transform data to bring it to a format\n# that can be used by dataloaders for training. Image size is (128, 385, 3)\nclass AudioMelSpecImgDataset(Dataset):\n    def __init__(self, df, file_name_col, target_col, img_root_folder, transform=None, target_transform=None):\n        self.df = df\n        self.file_name_col = file_name_col\n        self.target_col = target_col\n        self.img_root_folder = img_root_folder\n        self.transform = transform\n        self.target_transform = target_transform\n\n    def __getitem__(self, index):\n        if Config.USE_MEL_SPEC_AUG:\n            mel_spec_img = self.df.loc[index, self.file_name_col]\n            img_path = self.img_root_folder + mel_spec_img\n            id = self.df.loc[index, \"row_id\"] \n        else:            \n            file_name_noext = self.df.loc[index, self.file_name_col].split(\".\")[0]        \n            img_path = self.img_root_folder + \"/\" + file_name_noext + \".jpg\"\n            id = self.df.loc[index, \"song_id\"]        \n        img = Image.open(img_path)\n        img_arr = np.array(img)\n        img_label = self.df.loc[index, self.target_col]\n        if self.transform is not None:\n            augmented = self.transform(image=img_arr)\n            img_tfmd = augmented[\"image\"]\n            #img_tfmd = self.transform(img)            \n        if self.target_transform is not None:\n            img_label = self.target_transform(img_label)        \n        return id, img_tfmd, img_label\n\n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:22:50.789223Z","iopub.execute_input":"2022-04-22T05:22:50.790014Z","iopub.status.idle":"2022-04-22T05:22:50.804205Z","shell.execute_reply.started":"2022-04-22T05:22:50.789966Z","shell.execute_reply":"2022-04-22T05:22:50.803127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from torchvision.transforms import ToTensor, RandomResizedCrop, ToPILImage\n\n# train_transform = transforms.Compose([\n#         #ToPILImage(),\n#         RandomResizedCrop(size=(Config.INPUT_IMAGE_SIZE[0], Config.INPUT_IMAGE_SIZE[1])),                \n#         ToTensor()\n# ])\n\n# val_transform = transforms.Compose([\n#         #ToPILImage(),\n#         RandomResizedCrop(size=(Config.INPUT_IMAGE_SIZE[0], Config.INPUT_IMAGE_SIZE[1])),        \n#         ToTensor()        \n# ])","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:22:50.805934Z","iopub.execute_input":"2022-04-22T05:22:50.806537Z","iopub.status.idle":"2022-04-22T05:22:50.819761Z","shell.execute_reply.started":"2022-04-22T05:22:50.806495Z","shell.execute_reply":"2022-04-22T05:22:50.818630Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from albumentations.pytorch import ToTensorV2\nfrom albumentations.augmentations.transforms import Cutout\n\ntrain_transform = alb.Compose([\n        alb.RandomResizedCrop(Config.INPUT_IMAGE_SIZE[0], Config.INPUT_IMAGE_SIZE[1]),        \n        alb.HorizontalFlip(p=0.5),\n        alb.VerticalFlip(p=0.5),\n        #Cutout(),\n        alb.augmentations.transforms.JpegCompression(p=0.5),        \n        alb.augmentations.transforms.ImageCompression(\n            p=0.5, \n            compression_type=alb.augmentations.transforms.ImageCompression.ImageCompressionType.WEBP\n        ),\n        alb.Normalize(mean=Config.IMG_MEAN, std=Config.IMG_STD),\n        ToTensorV2()\n])\n\nval_transform = alb.Compose([\n        alb.CenterCrop(Config.INPUT_IMAGE_SIZE[0], Config.INPUT_IMAGE_SIZE[1]),\n        alb.Normalize(mean=Config.IMG_MEAN, std=Config.IMG_STD),\n        ToTensorV2()        \n])","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:22:50.822594Z","iopub.execute_input":"2022-04-22T05:22:50.823324Z","iopub.status.idle":"2022-04-22T05:22:50.839276Z","shell.execute_reply.started":"2022-04-22T05:22:50.823279Z","shell.execute_reply":"2022-04-22T05:22:50.838172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_fold_dls(fold, df_imgs, img_root_folder):\n    df_training = df_imgs[df_imgs[\"kfold\"] != fold].reset_index(drop=True)\n    df_val = df_imgs[df_imgs[\"kfold\"] == fold].reset_index(drop=True)\n    if Config.USE_MEL_SPEC_AUG:\n        file_name_col = \"mel_spec\"\n    else:\n        file_name_col = \"filename\"\n    ds_train = AudioMelSpecImgDataset(\n        df_training, \n        file_name_col=file_name_col,\n        target_col=\"genre_id\",\n        img_root_folder=img_root_folder,\n        transform=train_transform,\n        target_transform=torch.as_tensor\n    )\n    ds_val = AudioMelSpecImgDataset(\n        df_val, \n        file_name_col=file_name_col,\n        target_col=\"genre_id\",\n        img_root_folder=img_root_folder,\n        transform=val_transform,\n        target_transform=torch.as_tensor\n    )        \n    dl_train = DataLoader(ds_train, batch_size=Config.BATCH_SIZE, shuffle=True, num_workers=Config.NUM_WORKERS)    \n    dl_val = DataLoader(ds_val, batch_size=Config.BATCH_SIZE, num_workers=Config.NUM_WORKERS)\n    return dl_train, dl_val, ds_train, ds_val","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:22:50.841059Z","iopub.execute_input":"2022-04-22T05:22:50.841689Z","iopub.status.idle":"2022-04-22T05:22:50.851689Z","shell.execute_reply.started":"2022-04-22T05:22:50.841632Z","shell.execute_reply":"2022-04-22T05:22:50.850640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# display images along with their labels from a batch where images are in form of numpy arrays \n# if predictions are provided along with labels, these are displayed too\ndef show_batch(img_ds, num_items, num_rows, num_cols, predict_arr=None):\n    fig = plt.figure(figsize=(12, 6))    \n    img_index = np.random.randint(0, len(img_ds)-1, num_items)\n    for index, img_index in enumerate(img_index):  # list first 9 images\n        id, img, lb = img_ds[img_index]        \n        ax = fig.add_subplot(num_rows, num_cols, index + 1, xticks=[], yticks=[])\n        if isinstance(img, torch.Tensor):\n            img = img.detach().numpy()\n        if isinstance(img, np.ndarray):\n            # the image data has RGB channels at dim 0, the shape of 3, 64, 64 needs to be 64, 64, 3 for display            \n            img = img.transpose(1, 2, 0)\n            ax.imshow(img)        \n        if isinstance(lb, torch.Tensor):\n            # extract the label from label tensor\n            lb = lb.item()            \n        title = f\"Actual: {lb}\"\n        if predict_arr: \n            title += f\", Pred: {predict_arr[img_index]}\"        \n        ax.set_title(title)  ","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:22:51.009632Z","iopub.execute_input":"2022-04-22T05:22:51.010516Z","iopub.status.idle":"2022-04-22T05:22:51.020633Z","shell.execute_reply.started":"2022-04-22T05:22:51.010472Z","shell.execute_reply":"2022-04-22T05:22:51.019579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# img_path = \"./data/processed_train/mel_spec/000001.jpg\"\n# img = Image.open(img_path)\n# print(type(img))\n# img_arr = np.array(img)\n# img_arr = np.stack([img_arr]*3, axis=-1)\n# plt.imshow(img)\n# print(img_arr.shape)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:22:51.022793Z","iopub.execute_input":"2022-04-22T05:22:51.023362Z","iopub.status.idle":"2022-04-22T05:22:51.036689Z","shell.execute_reply.started":"2022-04-22T05:22:51.023309Z","shell.execute_reply":"2022-04-22T05:22:51.035015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(img_arr.shape)\n# img = transforms.ToPILImage()(img_arr)\n# print(type(img))\n# img_tfmd = val_transform(img)\n# img_np = img_tfmd.detach().numpy()\n# if isinstance(img_np, np.ndarray):\n#     # the image data has RGB channels at dim 0, the shape of 3, 64, 64 needs to be 64, 64, 3 for display            \n#     img_np = img_np.transpose(1, 2, 0)\n#     plt.imshow(img_np)\n#     print(img_np.shape)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:22:51.038926Z","iopub.execute_input":"2022-04-22T05:22:51.039401Z","iopub.status.idle":"2022-04-22T05:22:51.048027Z","shell.execute_reply.started":"2022-04-22T05:22:51.039352Z","shell.execute_reply":"2022-04-22T05:22:51.046591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dl_train, dl_val, ds_train, ds_val = get_fold_dls(0, df_train, Config.IMG_ROOT_FOLDER)\nshow_batch(ds_val, 8, 2, 4)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:22:51.051502Z","iopub.execute_input":"2022-04-22T05:22:51.052478Z","iopub.status.idle":"2022-04-22T05:22:51.739536Z","shell.execute_reply.started":"2022-04-22T05:22:51.052448Z","shell.execute_reply":"2022-04-22T05:22:51.738689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import get_linear_schedule_with_warmup\n\ndef get_linear_lr_scheduler(optimizer):\n    # Scheduler and math around the number of training steps.    \n    num_train_steps = Config.NUM_EPOCHS * SchedulerConfig.STEPS_PER_EPOCH\n    num_warmup_steps = int(Config.MODEL_PARAMS[\"warmup_prop\"] * Config.NUM_EPOCHS * SchedulerConfig.STEPS_PER_EPOCH)    \n    print(f\"num_train_steps = {num_train_steps}\")\n    print(f\"num_warmup_steps = {num_warmup_steps}\")\n    lr_scheduler = get_linear_schedule_with_warmup(\n            optimizer=optimizer,\n            num_warmup_steps=num_warmup_steps,\n            num_training_steps=num_train_steps\n        )\n    return lr_scheduler        ","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:24:22.586624Z","iopub.execute_input":"2022-04-22T05:24:22.587159Z","iopub.status.idle":"2022-04-22T05:24:22.598055Z","shell.execute_reply.started":"2022-04-22T05:24:22.587111Z","shell.execute_reply":"2022-04-22T05:24:22.596956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.optim.lr_scheduler import CosineAnnealingLR, CosineAnnealingWarmRestarts, ReduceLROnPlateau, OneCycleLR\n\ndef get_optimizer(lr, params):\n    model_optimizer = torch.optim.Adam(\n            filter(lambda p: p.requires_grad, params), \n            lr=lr,\n            weight_decay=Config.WEIGHT_DECAY\n        )\n    interval = \"epoch\"\n    if SchedulerConfig.SCHEDULER == \"CosineAnnealingWarmRestarts\":\n        lr_scheduler = CosineAnnealingWarmRestarts(\n                            model_optimizer, \n                            T_0=SchedulerConfig.T_0, \n                            T_mult=1, \n                            eta_min=SchedulerConfig.MIN_LR, \n                            last_epoch=-1\n                        )\n    elif SchedulerConfig.SCHEDULER == \"OneCycleLR\":\n        lr_scheduler = OneCycleLR(\n            optimizer=model_optimizer,\n            max_lr=SchedulerConfig.MAX_LR,\n            epochs=Config.NUM_EPOCHS,\n            steps_per_epoch=SchedulerConfig.STEPS_PER_EPOCH,\n            verbose=True\n        )\n        interval = \"step\"\n    elif SchedulerConfig.SCHEDULER == \"CosineAnnealingLR\":\n        lr_scheduler = CosineAnnealingLR(model_optimizer, eta_min=SchedulerConfig.MIN_LR, T_max=Config.NUM_EPOCHS)\n    elif SchedulerConfig.SCHEDULER == \"LinearWithWarmup\":\n        lr_scheduler = get_linear_lr_scheduler(model_optimizer)\n        interval = \"step\"\n    else:\n        # ReduceLROnPlateau throws an error is parameters are filtered, \n        # refer: https://github.com/PyTorchLightning/pytorch-lightning/issues/8720\n        model_optimizer = torch.optim.Adam(\n            params, \n            lr=lr,\n            weight_decay=Config.WEIGHT_DECAY\n        )\n        lr_scheduler = ReduceLROnPlateau(\n                            model_optimizer, \n                            mode=\"min\",                                                                \n                            factor=0.1,\n                            patience=SchedulerConfig.SCHEDULER_PATIENCE,\n                            min_lr=SchedulerConfig.MIN_LR,                                \n                            verbose=True\n                        )   \n    return {\n        \"optimizer\": model_optimizer, \n        \"lr_scheduler\": {\n            \"scheduler\": lr_scheduler,\n            \"interval\": interval,\n            \"monitor\": \"val_loss\",\n            \"frequency\": 1\n        }\n    }","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:24:22.614879Z","iopub.execute_input":"2022-04-22T05:24:22.615204Z","iopub.status.idle":"2022-04-22T05:24:22.640096Z","shell.execute_reply.started":"2022-04-22T05:24:22.615146Z","shell.execute_reply":"2022-04-22T05:24:22.635988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchtoolbox.tools import mixup_data, mixup_criterion\nimport torch.nn as nn\nfrom torch.nn.functional import cross_entropy\nimport torchmetrics\nimport timm\n\nclass MusicClfLitModel(pl.LightningModule):\n    def __init__(self, num_classes, hparams, model_to_use):\n        super().__init__()\n        self.save_hyperparameters()\n        self.lr = hparams[\"lr\"]\n        self.num_classes = num_classes\n        self.f1 = torchmetrics.F1(num_classes=num_classes)\n        self.backbone, self.classifier = self.get_backbone_classifier(model_to_use, hparams[\"drop_out\"], num_classes) \n\n    @staticmethod\n    def get_backbone_classifier(model_to_use, drop_out, num_classes):\n        pt_model = timm.create_model(model_to_use, pretrained=Config.PRETRAINED)\n        backbone = None\n        classifier = None\n        if model_to_use in [Models.RESNET34, Models.RESNET50, Models.RESNEXT50]:            \n            backbone = nn.Sequential(*list(pt_model.children())[:-1])\n            in_features = pt_model.fc.in_features\n            classifier = nn.Sequential(\n                nn.Dropout(drop_out),\n                nn.Linear(in_features, num_classes)\n            )\n        if model_to_use in [Models.EFFNET_B0, Models.EFFNET_B4]:\n            backbone = nn.Sequential(*list(pt_model.children())[:-1])\n            in_features = pt_model.classifier.in_features\n            classifier = nn.Linear(in_features, num_classes)\n                    \n        return backbone, classifier\n\n    def forward(self, x):\n        features = self.backbone(x)\n        features = torch.flatten(features, 1)                \n        x = self.classifier(features)\n        return x\n\n    def configure_optimizers(self):\n        return get_optimizer(lr=self.lr, params=self.parameters())\n\n    def train_with_mixup(self, X, y):\n        X, y_a, y_b, lam = mixup_data(X, y, alpha=Config.MIXUP_ALPHA)\n        y_pred = self(X)\n        loss_mixup = mixup_criterion(cross_entropy, y_pred, y_a, y_b, lam)\n        return loss_mixup\n\n    def training_step(self, batch, batch_idx):\n        id, X, y = batch        \n        if Config.USE_MIXUP:\n            loss = self.train_with_mixup(X, y)\n        else:\n            y_pred = self(X)\n            loss = cross_entropy(y_pred, y)                \n        #train_f1 = torchmetrics.functional.f1(preds=y_pred, target=y, num_classes=self.num_classes, average=\"micro\")\n        self.log(\"train_loss\", loss, on_step=True, on_epoch=True, logger=True, prog_bar=True)\n        #self.log(\"train_f1\", train_f1, on_step=True, on_epoch=True, logger=True, prog_bar=True)\n        return loss        \n\n    def validation_step(self, batch, batch_idx):\n        id, X, y = batch\n        y_pred = self(X)\n        val_loss = cross_entropy(y_pred, y)\n        current_lr = self.trainer.optimizers[0].param_groups[0]['lr']\n        #val_f1 = torchmetrics.functional.f1(preds=y_pred, target=y, num_classes=self.num_classes, average=\"micro\")\n        val_f1 = self.f1(preds=y_pred, target=y)\n        self.log(\"val_loss\", val_loss, on_step=True, on_epoch=True, logger=True, prog_bar=True)\n        self.log(\"val_f1\", val_f1, on_step=True, on_epoch=True, logger=True, prog_bar=True)\n        self.log(\"cur_lr\", current_lr, prog_bar=True, on_step=True, on_epoch=True, logger=True)\n        return {\"loss\": val_loss, \"val_f1\": val_f1, \"cur_lr\": current_lr}","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:24:22.646859Z","iopub.execute_input":"2022-04-22T05:24:22.651792Z","iopub.status.idle":"2022-04-22T05:24:22.695866Z","shell.execute_reply.started":"2022-04-22T05:24:22.651629Z","shell.execute_reply":"2022-04-22T05:24:22.694612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pytorch_lightning.callbacks import ModelCheckpoint, BackboneFinetuning, EarlyStopping\n\n# For results reproducibility \n# sets seeds for numpy, torch, python.random and PYTHONHASHSEED.\npl.seed_everything(Config.RANDOM_SEED, workers=True)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:24:22.701521Z","iopub.execute_input":"2022-04-22T05:24:22.702925Z","iopub.status.idle":"2022-04-22T05:24:22.718972Z","shell.execute_reply.started":"2022-04-22T05:24:22.702880Z","shell.execute_reply":"2022-04-22T05:24:22.717883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pytorch_lightning import LightningModule, Trainer\nfrom pytorch_lightning.callbacks import Callback\n\nclass MetricsAggCallback(Callback):\n    def __init__(self, metric_to_monitor, mode):\n        self.metric_to_monitor = metric_to_monitor\n        self.metrics = []\n        self.best_metric = None\n        self.mode = mode\n        self.best_metric_epoch = None\n        self.val_epoch_num = 0\n\n    def on_validation_epoch_end(self, trainer: Trainer, pl_module: LightningModule):\n        self.val_epoch_num += 1\n        metric_value = trainer.callback_metrics[self.metric_to_monitor].cpu().detach().item()\n        val_loss = trainer.callback_metrics[\"val_loss\"].cpu().detach().item()\n        current_lr = trainer.callback_metrics[\"cur_lr\"].cpu().detach().item()\n        print(f\"epoch = {self.val_epoch_num} => metric {self.metric_to_monitor} = {metric_value}, \" \\\n              f\"val_loss={val_loss}, lr={current_lr}\")\n        self.metrics.append(metric_value)\n        if self.mode == \"max\":\n            self.best_metric = max(self.metrics)\n            self.best_metric_epoch = self.metrics.index(self.best_metric)        ","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:24:22.834243Z","iopub.execute_input":"2022-04-22T05:24:22.835143Z","iopub.status.idle":"2022-04-22T05:24:22.848621Z","shell.execute_reply.started":"2022-04-22T05:24:22.835099Z","shell.execute_reply":"2022-04-22T05:24:22.847398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import wandb\nfrom kaggle_secrets import UserSecretsClient\n\ndef wandb_login():\n    user_secrets = UserSecretsClient()\n    wandb_secret = user_secrets.get_secret(\"wandb\")\n    wandb.login(key=wandb_secret)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:24:22.850945Z","iopub.execute_input":"2022-04-22T05:24:22.851751Z","iopub.status.idle":"2022-04-22T05:24:22.864555Z","shell.execute_reply.started":"2022-04-22T05:24:22.851704Z","shell.execute_reply":"2022-04-22T05:24:22.863504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_wandb_logger(fold):\n    logger = None\n    if WandbConfig.USE_WANDB:\n        if Config.RUNTIME == \"KAGGLE\":\n            wandb_login()\n        else:\n            wandb.login(key=WandbConfig.WANDB_KEY)\n        config_dict = config_to_dict(Config)\n        schd_config_dict = config_to_dict(SchedulerConfig)\n        merged_config_dict = {**config_dict, **schd_config_dict}\n        logger = WandbLogger(\n            name=WandbConfig.WANDB_RUN_NAME + f\"_fold{fold}\", \n            project=WandbConfig.WANDB_PROJECT,\n            config=merged_config_dict,\n            group=Config.MODEL_TO_USE\n        )\n    return logger","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:24:22.866626Z","iopub.execute_input":"2022-04-22T05:24:22.867374Z","iopub.status.idle":"2022-04-22T05:24:22.878099Z","shell.execute_reply.started":"2022-04-22T05:24:22.867328Z","shell.execute_reply":"2022-04-22T05:24:22.876575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pytorch_lightning.loggers import WandbLogger\nimport gc\n\ndef run_training(fold, dl_train, dl_val, fold_loss, fold_f1):\n    try:\n        fold_str = f\"fold{fold}\"\n        print(f\"Running training for {fold_str}\")\n        logger = None\n        val_loss_chkpt = \"best_model_{epoch}_{val_loss:.4f}\"\n        val_f1_chkpt = \"best_model_{epoch}_{val_f1:.4f}\"      \n        multiplicative = lambda epoch: 1.5\n        backbone_finetuning = BackboneFinetuning(Config.UNFREEZE_EPOCH_NO, multiplicative, verbose=True)\n        early_stopping_callback = EarlyStopping(monitor=\"val_loss\", patience=Config.PATIENCE, mode=\"min\", verbose=True)        \n        if fold is not None:       \n            val_loss_chkpt = fold_str + \"_\" + Config.MODEL_TO_USE + \"_\" + val_loss_chkpt\n            val_f1_chkpt = fold_str + \"_\" + Config.MODEL_TO_USE + \"_\" + val_f1_chkpt\n             \n        audio_model = MusicClfLitModel(\n            num_classes=Config.NUM_CLASSES, \n            hparams=Config.MODEL_PARAMS,        \n            model_to_use=Config.MODEL_TO_USE\n        )\n        logger = get_wandb_logger(fold)\n        print(\"Instantiated wandb logger\")    \n        val_loss_chkpt_callback = ModelCheckpoint(dirpath=\"./model\", verbose=True, \n                                                  monitor=\"val_loss\", mode=\"min\", filename=val_loss_chkpt)\n        val_f1_chkpt_callback = ModelCheckpoint(dirpath=\"./model\", verbose=True, \n                                                  monitor=\"val_f1\", mode=\"max\", filename=val_f1_chkpt)\n        acc_chkpt_callback = MetricsAggCallback(metric_to_monitor=\"val_f1\", mode=\"max\")\n        callbacks_to_use = [val_loss_chkpt_callback, val_f1_chkpt_callback, acc_chkpt_callback, early_stopping_callback]\n        if Config.RESUME_FROM_CHKPT is not None:\n            resume_from_checkpoint = Config.RESUME_FROM_CHKPT\n        else:\n            resume_from_checkpoint = None\n        if Config.PRETRAINED:\n            callbacks_to_use.append(backbone_finetuning)\n        trainer = pl.Trainer(\n            gpus=1,\n            # For results reproducibility \n            deterministic=True,\n            auto_select_gpus=True,\n            progress_bar_refresh_rate=20,\n            max_epochs=Config.NUM_EPOCHS,\n            logger=logger,\n            auto_lr_find=True,    \n            precision=Config.PRECISION,    \n            weights_summary=None, \n            fast_dev_run=Config.FAST_DEV_RUN,\n            resume_from_checkpoint=resume_from_checkpoint,                   \n            callbacks=callbacks_to_use\n        )\n        if Config.FIND_LR:\n            trainer.tune(model=audio_model, train_dataloaders=dl_train)\n            print(f\"Learning rate using trainer.tune = {audio_model.lr}\")\n        if Config.FIND_LR and Config.RUNTIME == \"COLAB\":\n            return\n        else:\n            print(\"Running trainer.fit\")\n            trainer.fit(audio_model, train_dataloaders=dl_train, val_dataloaders=dl_val)                \n            if not Config.FAST_DEV_RUN:\n                fold_loss.append((val_loss_chkpt_callback.best_model_score.cpu().detach().item(), val_loss_chkpt_callback.best_model_path))\n                fold_f1.append((acc_chkpt_callback.best_metric, val_f1_chkpt_callback.best_model_path))\n                print(f\"Loss for {fold_str} = {fold_loss[fold]}, f1 = {fold_f1[fold]}\")\n        del trainer, audio_model, backbone_finetuning, early_stopping_callback, acc_chkpt_callback, val_loss_chkpt_callback, val_f1_chkpt_callback \n        gc.collect()\n        torch.cuda.empty_cache()\n    except KeyboardInterrupt as e:\n        wandb.finish(exit_code=-1, quiet=True)\n        print(\"Marked the wandb run as failed\")","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:24:22.884571Z","iopub.execute_input":"2022-04-22T05:24:22.885880Z","iopub.status.idle":"2022-04-22T05:24:22.921814Z","shell.execute_reply.started":"2022-04-22T05:24:22.885797Z","shell.execute_reply":"2022-04-22T05:24:22.920598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.notebook import tqdm\n\n# For a specific fold get the predictions on oof (validation) data. We do this for each fold\n# We then use these oof predictions to calculate the cross validation score (using the evaluation metric)\ndef get_oof_preds(fold, fold_loss, dl_val):\n    # get the best model (having lowest val loss) for the fold\n    best_model_path_val_loss = fold_loss[fold][1]\n    print(f\"Using best model = {best_model_path_val_loss} for oof prediction on fold {fold} validation set\")\n    best_model = MusicClfLitModel.load_from_checkpoint(\n        checkpoint_path=best_model_path_val_loss,\n        num_classes=Config.NUM_CLASSES, \n        hparams=Config.MODEL_PARAMS,        \n        model_to_use=Config.MODEL_TO_USE\n    )\n    best_model.to(Config.DEVICE)        \n    if \"val_preds\" not in df_train.columns:\n        df_train[\"val_preds\"] = len(df_train) * [-100]\n    # For each class there is one predicted probability column\n    pred_proba_cols = [f\"proba_{i}\" for i in range(Config.NUM_CLASSES)]\n    with torch.no_grad():        \n        for id, X, y in tqdm(dl_val):\n            id = id.cpu().detach().numpy()            \n            # y_preds = [batch_size, num_classes]\n            y_preds_proba = best_model(X.to(Config.DEVICE))\n            y_preds = torch.argmax(y_preds_proba, dim=1)                            \n            y_preds_proba = y_preds_proba.cpu().detach().numpy()\n            y_preds = y_preds.cpu().detach().numpy().astype(int)\n            if Config.USE_MEL_SPEC_AUG:\n                df_train.loc[df_train.row_id.isin(id), \"val_preds\"] = y_preds            \n                df_train.loc[df_train.row_id.isin(id), pred_proba_cols] = y_preds_proba\n            else:\n                df_train.loc[df_train.song_id.isin(id), \"val_preds\"] = y_preds            \n                df_train.loc[df_train.song_id.isin(id), pred_proba_cols] = y_preds_proba","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:24:22.926791Z","iopub.execute_input":"2022-04-22T05:24:22.930035Z","iopub.status.idle":"2022-04-22T05:24:22.949656Z","shell.execute_reply.started":"2022-04-22T05:24:22.929987Z","shell.execute_reply":"2022-04-22T05:24:22.948598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import statistics\n\ndef print_exp_statistics(fold_loss, fold_f1):\n    fold_loss = [item[0] for item in fold_loss]\n    fold_f1 = [item[0] for item in fold_f1]\n    print(\"Loss across folds\")\n    print(fold_loss)\n    print(\"F1 across folds\")\n    print(fold_f1)\n    if len(fold_loss) > 1:\n        mean_loss = statistics.mean(fold_loss)\n        mean_f1 = statistics.mean(fold_f1)\n        std_loss = statistics.stdev(fold_loss)\n        std_f1 = statistics.stdev(fold_f1)\n        print(f\"mean loss across folds = {mean_loss}, loss stdev across fold = {std_loss}\")\n        print(f\"mean accuracy across folds = {mean_f1}, accuracy stdev across fold = {std_f1}\")","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:24:22.951855Z","iopub.execute_input":"2022-04-22T05:24:22.952681Z","iopub.status.idle":"2022-04-22T05:24:22.967203Z","shell.execute_reply.started":"2022-04-22T05:24:22.952582Z","shell.execute_reply":"2022-04-22T05:24:22.965616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import display\n\nfold_loss = []\nfold_f1 = []\nfor fold in range(Config.NUM_FOLDS):\n    dl_train, dl_val, ds_train, ds_val = get_fold_dls(fold, df_train, Config.IMG_ROOT_FOLDER)\n    SchedulerConfig.STEPS_PER_EPOCH = len(dl_train) + 1\n    run_training(fold, dl_train, dl_val, fold_loss, fold_f1)\n    if Config.FIND_LR and Config.RUNTIME == \"COLAB\":\n        break\n    else:\n        get_oof_preds(fold, fold_loss, dl_val)\n        # export the oof predictions to csv for later use in stacking\n        if Config.RUNTIME != \"KAGGLE\":\n            df_train.to_csv(Config.DATA_ROOT_FOLDER + \"df_train_oof_preds.csv\")\n        else:\n            df_train.to_csv(\"/kaggle/working/df_train_oof_preds.csv\")\n        print(f\"Saved OOF predictions for fold {fold}\")\n        display(df_train[df_train.kfold == fold].head())      \n    break\n\nwandb.finish(exit_code=0, quiet=True)\nprint(\"Marked the wandb run as successful\")\nprint_exp_statistics(fold_loss, fold_f1)       ","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:24:22.970056Z","iopub.execute_input":"2022-04-22T05:24:22.970651Z","iopub.status.idle":"2022-04-22T05:27:56.506650Z","shell.execute_reply.started":"2022-04-22T05:24:22.970607Z","shell.execute_reply":"2022-04-22T05:27:56.504682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train[df_train.val_preds == -100]","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:27:56.508505Z","iopub.status.idle":"2022-04-22T05:27:56.510018Z","shell.execute_reply.started":"2022-04-22T05:27:56.509637Z","shell.execute_reply":"2022-04-22T05:27:56.509683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import f1_score\n\ndef get_cv_score():\n    # export the oof predictions to csv for later use in stacking\n    if Config.RUNTIME != \"KAGGLE\":\n        df_train.to_csv(Config.DATA_ROOT_FOLDER + \"df_train_oof_preds.csv\")\n    else:\n        df_train.to_csv(\"/kaggle/working/df_train_oof_preds.csv\")\n    print(f\"Saved OOF predictions for fold {fold}\")\n    df_oof = df_train[df_train.val_preds != -100]\n    cv_f1 = f1_score(y_pred=df_oof.val_preds, y_true=df_oof.genre_id, average=\"micro\")\n    print(f\"Cross validation F1 score across {len(fold_loss)} folds = {cv_f1}\")\n\nif Config.FIND_LR and Config.RUNTIME == \"COLAB\":\n    pass\nelse:\n    get_cv_score()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:27:56.511588Z","iopub.status.idle":"2022-04-22T05:27:56.512653Z","shell.execute_reply.started":"2022-04-22T05:27:56.512311Z","shell.execute_reply":"2022-04-22T05:27:56.512346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class labels\ndf_genre = pd.read_csv(Config.DATA_ROOT_FOLDER + \"genres.csv\")\ndf_genre.genre","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:27:56.514565Z","iopub.status.idle":"2022-04-22T05:27:56.515323Z","shell.execute_reply.started":"2022-04-22T05:27:56.514897Z","shell.execute_reply":"2022-04-22T05:27:56.514955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay\n\ndef plot_confusion_matrix():\n    df_oof = df_train[df_train.val_preds != -100]\n    cm = confusion_matrix(y_true=df_oof.genre_id.values, y_pred=df_oof.val_preds.values) #, labels=df_genre.genre.values)\n    fig, ax = plt.subplots(figsize=(20, 7))\n    disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=df_genre.genre)\n    disp.plot(xticks_rotation=\"vertical\", ax=ax)\n    plt.show()\n\nif Config.FIND_LR and Config.RUNTIME == \"COLAB\":\n    pass\nelse:\n    plot_confusion_matrix()  ","metadata":{"execution":{"iopub.status.busy":"2022-04-22T05:27:56.517446Z","iopub.status.idle":"2022-04-22T05:27:56.518347Z","shell.execute_reply.started":"2022-04-22T05:27:56.518010Z","shell.execute_reply":"2022-04-22T05:27:56.518042Z"},"trusted":true},"execution_count":null,"outputs":[]}]}