{"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":"## Using [motono0223](https://www.kaggle.com/motono0223)'s notebook as baseline. I am adding fetaures such as \n* Class weights for cross entropy loss\n* adding aditional images thats isnt in thumbnails folder\n* balanced accuracy as metric for saving model as there is huge class imbalance\n","metadata":{}},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Import Required Libraries 📚</h1></span>","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport math\nimport copy\nimport time\nimport random\nimport glob\nfrom matplotlib import pyplot as plt\n\n# For data manipulation\nimport numpy as np\nimport pandas as pd\n\n# Pytorch Imports\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda import amp\nimport torchvision\n\n# Utils\nimport joblib\nfrom tqdm import tqdm\nfrom collections import defaultdict\n\n# Sklearn Imports\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import StratifiedKFold\n\n# For Image Models\nimport timm\n\n# Albumentations for augmentations\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# For colored terminal text\nfrom colorama import Fore, Back, Style\nb_ = Fore.BLUE\nsr_ = Style.RESET_ALL\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n# For descriptive error messages\nos.environ['CUDA_LAUNCH_BLOCKING'] = \"1\"","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:08.378506Z","iopub.execute_input":"2023-10-08T18:20:08.379285Z","iopub.status.idle":"2023-10-08T18:20:15.552715Z","shell.execute_reply.started":"2023-10-08T18:20:08.379254Z","shell.execute_reply":"2023-10-08T18:20:15.551699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Training Configuration ⚙️</h1></span>","metadata":{}},{"cell_type":"code","source":"CONFIG = {\n    \"seed\": 42,\n    \"epochs\": 20,\n    \"img_size\": 512,\n    \"model_name\": \"tf_efficientnet_b0_ns\",\n    \"checkpoint_path\" : \"/kaggle/input/tf-efficientnet/pytorch/tf-efficientnet-b0/1/tf_efficientnet_b0_aa-827b6e33.pth\",\n    \"num_classes\": 5,\n    \"train_batch_size\": 32,\n    \"valid_batch_size\": 64,\n    \"learning_rate\": 1e-4,\n    \"scheduler\": 'CosineAnnealingLR',\n    \"min_lr\": 1e-6,\n    \"T_max\": 500,\n    \"weight_decay\": 1e-6,\n    \"fold\" : 0,\n    \"n_fold\": 5,\n    \"n_accumulate\": 1,\n    \"device\": torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\"),\n}","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:15.554753Z","iopub.execute_input":"2023-10-08T18:20:15.555268Z","iopub.status.idle":"2023-10-08T18:20:15.583673Z","shell.execute_reply.started":"2023-10-08T18:20:15.555236Z","shell.execute_reply":"2023-10-08T18:20:15.582784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Set Seed for Reproducibility</h1></span>","metadata":{}},{"cell_type":"code","source":"def set_seed(seed=42):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed(CONFIG['seed'])","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:15.584947Z","iopub.execute_input":"2023-10-08T18:20:15.585811Z","iopub.status.idle":"2023-10-08T18:20:15.608074Z","shell.execute_reply.started":"2023-10-08T18:20:15.585776Z","shell.execute_reply":"2023-10-08T18:20:15.607153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT_DIR = '/kaggle/input/UBC-OCEAN'\nTRAIN_DIR = '/kaggle/input/UBC-OCEAN/train_thumbnails'\nALT_TEST_DIR = '/kaggle/input/UBC-OCEAN/test_images'\nTEST_DIR = '/kaggle/input/UBC-OCEAN/test_thumbnails'\nALT_TRAIN_DIR = '/kaggle/input/UBC-OCEAN/train_images'","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:15.610834Z","iopub.execute_input":"2023-10-08T18:20:15.611292Z","iopub.status.idle":"2023-10-08T18:20:15.620469Z","shell.execute_reply.started":"2023-10-08T18:20:15.611269Z","shell.execute_reply":"2023-10-08T18:20:15.619546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# this should load few extra images which arent thumbnails\ndef get_train_file_path(image_id):\n    if os.path.exists(f\"{TRAIN_DIR}/{image_id}_thumbnail.png\"):\n        return f\"{TRAIN_DIR}/{image_id}_thumbnail.png\"\n    else:\n        return f\"{ALT_TRAIN_DIR}/{image_id}.png\"\n#    return f\"{TRAIN_DIR}/{image_id}.png\"","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:15.62203Z","iopub.execute_input":"2023-10-08T18:20:15.622289Z","iopub.status.idle":"2023-10-08T18:20:15.63316Z","shell.execute_reply.started":"2023-10-08T18:20:15.622268Z","shell.execute_reply":"2023-10-08T18:20:15.632128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Read the Data 📖</h1>","metadata":{}},{"cell_type":"code","source":"train_images = sorted(glob.glob(f\"{TRAIN_DIR}/*.png\"))","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:15.634376Z","iopub.execute_input":"2023-10-08T18:20:15.635131Z","iopub.status.idle":"2023-10-08T18:20:15.692554Z","shell.execute_reply.started":"2023-10-08T18:20:15.6351Z","shell.execute_reply":"2023-10-08T18:20:15.691726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(f\"{ROOT_DIR}/train.csv\")\ndf['file_path'] = df['image_id'].apply(get_train_file_path)\n# df = df[ df[\"file_path\"].isin(train_images) ].reset_index(drop=True)\ndf","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:15.69372Z","iopub.execute_input":"2023-10-08T18:20:15.694445Z","iopub.status.idle":"2023-10-08T18:20:15.75487Z","shell.execute_reply.started":"2023-10-08T18:20:15.694423Z","shell.execute_reply":"2023-10-08T18:20:15.753942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder = LabelEncoder()\ndf['label'] = encoder.fit_transform(df['label'])\n\nwith open(\"label_encoder.pkl\", \"wb\") as fp:\n    joblib.dump(encoder, fp)","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:15.756241Z","iopub.execute_input":"2023-10-08T18:20:15.756565Z","iopub.status.idle":"2023-10-08T18:20:15.765027Z","shell.execute_reply.started":"2023-10-08T18:20:15.756523Z","shell.execute_reply":"2023-10-08T18:20:15.764067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONFIG['T_max'] = df.shape[0] * (CONFIG[\"n_fold\"]-1) * CONFIG['epochs'] // CONFIG['train_batch_size'] // CONFIG[\"n_fold\"]\nCONFIG['T_max']","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:15.766102Z","iopub.execute_input":"2023-10-08T18:20:15.76673Z","iopub.status.idle":"2023-10-08T18:20:15.775405Z","shell.execute_reply.started":"2023-10-08T18:20:15.766706Z","shell.execute_reply":"2023-10-08T18:20:15.774333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Class weights</h1></span>","metadata":{}},{"cell_type":"code","source":"def compute_class_weights(df, label_column):\n    \"\"\"\n    Compute class weights based on the inverse of class frequencies.\n    \n    Parameters:\n        df (pd.DataFrame): DataFrame containing the data.\n        label_column (str): Name of the column containing class labels.\n    \n    Returns:\n        class_weights (dict): Dictionary containing weights for each class.\n    \"\"\"\n    # Get the total number of samples\n    total_samples = len(df)\n    \n    # Get the number of classes\n    num_classes = df[label_column].nunique()\n    \n    # Get the count of each class\n    class_counts = df[label_column].value_counts().to_dict()\n    \n    # Compute class weights\n    class_weights = {class_label: total_samples / (num_classes * count) \n                     for class_label, count in class_counts.items()}\n    \n    return class_weights\n\n# Compute class weights for the provided dataset\nclass_weights = compute_class_weights(df, 'label')\nclass_weights = np.array([1.0868686868686868,0.867741935483871,0.4846846846846847,2.2893617021276595,2.3391304347826085])\nclass_weights_tensor = torch.tensor(class_weights, dtype=torch.float32)","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:15.779514Z","iopub.execute_input":"2023-10-08T18:20:15.779773Z","iopub.status.idle":"2023-10-08T18:20:15.820948Z","shell.execute_reply.started":"2023-10-08T18:20:15.779752Z","shell.execute_reply":"2023-10-08T18:20:15.820156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Create Folds</h1></span>","metadata":{}},{"cell_type":"code","source":"skf = StratifiedKFold(n_splits=CONFIG['n_fold'])\n\nfor fold, ( _, val_) in enumerate(skf.split(X=df, y=df.label)):\n      df.loc[val_ , \"kfold\"] = int(fold)","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:15.822136Z","iopub.execute_input":"2023-10-08T18:20:15.822439Z","iopub.status.idle":"2023-10-08T18:20:15.834077Z","shell.execute_reply.started":"2023-10-08T18:20:15.82241Z","shell.execute_reply":"2023-10-08T18:20:15.832978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Dataset Class</h1></span>","metadata":{}},{"cell_type":"code","source":"class UBCDataset(Dataset):\n    def __init__(self, df, transforms=None):\n        self.df = df\n        self.file_names = df['file_path'].values\n        self.labels = df['label'].values\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_path = self.file_names[index]\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        label = self.labels[index]\n        \n        if self.transforms:\n            img = self.transforms(image=img)[\"image\"]\n            \n        return {\n            'image': img,\n            'label': torch.tensor(label, dtype=torch.long)\n        }","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:15.83639Z","iopub.execute_input":"2023-10-08T18:20:15.836629Z","iopub.status.idle":"2023-10-08T18:20:15.843061Z","shell.execute_reply.started":"2023-10-08T18:20:15.836609Z","shell.execute_reply":"2023-10-08T18:20:15.84176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Augmentations</h1></span>","metadata":{}},{"cell_type":"code","source":"data_transforms = {\n    \"train\": A.Compose([\n        A.Resize(CONFIG['img_size'], CONFIG['img_size']),\n        A.ShiftScaleRotate(shift_limit=0.1, \n                           scale_limit=0.15, \n                           rotate_limit=60, \n                           p=0.5),\n        A.HueSaturationValue(\n                hue_shift_limit=0.2, \n                sat_shift_limit=0.2, \n                val_shift_limit=0.2, \n                p=0.5\n            ),\n        A.RandomBrightnessContrast(\n                brightness_limit=(-0.1,0.1), \n                contrast_limit=(-0.1, 0.1), \n                p=0.5\n            ),\n        A.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                max_pixel_value=255.0, \n                p=1.0\n            ),\n        ToTensorV2()], p=1.),\n    \n    \"valid\": A.Compose([\n        A.Resize(CONFIG['img_size'], CONFIG['img_size']),\n        A.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                max_pixel_value=255.0, \n                p=1.0\n            ),\n        ToTensorV2()], p=1.)\n}","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:15.844308Z","iopub.execute_input":"2023-10-08T18:20:15.845204Z","iopub.status.idle":"2023-10-08T18:20:15.855157Z","shell.execute_reply.started":"2023-10-08T18:20:15.845181Z","shell.execute_reply":"2023-10-08T18:20:15.854217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">GeM Pooling</h1></span>\n\n<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Code taken from <a href=\"https://amaarora.github.io/2020/08/30/gempool.html\">GeM Pooling Explained</a></span>\n\n![](https://i.imgur.com/thTgYWG.jpg)","metadata":{}},{"cell_type":"code","source":"class GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super(GeM, self).__init__()\n        self.p = nn.Parameter(torch.ones(1)*p)\n        self.eps = eps\n\n    def forward(self, x):\n        return self.gem(x, p=self.p, eps=self.eps)\n        \n    def gem(self, x, p=3, eps=1e-6):\n        return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1./p)\n        \n    def __repr__(self):\n        return self.__class__.__name__ + \\\n                '(' + 'p=' + '{:.4f}'.format(self.p.data.tolist()[0]) + \\\n                ', ' + 'eps=' + str(self.eps) + ')'","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:15.858309Z","iopub.execute_input":"2023-10-08T18:20:15.858556Z","iopub.status.idle":"2023-10-08T18:20:15.869906Z","shell.execute_reply.started":"2023-10-08T18:20:15.858536Z","shell.execute_reply":"2023-10-08T18:20:15.868968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Create Model</h1></span>","metadata":{}},{"cell_type":"code","source":"class UBCModel(nn.Module):\n    def __init__(self, model_name, num_classes, pretrained=True, checkpoint_path=None):\n        super(UBCModel, self).__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained, checkpoint_path=checkpoint_path)\n\n        in_features = self.model.classifier.in_features\n        self.model.classifier = nn.Identity()\n        self.model.global_pool = nn.Identity()\n        self.pooling = GeM()\n        self.linear = nn.Linear(in_features, num_classes)\n        self.softmax = nn.Softmax(dim=1)\n\n    def forward(self, images):\n        features = self.model(images)\n        pooled_features = self.pooling(features).flatten(1)\n        output = self.linear(pooled_features)\n        return output\n\n    \nmodel = UBCModel(CONFIG['model_name'], CONFIG['num_classes'], checkpoint_path=CONFIG['checkpoint_path'])\nmodel.to(CONFIG['device']);","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:15.87279Z","iopub.execute_input":"2023-10-08T18:20:15.873044Z","iopub.status.idle":"2023-10-08T18:20:20.299647Z","shell.execute_reply.started":"2023-10-08T18:20:15.873024Z","shell.execute_reply":"2023-10-08T18:20:20.298727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Loss Function</h1></span>","metadata":{}},{"cell_type":"code","source":"loss = nn.CrossEntropyLoss(weight = class_weights_tensor.to(CONFIG['device']))\ndef criterion(outputs, labels):\n    return loss(outputs, labels)","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:20.300883Z","iopub.execute_input":"2023-10-08T18:20:20.301239Z","iopub.status.idle":"2023-10-08T18:20:20.307543Z","shell.execute_reply.started":"2023-10-08T18:20:20.301207Z","shell.execute_reply":"2023-10-08T18:20:20.305465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Training Function</h1></span>","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import balanced_accuracy_score\ndef train_one_epoch(model, optimizer, scheduler, dataloader, device, epoch):\n    model.train()\n    \n    dataset_size = 0\n    running_loss = 0.0\n    running_acc  = 0.0\n    all_preds=[]\n    all_labels=[]\n    bar = tqdm(enumerate(dataloader), total=len(dataloader))\n    for step, data in bar:\n        images = data['image'].to(device, dtype=torch.float)\n        labels = data['label'].to(device, dtype=torch.long)\n        \n        batch_size = images.size(0)\n        \n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss = loss / CONFIG['n_accumulate']\n            \n        loss.backward()\n    \n        if (step + 1) % CONFIG['n_accumulate'] == 0:\n            optimizer.step()\n\n            # zero the parameter gradients\n            optimizer.zero_grad()\n\n            if scheduler is not None:\n                scheduler.step()\n                \n        _, predicted = torch.max(torch.nn.Softmax(dim=1)(outputs), 1)\n        acc = torch.sum( predicted == labels )\n        all_preds.extend(predicted.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n        running_loss += (loss.item() * batch_size)\n        running_acc  += acc.item()\n        dataset_size += batch_size\n        \n        epoch_loss = running_loss / dataset_size\n        epoch_acc = running_acc / dataset_size\n        \n        bar.set_postfix(Epoch=epoch, Train_Loss=epoch_loss, Train_Acc=epoch_acc,\n                        LR=optimizer.param_groups[0]['lr'])\n    gc.collect()\n    bl_accuracy_score = balanced_accuracy_score(all_labels, all_preds)\n    return epoch_loss, epoch_acc,bl_accuracy_score","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:20.308966Z","iopub.execute_input":"2023-10-08T18:20:20.309499Z","iopub.status.idle":"2023-10-08T18:20:20.320209Z","shell.execute_reply.started":"2023-10-08T18:20:20.309468Z","shell.execute_reply":"2023-10-08T18:20:20.319189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Validation Function</h1></span>","metadata":{}},{"cell_type":"code","source":"@torch.inference_mode()\ndef valid_one_epoch(model, dataloader, device, epoch):\n    model.eval()\n    \n    dataset_size = 0\n    running_loss = 0.0\n    running_acc = 0.0\n    all_preds=[]\n    all_labels=[]\n    bar = tqdm(enumerate(dataloader), total=len(dataloader))\n    for step, data in bar:        \n        images = data['image'].to(device, dtype=torch.float)\n        labels = data['label'].to(device, dtype=torch.long)\n        \n        batch_size = images.size(0)\n\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n\n        _, predicted = torch.max(torch.nn.Softmax(dim=1)(outputs), 1)\n        all_preds.extend(predicted.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n        acc = torch.sum( predicted == labels )\n\n        running_loss += (loss.item() * batch_size)\n        running_acc  += acc.item()\n        dataset_size += batch_size\n        \n        epoch_loss = running_loss / dataset_size\n        epoch_acc = running_acc / dataset_size\n        \n        bar.set_postfix(Epoch=epoch, Valid_Loss=epoch_loss, Valid_Acc=epoch_acc,\n                        LR=optimizer.param_groups[0]['lr'])   \n    \n    gc.collect()\n    bl_accuracy_score = balanced_accuracy_score(all_labels, all_preds)\n    return epoch_loss, epoch_acc,bl_accuracy_score","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:20.321553Z","iopub.execute_input":"2023-10-08T18:20:20.321863Z","iopub.status.idle":"2023-10-08T18:20:20.335722Z","shell.execute_reply.started":"2023-10-08T18:20:20.321836Z","shell.execute_reply":"2023-10-08T18:20:20.334845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Run Training</h1></span>","metadata":{}},{"cell_type":"code","source":"def run_training(model, optimizer, scheduler, device, num_epochs):\n    if torch.cuda.is_available():\n        print(\"[INFO] Using GPU: {}\\n\".format(torch.cuda.get_device_name()))\n    \n    start = time.time()\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_epoch_acc = -np.inf\n    history = defaultdict(list)\n    \n    for epoch in range(1, num_epochs + 1): \n        gc.collect()\n        train_epoch_loss, train_epoch_acc, train_bl_accuracy_score = train_one_epoch(model, optimizer, scheduler, \n                                           dataloader=train_loader, \n                                           device=CONFIG['device'], epoch=epoch)\n        \n        val_epoch_loss, val_epoch_acc, val_bl_accuracy_score = valid_one_epoch(model, valid_loader, device=CONFIG['device'], \n                                         epoch=epoch)\n    \n        history['Train Loss'].append(train_epoch_loss)\n        history['Valid Loss'].append(val_epoch_loss)\n        history['Train Accuracy'].append(train_epoch_acc)\n        history['Valid Accuracy'].append(val_epoch_acc)\n        history['Train balanced Accuracy'].append(train_bl_accuracy_score)\n        history['Valid balanced Accuracy'].append(val_bl_accuracy_score)\n        history['lr'].append( scheduler.get_lr()[0] )\n        \n        # deep copy the model\n        if best_epoch_acc <= val_bl_accuracy_score:\n            print(f\"{b_}Validation Balanced Accuracy Improved ({best_epoch_acc} ---> {val_bl_accuracy_score})\")\n            best_epoch_acc = val_bl_accuracy_score\n            best_model_wts = copy.deepcopy(model.state_dict())\n            PATH = \"Acc{:.2f}_Loss{:.4f}_epoch{:.0f}.bin\".format(best_epoch_acc, val_epoch_loss, epoch)\n            torch.save(model.state_dict(), PATH)\n            # Save a model file from the current directory\n            print(f\"Model Saved{sr_}\")\n            \n        print()\n    \n    end = time.time()\n    time_elapsed = end - start\n    print('Training complete in {:.0f}h {:.0f}m {:.0f}s'.format(\n        time_elapsed // 3600, (time_elapsed % 3600) // 60, (time_elapsed % 3600) % 60))\n    print(\"Best Accuracy: {:.4f}\".format(best_epoch_acc))\n    \n    # load best model weights\n    model.load_state_dict(best_model_wts)\n    \n    return model, history","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:20.338685Z","iopub.execute_input":"2023-10-08T18:20:20.338986Z","iopub.status.idle":"2023-10-08T18:20:20.350073Z","shell.execute_reply.started":"2023-10-08T18:20:20.33895Z","shell.execute_reply":"2023-10-08T18:20:20.349201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fetch_scheduler(optimizer):\n    if CONFIG['scheduler'] == 'CosineAnnealingLR':\n        scheduler = lr_scheduler.CosineAnnealingLR(optimizer,T_max=CONFIG['T_max'], \n                                                   eta_min=CONFIG['min_lr'])\n    elif CONFIG['scheduler'] == 'CosineAnnealingWarmRestarts':\n        scheduler = lr_scheduler.CosineAnnealingWarmRestarts(optimizer,T_0=CONFIG['T_0'], \n                                                             eta_min=CONFIG['min_lr'])\n    elif CONFIG['scheduler'] == None:\n        return None\n        \n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:20.35146Z","iopub.execute_input":"2023-10-08T18:20:20.351743Z","iopub.status.idle":"2023-10-08T18:20:20.366554Z","shell.execute_reply.started":"2023-10-08T18:20:20.351716Z","shell.execute_reply":"2023-10-08T18:20:20.365608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_loaders(df, fold):\n    df_train = df[df.kfold != fold].reset_index(drop=True)\n    df_valid = df[df.kfold == fold].reset_index(drop=True)\n    \n    train_dataset = UBCDataset(df_train, transforms=data_transforms[\"train\"])\n    valid_dataset = UBCDataset(df_valid, transforms=data_transforms[\"valid\"])\n\n    train_loader = DataLoader(train_dataset, batch_size=CONFIG['train_batch_size'], \n                              num_workers=2, shuffle=True, pin_memory=True, drop_last=True)\n    valid_loader = DataLoader(valid_dataset, batch_size=CONFIG['valid_batch_size'], \n                              num_workers=2, shuffle=False, pin_memory=True)\n    \n    return train_loader, valid_loader","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:20.367463Z","iopub.execute_input":"2023-10-08T18:20:20.367676Z","iopub.status.idle":"2023-10-08T18:20:20.380903Z","shell.execute_reply.started":"2023-10-08T18:20:20.367657Z","shell.execute_reply":"2023-10-08T18:20:20.380051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Prepare Dataloaders</h1></span>","metadata":{}},{"cell_type":"code","source":"train_loader, valid_loader = prepare_loaders(df, fold=CONFIG[\"fold\"])","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:20.381896Z","iopub.execute_input":"2023-10-08T18:20:20.382146Z","iopub.status.idle":"2023-10-08T18:20:20.396967Z","shell.execute_reply.started":"2023-10-08T18:20:20.382126Z","shell.execute_reply":"2023-10-08T18:20:20.396015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Define Optimizer and Scheduler</h1></span>\n","metadata":{}},{"cell_type":"code","source":"optimizer = optim.Adam(model.parameters(), lr=CONFIG['learning_rate'], \n                       weight_decay=CONFIG['weight_decay'])\nscheduler = fetch_scheduler(optimizer)","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:20.397974Z","iopub.execute_input":"2023-10-08T18:20:20.398223Z","iopub.status.idle":"2023-10-08T18:20:20.408575Z","shell.execute_reply.started":"2023-10-08T18:20:20.398204Z","shell.execute_reply":"2023-10-08T18:20:20.407493Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Start Training</h1></span>\n","metadata":{}},{"cell_type":"code","source":"model, history = run_training(model, optimizer, scheduler,\n                              device=CONFIG['device'],\n                              num_epochs=CONFIG['epochs'])","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:20:20.409607Z","iopub.execute_input":"2023-10-08T18:20:20.409832Z","iopub.status.idle":"2023-10-08T18:24:45.071987Z","shell.execute_reply.started":"2023-10-08T18:20:20.409813Z","shell.execute_reply":"2023-10-08T18:24:45.069583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = pd.DataFrame.from_dict(history)\nhistory.to_csv(\"history.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:24:48.657317Z","iopub.execute_input":"2023-10-08T18:24:48.657974Z","iopub.status.idle":"2023-10-08T18:24:48.688066Z","shell.execute_reply.started":"2023-10-08T18:24:48.657932Z","shell.execute_reply":"2023-10-08T18:24:48.686787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Logs</h1></span>","metadata":{}},{"cell_type":"code","source":"plt.plot( range(history.shape[0]), history[\"Train Loss\"].values, label=\"Train Loss\")\nplt.plot( range(history.shape[0]), history[\"Valid Loss\"].values, label=\"Valid Loss\")\nplt.xlabel(\"epochs\")\nplt.ylabel(\"Loss\")\nplt.grid()\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:24:51.307678Z","iopub.execute_input":"2023-10-08T18:24:51.308331Z","iopub.status.idle":"2023-10-08T18:24:51.338486Z","shell.execute_reply.started":"2023-10-08T18:24:51.308298Z","shell.execute_reply":"2023-10-08T18:24:51.336997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot( range(history.shape[0]), history[\"Train Accuracy\"].values, label=\"Train Accuracy\")\nplt.plot( range(history.shape[0]), history[\"Valid Accuracy\"].values, label=\"Valid Accuracy\")\nplt.xlabel(\"epochs\")\nplt.ylabel(\"Accuracy\")\nplt.grid()\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:24:52.086048Z","iopub.execute_input":"2023-10-08T18:24:52.086691Z","iopub.status.idle":"2023-10-08T18:24:52.115594Z","shell.execute_reply.started":"2023-10-08T18:24:52.086661Z","shell.execute_reply":"2023-10-08T18:24:52.114435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot( range(history.shape[0]), history[\"lr\"].values, label=\"lr\")\nplt.xlabel(\"epochs\")\nplt.ylabel(\"lr\")\nplt.grid()\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:24:52.836704Z","iopub.execute_input":"2023-10-08T18:24:52.837378Z","iopub.status.idle":"2023-10-08T18:24:52.86597Z","shell.execute_reply.started":"2023-10-08T18:24:52.837346Z","shell.execute_reply":"2023-10-08T18:24:52.864412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_test_file_path(image_id):\n    if os.path.exists(f\"{TEST_DIR}/{image_id}_thumbnail.png\"):\n        return f\"{TEST_DIR}/{image_id}_thumbnail.png\"\n    else:\n        return f\"{ALT_TEST_DIR}/{image_id}.png\"","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:24:53.663356Z","iopub.execute_input":"2023-10-08T18:24:53.663697Z","iopub.status.idle":"2023-10-08T18:24:53.668401Z","shell.execute_reply.started":"2023-10-08T18:24:53.663661Z","shell.execute_reply":"2023-10-08T18:24:53.667447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(f\"{ROOT_DIR}/test.csv\")\ndf['file_path'] = df['image_id'].apply(get_test_file_path)\ndf['label'] = 0 # dummy\ndf","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:24:55.237287Z","iopub.execute_input":"2023-10-08T18:24:55.238215Z","iopub.status.idle":"2023-10-08T18:24:55.259477Z","shell.execute_reply.started":"2023-10-08T18:24:55.238175Z","shell.execute_reply":"2023-10-08T18:24:55.258516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub = pd.read_csv(f\"{ROOT_DIR}/sample_submission.csv\")\ndf_sub","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:24:56.356265Z","iopub.execute_input":"2023-10-08T18:24:56.35725Z","iopub.status.idle":"2023-10-08T18:24:56.372726Z","shell.execute_reply.started":"2023-10-08T18:24:56.357206Z","shell.execute_reply":"2023-10-08T18:24:56.371734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"LABEL_ENCODER_BIN = \"/kaggle/working/label_encoder.pkl\"","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:25:11.469872Z","iopub.execute_input":"2023-10-08T18:25:11.470243Z","iopub.status.idle":"2023-10-08T18:25:11.474476Z","shell.execute_reply.started":"2023-10-08T18:25:11.470216Z","shell.execute_reply":"2023-10-08T18:25:11.473577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder = joblib.load( LABEL_ENCODER_BIN )","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:25:13.295014Z","iopub.execute_input":"2023-10-08T18:25:13.295333Z","iopub.status.idle":"2023-10-08T18:25:13.301026Z","shell.execute_reply.started":"2023-10-08T18:25:13.295305Z","shell.execute_reply":"2023-10-08T18:25:13.300048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = UBCDataset(df, transforms=data_transforms[\"valid\"])\ntest_loader = DataLoader(test_dataset, batch_size=CONFIG['valid_batch_size'], \n                          num_workers=2, shuffle=False, pin_memory=True)","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:25:14.445508Z","iopub.execute_input":"2023-10-08T18:25:14.446463Z","iopub.status.idle":"2023-10-08T18:25:14.452023Z","shell.execute_reply.started":"2023-10-08T18:25:14.446421Z","shell.execute_reply":"2023-10-08T18:25:14.451166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = []\nwith torch.no_grad():\n    bar = tqdm(enumerate(test_loader), total=len(test_loader))\n    for step, data in bar:        \n        images = data['image'].to(CONFIG[\"device\"], dtype=torch.float)        \n        batch_size = images.size(0)\n        outputs = model(images)\n        _, predicted = torch.max(model.softmax(outputs), 1)\n        preds.append( predicted.detach().cpu().numpy() )\npreds = np.concatenate(preds).flatten()\npred_labels = encoder.inverse_transform( preds )","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:25:15.7511Z","iopub.execute_input":"2023-10-08T18:25:15.751426Z","iopub.status.idle":"2023-10-08T18:25:16.182422Z","shell.execute_reply.started":"2023-10-08T18:25:15.7514Z","shell.execute_reply":"2023-10-08T18:25:16.18132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub[\"label\"] = pred_labels\ndf_sub.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:25:17.324349Z","iopub.execute_input":"2023-10-08T18:25:17.324701Z","iopub.status.idle":"2023-10-08T18:25:17.335001Z","shell.execute_reply.started":"2023-10-08T18:25:17.324673Z","shell.execute_reply":"2023-10-08T18:25:17.334034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub","metadata":{"execution":{"iopub.status.busy":"2023-10-08T18:25:18.184376Z","iopub.execute_input":"2023-10-08T18:25:18.184704Z","iopub.status.idle":"2023-10-08T18:25:18.193571Z","shell.execute_reply.started":"2023-10-08T18:25:18.184677Z","shell.execute_reply":"2023-10-08T18:25:18.192512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}