{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":145619039,"sourceType":"kernelVersion"},{"sourceId":155318053,"sourceType":"kernelVersion"},{"sourceId":3729,"sourceType":"modelInstanceVersion","modelInstanceId":2656}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"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-12-15T11:42:16.757554Z","iopub.execute_input":"2023-12-15T11:42:16.757992Z","iopub.status.idle":"2023-12-15T11:42:26.912072Z","shell.execute_reply.started":"2023-12-15T11:42:16.757954Z","shell.execute_reply":"2023-12-15T11:42:26.910803Z"},"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    \"img_size\": 512,\n    \"model_name\": \"tf_efficientnet_b0_ns\",\n    \"num_classes\": 8,\n    \"valid_batch_size\": 64,\n    \"device\": torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\"),\n}","metadata":{"execution":{"iopub.status.busy":"2023-12-15T11:42:56.785396Z","iopub.execute_input":"2023-12-15T11:42:56.786351Z","iopub.status.idle":"2023-12-15T11:42:56.795149Z","shell.execute_reply.started":"2023-12-15T11:42:56.7863Z","shell.execute_reply":"2023-12-15T11:42:56.793726Z"},"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-12-15T11:42:59.91435Z","iopub.execute_input":"2023-12-15T11:42:59.915548Z","iopub.status.idle":"2023-12-15T11:42:59.929533Z","shell.execute_reply.started":"2023-12-15T11:42:59.915476Z","shell.execute_reply":"2023-12-15T11:42:59.928111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ROOT_DIR = '/kaggle/input/UBC-OCEAN'\n# TEST_DIR = '/kaggle/input/UBC-OCEAN/test_thumbnails'\n\n# LABEL_ENCODER_BIN = \"/kaggle/input/ubc-pytorch-cnn-training-fold1of5/label_encoder.pkl\"\n# BEST_WEIGHT = \"/kaggle/input/ubc-pytorch-cnn-training-fold1of5/Acc0.69_Loss0.9592_epoch20.bin\"","metadata":{"execution":{"iopub.status.busy":"2023-12-15T11:43:04.944456Z","iopub.execute_input":"2023-12-15T11:43:04.944953Z","iopub.status.idle":"2023-12-15T11:43:04.951211Z","shell.execute_reply.started":"2023-12-15T11:43:04.944918Z","shell.execute_reply":"2023-12-15T11:43:04.949548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT_DIR = '/kaggle/input/UBC-OCEAN'\nTEST_DIR = '/kaggle/input/UBC-OCEAN/test_thumbnails'\n\nLABEL_ENCODER_BIN = \"/kaggle/input/ubcpytorchcnn-trainingfold1of5/label_encoder.pkl\"\nBEST_WEIGHT = \"/kaggle/input/ubcpytorchcnn-trainingfold1of5/Acc0.62_Loss1.0335_epoch10.bin\"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_test_file_path(image_id):\n    return f\"{TEST_DIR}/{image_id}_thumbnail.png\"","metadata":{"execution":{"iopub.status.busy":"2023-12-15T11:43:14.52499Z","iopub.execute_input":"2023-12-15T11:43:14.525435Z","iopub.status.idle":"2023-12-15T11:43:14.531367Z","shell.execute_reply.started":"2023-12-15T11:43:14.525402Z","shell.execute_reply":"2023-12-15T11:43:14.530181Z"},"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":"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-12-15T11:43:16.994675Z","iopub.execute_input":"2023-12-15T11:43:16.995131Z","iopub.status.idle":"2023-12-15T11:43:17.037769Z","shell.execute_reply.started":"2023-12-15T11:43:16.995097Z","shell.execute_reply":"2023-12-15T11:43:17.036558Z"},"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-12-15T11:43:20.684539Z","iopub.execute_input":"2023-12-15T11:43:20.68499Z","iopub.status.idle":"2023-12-15T11:43:20.701883Z","shell.execute_reply.started":"2023-12-15T11:43:20.684958Z","shell.execute_reply":"2023-12-15T11:43:20.700582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder = joblib.load( LABEL_ENCODER_BIN )","metadata":{"execution":{"iopub.status.busy":"2023-12-15T11:43:37.149746Z","iopub.execute_input":"2023-12-15T11:43:37.150202Z","iopub.status.idle":"2023-12-15T11:43:37.162672Z","shell.execute_reply.started":"2023-12-15T11:43:37.150168Z","shell.execute_reply":"2023-12-15T11:43:37.161743Z"},"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-12-15T11:43:40.344812Z","iopub.execute_input":"2023-12-15T11:43:40.345285Z","iopub.status.idle":"2023-12-15T11:43:40.355899Z","shell.execute_reply.started":"2023-12-15T11:43:40.345247Z","shell.execute_reply":"2023-12-15T11:43:40.354273Z"},"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    \"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-12-15T11:43:45.184624Z","iopub.execute_input":"2023-12-15T11:43:45.185107Z","iopub.status.idle":"2023-12-15T11:43:45.193436Z","shell.execute_reply.started":"2023-12-15T11:43:45.185069Z","shell.execute_reply":"2023-12-15T11:43:45.192094Z"},"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-12-15T11:43:48.724427Z","iopub.execute_input":"2023-12-15T11:43:48.724913Z","iopub.status.idle":"2023-12-15T11:43:48.735891Z","shell.execute_reply.started":"2023-12-15T11:43:48.724879Z","shell.execute_reply":"2023-12-15T11:43:48.734542Z"},"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=False, checkpoint_path=None):\n#         super(UBCModel, self).__init__()\n#         self.model = timm.create_model(model_name, pretrained=pretrained)\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    \n# model = UBCModel(CONFIG['model_name'], CONFIG['num_classes'])\n# model.load_state_dict(torch.load( BEST_WEIGHT ))\n# model.to(CONFIG['device']);","metadata":{"execution":{"iopub.status.busy":"2023-10-07T18:51:46.11481Z","iopub.execute_input":"2023-10-07T18:51:46.115083Z","iopub.status.idle":"2023-10-07T18:51:46.301446Z","shell.execute_reply.started":"2023-10-07T18:51:46.115059Z","shell.execute_reply":"2023-10-07T18:51:46.300512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport timm\n\nclass UBCModel(nn.Module):\n    def __init__(self, model_name, num_classes, pretrained=False, checkpoint_path=None):\n        super(UBCModel, self).__init__()\n        \n        # base model using timm\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n\n        # Update classifier and global_pool layers\n        in_features = self.model.classifier.in_features\n        self.model.classifier = nn.Identity()\n        self.model.global_pool = nn.Identity()\n        \n        # custom pooling \n        self.pooling = GeM()\n        \n        # Linear layer for classification\n        self.linear = nn.Linear(in_features, num_classes)\n        \n        # Softmax for output probabilities\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\nmodel = UBCModel(CONFIG['model_name'], CONFIG['num_classes'])\n\n# pre-trained weights, explicitly mapping to CPU if necessary\nif torch.cuda.is_available():\n    model.load_state_dict(torch.load(BEST_WEIGHT, map_location=torch.device('cuda')))\nelse:\n    model.load_state_dict(torch.load(BEST_WEIGHT, map_location=torch.device('cpu')))\n\n# model to the specified device\nmodel.to(CONFIG['device'])","metadata":{"execution":{"iopub.status.busy":"2023-12-15T13:43:13.035048Z","iopub.execute_input":"2023-12-15T13:43:13.035369Z","iopub.status.idle":"2023-12-15T13:43:13.039884Z","shell.execute_reply.started":"2023-12-15T13:43:13.035346Z","shell.execute_reply":"2023-12-15T13:43:13.039295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Prepare Dataloaders</span>","metadata":{}},{"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-12-15T11:51:37.724329Z","iopub.execute_input":"2023-12-15T11:51:37.724777Z","iopub.status.idle":"2023-12-15T11:51:37.732054Z","shell.execute_reply.started":"2023-12-15T11:51:37.724741Z","shell.execute_reply":"2023-12-15T11:51:37.730591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Start Inference</span>","metadata":{}},{"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-12-15T11:51:40.907933Z","iopub.execute_input":"2023-12-15T11:51:40.908394Z","iopub.status.idle":"2023-12-15T11:51:41.888809Z","shell.execute_reply.started":"2023-12-15T11:51:40.908359Z","shell.execute_reply":"2023-12-15T11:51:41.887275Z"},"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-12-15T11:51:46.440192Z","iopub.execute_input":"2023-12-15T11:51:46.441016Z","iopub.status.idle":"2023-12-15T11:51:46.452508Z","shell.execute_reply.started":"2023-12-15T11:51:46.440962Z","shell.execute_reply":"2023-12-15T11:51:46.451569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub","metadata":{"execution":{"iopub.status.busy":"2023-12-15T11:51:49.514627Z","iopub.execute_input":"2023-12-15T11:51:49.515348Z","iopub.status.idle":"2023-12-15T11:51:49.528649Z","shell.execute_reply.started":"2023-12-15T11:51:49.5153Z","shell.execute_reply":"2023-12-15T11:51:49.527333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}